diff --git a/docs/source/policy_diffusion_README.md b/docs/source/policy_diffusion_README.md index 9ec934add..ab6cf134a 100644 --- a/docs/source/policy_diffusion_README.md +++ b/docs/source/policy_diffusion_README.md @@ -2,6 +2,25 @@ https://diffusion-policy.cs.columbia.edu +## Training + +The reference implementation maintains an exponential moving average (EMA) of the policy weights during training and evaluates the EMA weights. To reproduce this behavior, enable the trainer's EMA shadow: + +```bash +lerobot-train \ + --policy.type=diffusion \ + --ema.enable=true \ + ... +``` + +Checkpoints then contain a directly loadable copy of the EMA weights next to the live ones, e.g. for evaluation: + +```bash +lerobot-eval --policy.path=outputs/train/.../checkpoints/last/pretrained_model_ema ... +``` + +The EMA decay schedule (`--ema.inv_gamma`, `--ema.power`, ...) defaults to the reference implementation's values. For a constant decay instead of the warmup schedule (e.g. to match openpi's pi0/pi05 training), set `--ema.decay=0.99`. + ## Citation ```bibtex diff --git a/docs/source/policy_pi05_README.md b/docs/source/policy_pi05_README.md index 9abec99fa..800ed9071 100644 --- a/docs/source/policy_pi05_README.md +++ b/docs/source/policy_pi05_README.md @@ -59,6 +59,22 @@ When `use_relative_actions=true`, the training script automatically: --- +## EMA of the policy weights + +OpenPI maintains an exponential moving average of the weights during training (`ema_decay=0.99` by default) and keeps the EMA copy for inference. To reproduce this with the LeRobot trainer, enable the EMA shadow with a constant decay: + +```bash +python -m lerobot.scripts.lerobot_train \ + --policy.type=pi05 \ + --dataset.repo_id=your_org/your_dataset \ + --ema.enable=true \ + --ema.decay=0.99 +``` + +Checkpoints then contain a directly loadable copy of the EMA weights in `pretrained_model_ema/` next to the live ones. Note that the shadow is a full extra copy of the parameters on the GPU. Like OpenPI (which disables EMA in its LoRA configs), EMA is not supported together with PEFT adapters. + +--- + ## Citation If you use this work, please cite both **OpenPI** and the π₀.₅ paper: diff --git a/src/lerobot/configs/__init__.py b/src/lerobot/configs/__init__.py index c32e3368b..a3b2190bb 100644 --- a/src/lerobot/configs/__init__.py +++ b/src/lerobot/configs/__init__.py @@ -22,7 +22,7 @@ Import them directly: ``from lerobot.configs.train import TrainPipelineConfig`` """ from .dataset import DatasetRecordConfig -from .default import DatasetConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig +from .default import DatasetConfig, EMAConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig from .policies import PreTrainedConfig from .recipe import MessageTurn, TrainingRecipe, load_recipe from .types import ( @@ -57,6 +57,7 @@ __all__ = [ # Config classes "DatasetRecordConfig", "DatasetConfig", + "EMAConfig", "EvalConfig", "JobConfig", "MessageTurn", diff --git a/src/lerobot/configs/default.py b/src/lerobot/configs/default.py index 6cbda669a..1fbc969b6 100644 --- a/src/lerobot/configs/default.py +++ b/src/lerobot/configs/default.py @@ -139,6 +139,59 @@ class EvalConfig: return min(by_cpu, self.n_episodes, 64) +@dataclass +class EMAConfig: + """Exponential moving average (EMA) of the policy weights. + + Standard practice for diffusion-style policies (Chi et al. 2023, "Diffusion Policy", section V.D): + the reference implementation enables it in every config and evaluates the EMA weights. Off by + default here because it keeps a second full copy of the parameters in memory. + + The decay follows the warmup schedule from diffusers' `EMAModel`: + `decay_t = 1 - (1 + t / inv_gamma) ** -power`, clamped to `[min_decay, max_decay]`. + The defaults mirror the reference implementation. Alternatively, set `decay` for a constant + decay at every step, as used by openpi for pi0/pi05 (`ema_decay=0.99`). + """ + + enable: bool = False + # Constant decay coefficient (openpi-style, e.g. 0.99 for pi0/pi05). When set, the warmup + # schedule below is bypassed and the shadow uses this decay at every step. + decay: float | None = None + # Number of optimizer steps during which the shadow stays a hard copy of the live weights. + update_after_step: int = 0 + # Warmup schedule parameters (see class docstring). + inv_gamma: float = 1.0 + power: float = 0.75 + min_decay: float = 0.0 + max_decay: float = 0.9999 + # Evaluate the EMA weights (instead of the live ones) during periodic env eval. + # Offline eval-loss (--eval_steps) always uses the live weights: it runs on every rank + # while the EMA shadow only lives on the main process. + use_for_eval: bool = True + + def __post_init__(self) -> None: + if not (0.0 <= self.min_decay <= self.max_decay <= 1.0): + raise ValueError( + "Expected 0 <= ema.min_decay <= ema.max_decay <= 1, got " + f"min_decay={self.min_decay} and max_decay={self.max_decay}." + ) + if self.inv_gamma <= 0: + raise ValueError(f"ema.inv_gamma must be positive, got {self.inv_gamma}.") + if self.power <= 0: + raise ValueError(f"ema.power must be positive, got {self.power}.") + if self.update_after_step < 0: + raise ValueError(f"ema.update_after_step must be >= 0, got {self.update_after_step}.") + if self.decay is not None: + if not 0.0 <= self.decay <= 1.0: + raise ValueError(f"ema.decay must be in [0, 1], got {self.decay}.") + # Keep the literals in sync with the field defaults above. + if self.min_decay != 0.0 or self.max_decay != 0.9999: + raise ValueError( + "ema.decay (constant decay) and ema.min_decay/ema.max_decay (schedule clamp) are " + "mutually exclusive: set one or the other." + ) + + @dataclass class PeftConfig: # PEFT offers many fine-tuning methods, layer adapters being the most common and currently also the most diff --git a/src/lerobot/configs/train.py b/src/lerobot/configs/train.py index cf9b7817e..9381a0370 100644 --- a/src/lerobot/configs/train.py +++ b/src/lerobot/configs/train.py @@ -35,7 +35,7 @@ from lerobot.utils.hub import HubMixin, find_latest_hub_checkpoint from lerobot.utils.sample_weighting import SampleWeightingConfig from . import parser -from .default import DatasetConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig +from .default import DatasetConfig, EMAConfig, EvalConfig, JobConfig, PeftConfig, WandBConfig from .policies import PreTrainedConfig from .rewards import RewardModelConfig @@ -163,6 +163,8 @@ class TrainPipelineConfig(HubMixin): # FSDP/DDP tuning knobs, compile & activation-checkpointing placeholders. accelerator: AcceleratorConfig = field(default_factory=AcceleratorConfig) eval: EvalConfig = field(default_factory=EvalConfig) + # Maintain an EMA shadow of the policy weights during training (see EMAConfig). + ema: EMAConfig = field(default_factory=EMAConfig) wandb: WandBConfig = field(default_factory=WandBConfig) peft: PeftConfig | None = None diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index 8a4cefd35..98601d303 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -76,7 +76,7 @@ from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_proces from lerobot.policies.factory import ProcessorConfigKwargs from lerobot.rewards import make_reward_pre_post_processors from lerobot.utils.collate import lerobot_collate_fn -from lerobot.utils.constants import TRAINING_STATE_DIR +from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR from lerobot.utils.import_utils import _peft_available, register_third_party_plugins, require_package from lerobot.utils.logging_utils import AverageMeter, MetricsTracker from lerobot.utils.random_utils import set_seed @@ -95,6 +95,20 @@ else: from .lerobot_eval import eval_policy_all +EMA_STATE_FILENAME = "ema_state.pt" + + +@contextmanager +def _ema_weights(ema: Any, policy: PreTrainedPolicy) -> Iterator[None]: + """Temporarily swap the EMA shadow weights into `policy`, restoring the live ones on exit.""" + params = list(policy.parameters()) + ema.store(params) + ema.copy_to(params) + try: + yield + finally: + ema.restore(params) + @contextmanager def _make_eval_envs(cfg: TrainPipelineConfig) -> Iterator[dict[str, dict[int, Any]]]: @@ -592,6 +606,65 @@ def train(cfg: TrainPipelineConfig): dl_iter = cycle(dataloader) policy.train() + # EMA shadow of the policy weights (Chi et al. 2023, Diffusion Policy, section V.D). The shadow + # lives on the main process only, which is safe under DDP where every rank holds identical + # weights after each gradient sync. diffusers is imported lazily so the base training path does + # not depend on it. + ema = None + if cfg.ema.enable: + if parallel_dims.is_sharded: + raise NotImplementedError( + "--ema.enable=true is not supported with sharded training (FSDP2/HSDP/CP): the " + "parameters are sharded across ranks. Use a replicated (DDP) or single-GPU run." + ) + if cfg.peft is not None: + raise NotImplementedError("--ema.enable=true is not supported together with PEFT adapters.") + require_package("diffusers", extra="diffusion") + if is_main_process(): + from diffusers.training_utils import EMAModel # noqa: PLC0415 + + # A constant --ema.decay is expressed through the schedule clamp: with + # min_decay == max_decay, the warmup curve is pinned to that value at every step. + min_decay = cfg.ema.min_decay if cfg.ema.decay is None else cfg.ema.decay + max_decay = cfg.ema.max_decay if cfg.ema.decay is None else cfg.ema.decay + ema = EMAModel( + accelerator.unwrap_model(policy).parameters(), + decay=max_decay, + min_decay=min_decay, + update_after_step=cfg.ema.update_after_step, + use_ema_warmup=True, + inv_gamma=cfg.ema.inv_gamma, + power=cfg.ema.power, + ) + ema.to(device) + if cfg.ema.decay is not None: + logging.info( + "EMA enabled: decay=%g (constant), update_after_step=%d, use_for_eval=%s", + cfg.ema.decay, + cfg.ema.update_after_step, + cfg.ema.use_for_eval, + ) + else: + logging.info( + "EMA enabled: max_decay=%g, inv_gamma=%g, power=%g, update_after_step=%d, use_for_eval=%s", + cfg.ema.max_decay, + cfg.ema.inv_gamma, + cfg.ema.power, + cfg.ema.update_after_step, + cfg.ema.use_for_eval, + ) + if cfg.checkpoint_path is not None: + ema_path = cfg.checkpoint_path / TRAINING_STATE_DIR / EMA_STATE_FILENAME + if ema_path.exists(): + ema.load_state_dict(torch.load(ema_path, map_location=device, weights_only=True)) + logging.info("Resumed EMA shadow from %s", ema_path) + else: + logging.warning( + "Resuming with --ema.enable=true but %s is missing; " + "restarting the shadow from the current weights.", + ema_path, + ) + train_metrics = { # Per-rank loss reflects only one shard of the global batch; mean recovers the loss the # data-parallel group is actually optimizing. grad_norm and lr are already identical on @@ -657,6 +730,12 @@ def train(cfg: TrainPipelineConfig): ) train_tracker.step_s = time.perf_counter() - step_start + # Pull one optimizer step of the live weights into the EMA shadow (main process only). + # The shadow tracks optimizer updates, not micro-batches: gate on the sync step under + # gradient accumulation. + if ema is not None and accelerator.sync_gradients: + ema.step(accelerator.unwrap_model(policy).parameters()) + # Note: eval and checkpoint happens *after* the `step`th training update has completed, so we # increment `step` here. step += 1 @@ -684,6 +763,9 @@ def train(cfg: TrainPipelineConfig): if sample_weighter is not None: weighter_stats = sample_weighter.get_stats() wandb_log_dict.update({f"sample_weighting/{k}": v for k, v in weighter_stats.items()}) + if ema is not None and ema.cur_decay_value is not None: + wandb_log_dict["ema/decay"] = ema.cur_decay_value + wandb_log_dict["ema/step"] = ema.optimization_step wandb_logger.log_dict(wandb_log_dict, step) train_tracker.reset_averages() @@ -728,6 +810,17 @@ def train(cfg: TrainPipelineConfig): accelerator=accelerator, ) if is_main_process(): + if ema is not None: + # Save the shadow for exact resume, plus a directly loadable copy of the EMA + # weights (lerobot-eval --policy.path=/pretrained_model_ema). + torch.save(ema.state_dict(), checkpoint_dir / TRAINING_STATE_DIR / EMA_STATE_FILENAME) + unwrapped_policy = accelerator.unwrap_model(policy) + ema_dir = checkpoint_dir / f"{PRETRAINED_MODEL_DIR}_ema" + with _ema_weights(ema, unwrapped_policy): + unwrapped_policy.save_pretrained(ema_dir) + cfg.save_pretrained(ema_dir) + preprocessor.save_pretrained(ema_dir) + postprocessor.save_pretrained(ema_dir) update_last_checkpoint(checkpoint_dir) if cfg.save_checkpoint_to_hub: push_checkpoint_to_hub( @@ -743,10 +836,18 @@ def train(cfg: TrainPipelineConfig): if is_main_process(): step_id = get_step_identifier(step, cfg.steps) logging.info(f"Eval policy at step {step}") - with _make_eval_envs(cfg) as eval_env, torch.no_grad(), accelerator.autocast(): + eval_policy_model = accelerator.unwrap_model(policy) + # Evaluate the EMA weights when enabled: the swap happens only on the main + # process (the other ranks wait at the barrier below) and is exactly undone + # afterwards, so the live weights stay in sync across ranks. + use_ema_for_eval = ema is not None and cfg.ema.use_for_eval + if use_ema_for_eval: + logging.info("Evaluating the EMA weights") + weights_cm = _ema_weights(ema, eval_policy_model) if use_ema_for_eval else nullcontext() + with weights_cm, _make_eval_envs(cfg) as eval_env, torch.no_grad(), accelerator.autocast(): eval_info = eval_policy_all( envs=eval_env, # dict[suite][task_id] -> vec_env - policy=accelerator.unwrap_model(policy), + policy=eval_policy_model, env_preprocessor=env_preprocessor, env_postprocessor=env_postprocessor, preprocessor=preprocessor, @@ -805,6 +906,25 @@ def train(cfg: TrainPipelineConfig): peft_model=unwrapped if peft_model is not None else None, ) + # The push above ships the live weights; when EMA is on, the weights that were + # evaluated are the shadow, so push those too under a sibling `-ema` repo. + # The shadow lives on the main process only, so this is rank-0-only by construction. + # Non-fatal: the live model is already up if this fails. + if ema is not None: + ema_repo_id = f"{active_cfg.repo_id}-ema" + orig_repo_id = unwrapped.config.repo_id + try: + unwrapped.config.repo_id = ema_repo_id + with _ema_weights(ema, unwrapped): + unwrapped.push_model_to_hub(cfg, dataset_meta=dataset.meta) + preprocessor.push_to_hub(ema_repo_id) + postprocessor.push_to_hub(ema_repo_id) + logging.info("Pushed EMA weights to %s", ema_repo_id) + except Exception as exc: # noqa: BLE001 + logging.warning("Failed to push EMA weights to %s: %s", ema_repo_id, exc) + finally: + unwrapped.config.repo_id = orig_repo_id + # Properly clean up the distributed process group accelerator.wait_for_everyone() accelerator.end_training() diff --git a/tests/training/test_ema.py b/tests/training/test_ema.py new file mode 100644 index 000000000..7b5566be6 --- /dev/null +++ b/tests/training/test_ema.py @@ -0,0 +1,314 @@ +#!/usr/bin/env python + +# Copyright 2026 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for the opt-in EMA shadow maintained by the training pipeline (--ema.enable=true).""" + +import draccus +import numpy as np +import pytest +import torch + +pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])") + +from lerobot.configs.default import EMAConfig +from lerobot.configs.train import TrainPipelineConfig +from lerobot.datasets.lerobot_dataset import LeRobotDataset +from lerobot.utils.constants import PRETRAINED_MODEL_DIR, TRAINING_STATE_DIR + +DUMMY_REPO_ID = "dummy/repo" +DUMMY_STATE_DIM = 6 +DUMMY_ACTION_DIM = 6 +IMAGE_SIZE = 32 +N_EPISODES = 2 +EPISODE_LENGTH = 12 + + +def test_ema_config_defaults_match_the_reference(): + cfg = EMAConfig() + assert not cfg.enable + assert cfg.inv_gamma == 1.0 + assert cfg.power == 0.75 + assert cfg.update_after_step == 0 + + +@pytest.mark.parametrize( + "kwargs", + [ + {"min_decay": 0.5, "max_decay": 0.1}, + {"max_decay": 1.5}, + {"min_decay": -0.1}, + {"inv_gamma": 0.0}, + {"power": -1.0}, + {"update_after_step": -1}, + {"decay": 1.5}, + {"decay": -0.1}, + {"decay": 0.99, "min_decay": 0.5}, + {"decay": 0.99, "max_decay": 0.9}, + ], +) +def test_ema_config_rejects_invalid_values(kwargs): + with pytest.raises(ValueError): + EMAConfig(**kwargs) + + +def test_ema_config_cli_parsing(): + cfg = draccus.parse( + TrainPipelineConfig, + None, + args=[ + f"--dataset.repo_id={DUMMY_REPO_ID}", + "--ema.enable=true", + "--ema.power=0.8", + "--ema.update_after_step=10", + ], + ) + assert cfg.ema.enable + assert cfg.ema.power == 0.8 + assert cfg.ema.update_after_step == 10 + + +def test_ema_config_cli_parsing_constant_decay(): + cfg = draccus.parse( + TrainPipelineConfig, + None, + args=[ + f"--dataset.repo_id={DUMMY_REPO_ID}", + "--ema.enable=true", + "--ema.decay=0.99", + ], + ) + assert cfg.ema.enable + assert cfg.ema.decay == 0.99 + + +def test_ema_constant_decay_pins_the_schedule(): + """min_decay == max_decay clamps the warmup curve to a constant (how --ema.decay is implemented).""" + pytest.importorskip("diffusers") + from diffusers.training_utils import EMAModel + + model = torch.nn.Linear(4, 4) + ema = EMAModel( + model.parameters(), decay=0.99, min_decay=0.99, use_ema_warmup=True, inv_gamma=1.0, power=0.75 + ) + # The first update is a hard copy (decay 0); every one after uses the constant decay. + for step in range(1, 6): + ema.step(model.parameters()) + if step > 1: + assert ema.cur_decay_value == 0.99 + + +def test_ema_weights_context_swaps_and_restores(): + pytest.importorskip("diffusers") + from diffusers.training_utils import EMAModel + + from lerobot.scripts.lerobot_train import _ema_weights + + torch.manual_seed(0) + model = torch.nn.Linear(4, 4) + ema = EMAModel(model.parameters(), decay=0.9999, use_ema_warmup=True, inv_gamma=1.0, power=0.75) + + optimizer = torch.optim.SGD(model.parameters(), lr=0.1) + for _ in range(3): + model(torch.randn(2, 4)).sum().backward() + optimizer.step() + optimizer.zero_grad() + ema.step(model.parameters()) + + live = [p.detach().clone() for p in model.parameters()] + with _ema_weights(ema, model): + swapped = [p.detach().clone() for p in model.parameters()] + restored = list(model.parameters()) + + assert any(not torch.equal(a, b) for a, b in zip(live, swapped, strict=True)) + assert all(torch.equal(a, b.detach()) for a, b in zip(live, restored, strict=True)) + + +def make_dummy_dataset(tmp_path): + features = { + "action": {"dtype": "float32", "shape": (DUMMY_ACTION_DIM,), "names": None}, + "observation.state": {"dtype": "float32", "shape": (DUMMY_STATE_DIM,), "names": None}, + "observation.images.top": { + "dtype": "image", + "shape": (IMAGE_SIZE, IMAGE_SIZE, 3), + "names": ["height", "width", "channel"], + }, + } + root = tmp_path / "_dataset" + dataset = LeRobotDataset.create(repo_id=DUMMY_REPO_ID, fps=30, features=features, root=root) + rng = np.random.default_rng(0) + for ep_idx in range(N_EPISODES): + for _ in range(EPISODE_LENGTH): + dataset.add_frame( + { + "action": rng.standard_normal(DUMMY_ACTION_DIM).astype(np.float32), + "observation.state": rng.standard_normal(DUMMY_STATE_DIM).astype(np.float32), + "observation.images.top": rng.integers( + 0, 255, size=(IMAGE_SIZE, IMAGE_SIZE, 3), dtype=np.uint8 + ), + "task": f"task_{ep_idx}", + } + ) + dataset.save_episode() + dataset.finalize() + return root + + +def make_train_config(root, output_dir, steps, ema_enable, ema_decay=None): + from lerobot.configs.default import DatasetConfig + from lerobot.policies.factory import make_policy_config + + policy_config = make_policy_config( + "diffusion", + device="cpu", + push_to_hub=False, + n_obs_steps=2, + horizon=8, + n_action_steps=4, + drop_n_last_frames=0, + down_dims=(32, 64), + diffusion_step_embed_dim=32, + spatial_softmax_num_keypoints=8, + num_inference_steps=2, + pretrained_backbone_weights=None, + use_group_norm=True, + ) + cfg = TrainPipelineConfig( + dataset=DatasetConfig(repo_id=DUMMY_REPO_ID, root=str(root)), + policy=policy_config, + output_dir=output_dir, + steps=steps, + batch_size=2, + num_workers=0, + seed=42, + log_freq=0, + env_eval_freq=0, + save_freq=2, + ema=EMAConfig(enable=ema_enable, decay=ema_decay), + ) + cfg.optimizer = policy_config.get_optimizer_preset() + cfg.scheduler = policy_config.get_scheduler_preset() + # The config is built in-process, so skip the CLI-oriented validation. + cfg.validate = lambda: None + return cfg + + +def load_safetensors(path): + from safetensors.torch import load_file + + return load_file(path) + + +def test_train_diffusion_with_ema_checkpoint_and_resume(tmp_path): + pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])") + pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])") + from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train + + root = make_dummy_dataset(tmp_path) + output_dir = tmp_path / "_output" + + cfg = make_train_config(root, output_dir, steps=4, ema_enable=True) + train(cfg) + + checkpoint_dir = output_dir / "checkpoints" / "000004" + ema_state_path = checkpoint_dir / TRAINING_STATE_DIR / EMA_STATE_FILENAME + ema_model_dir = checkpoint_dir / f"{PRETRAINED_MODEL_DIR}_ema" + + # The shadow state is saved for resume and tracks every optimizer step. + assert ema_state_path.exists() + ema_state = torch.load(ema_state_path, weights_only=True) + assert ema_state["optimization_step"] == 4 + + # A directly loadable EMA model is saved next to the live one, with different weights. + live_weights = load_safetensors(checkpoint_dir / PRETRAINED_MODEL_DIR / "model.safetensors") + ema_weights = load_safetensors(ema_model_dir / "model.safetensors") + assert set(live_weights) == set(ema_weights) + assert any(not torch.equal(live_weights[k], ema_weights[k]) for k in live_weights) + + from lerobot.policies.diffusion.modeling_diffusion import DiffusionPolicy + + policy = DiffusionPolicy.from_pretrained(str(ema_model_dir)) + assert isinstance(policy, DiffusionPolicy) + + # Resuming picks the shadow up where it left off instead of restarting it. + resume_cfg = make_train_config(root, output_dir, steps=6, ema_enable=True) + resume_cfg.resume = True + resume_cfg.checkpoint_path = checkpoint_dir + train(resume_cfg) + + resumed_state = torch.load( + output_dir / "checkpoints" / "000006" / TRAINING_STATE_DIR / EMA_STATE_FILENAME, + weights_only=True, + ) + assert resumed_state["optimization_step"] == 6 + + +def test_train_with_constant_ema_decay(tmp_path): + pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])") + pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])") + from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train + + root = make_dummy_dataset(tmp_path) + output_dir = tmp_path / "_output" + + cfg = make_train_config(root, output_dir, steps=2, ema_enable=True, ema_decay=0.99) + train(cfg) + + ema_state = torch.load( + output_dir / "checkpoints" / "000002" / TRAINING_STATE_DIR / EMA_STATE_FILENAME, + weights_only=True, + ) + # The constant decay is implemented by pinning the schedule clamp to that value. + assert ema_state["decay"] == 0.99 + assert ema_state["min_decay"] == 0.99 + assert ema_state["optimization_step"] == 2 + + +def test_train_with_ema_and_gradient_accumulation(tmp_path): + """The shadow tracks optimizer steps, not micro-batches, under gradient accumulation.""" + pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])") + pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])") + from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train + + root = make_dummy_dataset(tmp_path) + output_dir = tmp_path / "_output" + + cfg = make_train_config(root, output_dir, steps=4, ema_enable=True) + cfg.accelerator.gradient_accumulation.steps = 2 + train(cfg) + + ema_state = torch.load( + output_dir / "checkpoints" / "000004" / TRAINING_STATE_DIR / EMA_STATE_FILENAME, + weights_only=True, + ) + # 4 micro-batches / 2 accumulation steps = 2 optimizer updates. + assert ema_state["optimization_step"] == 2 + + +def test_train_without_ema_writes_no_ema_files(tmp_path): + pytest.importorskip("accelerate", reason="accelerate is required (install lerobot[training])") + pytest.importorskip("diffusers", reason="diffusers is required (install lerobot[diffusion])") + from lerobot.scripts.lerobot_train import EMA_STATE_FILENAME, train + + root = make_dummy_dataset(tmp_path) + output_dir = tmp_path / "_output" + + cfg = make_train_config(root, output_dir, steps=2, ema_enable=False) + train(cfg) + + checkpoint_dir = output_dir / "checkpoints" / "000002" + assert (checkpoint_dir / PRETRAINED_MODEL_DIR / "model.safetensors").exists() + assert not (checkpoint_dir / TRAINING_STATE_DIR / EMA_STATE_FILENAME).exists() + assert not (checkpoint_dir / f"{PRETRAINED_MODEL_DIR}_ema").exists()