From 663971aa5d26c9b410b90d70340912161b8a7eaf Mon Sep 17 00:00:00 2001 From: Pepijn Date: Wed, 15 Jul 2026 23:34:10 +0200 Subject: [PATCH] feat(training): support gradient accumulation --- docs/source/multi_gpu_training.mdx | 13 +++ docs/source/robomme.mdx | 8 +- .../robomme/smolvla_visual_memory_ablation.sh | 4 +- src/lerobot/configs/train.py | 5 + src/lerobot/scripts/lerobot_train.py | 110 +++++++++++------- tests/training/test_gradient_accumulation.py | 84 +++++++++++++ 6 files changed, 180 insertions(+), 44 deletions(-) create mode 100644 tests/training/test_gradient_accumulation.py diff --git a/docs/source/multi_gpu_training.mdx b/docs/source/multi_gpu_training.mdx index 7c212364e..dd1d51891 100644 --- a/docs/source/multi_gpu_training.mdx +++ b/docs/source/multi_gpu_training.mdx @@ -113,6 +113,19 @@ accelerate launch --num_processes=2 $(which lerobot-train) \ --policy.type=act ``` +When the desired global batch is larger than the per-GPU batch that fits in memory, use gradient +accumulation. `steps` continues to count optimizer updates: + +```bash +# 8 samples/GPU × 4 GPUs × 2 microbatches = effective global batch 64. +accelerate launch --num_processes=4 $(which lerobot-train) \ + --batch_size=8 \ + --gradient_accumulation_steps=2 \ + --steps=80000 \ + --dataset.repo_id=lerobot/pusht \ + --policy.type=act +``` + ## Training Large Models with FSDP DDP replicates the full model on every GPU, so a model that doesn't fit on one GPU won't fit under diff --git a/docs/source/robomme.mdx b/docs/source/robomme.mdx index 9cf9e2c58..b0d079ea3 100644 --- a/docs/source/robomme.mdx +++ b/docs/source/robomme.mdx @@ -149,18 +149,20 @@ For a sample-matched SmolVLA visual-memory ablation, use `examples/robomme/smolvla_visual_memory_ablation.sh`. `TARGET_SAMPLES` counts examples across all GPUs, and `NUM_PROCESSES` is included when the script converts that target into optimizer steps. For example, the following reproduces 5.12 million example exposures (the exposure of RoboMME's -80,000-step, global-batch-64 memory-policy recipe) with four GPUs and a global batch of 48: +80,000-step, global-batch-64 memory-policy recipe) with four GPUs, two accumulated microbatches, +and an effective global batch of 64: ```bash TARGET_SAMPLES=5120000 \ NUM_PROCESSES=4 \ -BATCH_SIZE=12 \ +BATCH_SIZE=8 \ +GRADIENT_ACCUMULATION_STEPS=2 \ VARIANT=visual-memory \ RUN_EVAL=false \ bash examples/robomme/smolvla_visual_memory_ablation.sh ``` -This becomes 106,667 optimizer steps, or about 10.74 execution-target epochs. Run the baseline with +This becomes 80,000 optimizer steps, or about 10.74 execution-target epochs. Run the baseline with the same `TARGET_SAMPLES`, effective batch size, seed, and scheduler settings to isolate the visual memory change. RoboMME's released baseline uses a different global batch (128), so its nominal 80,000-step recipe is not sample-matched to its global-batch-64 memory recipe. diff --git a/examples/robomme/smolvla_visual_memory_ablation.sh b/examples/robomme/smolvla_visual_memory_ablation.sh index bb22d7fe6..bcd043ded 100755 --- a/examples/robomme/smolvla_visual_memory_ablation.sh +++ b/examples/robomme/smolvla_visual_memory_ablation.sh @@ -6,10 +6,11 @@ set -euo pipefail BATCH_SIZE="${BATCH_SIZE:-4}" NUM_PROCESSES="${NUM_PROCESSES:-1}" +GRADIENT_ACCUMULATION_STEPS="${GRADIENT_ACCUMULATION_STEPS:-1}" # The published training split has 476,857 execution frames. By default, train on one # execution-frame epoch; set STEPS explicitly to use a different optimizer-step budget. TARGET_SAMPLES="${TARGET_SAMPLES:-476857}" -EFFECTIVE_BATCH_SIZE=$((BATCH_SIZE * NUM_PROCESSES)) +EFFECTIVE_BATCH_SIZE=$((BATCH_SIZE * NUM_PROCESSES * GRADIENT_ACCUMULATION_STEPS)) STEPS="${STEPS:-$(((TARGET_SAMPLES + EFFECTIVE_BATCH_SIZE - 1) / EFFECTIVE_BATCH_SIZE))}" SCHEDULER_WARMUP_STEPS="${SCHEDULER_WARMUP_STEPS:-$(((STEPS + 29) / 30))}" SCHEDULER_DECAY_STEPS="${SCHEDULER_DECAY_STEPS:-${STEPS}}" @@ -58,6 +59,7 @@ COMMON_TRAIN_ARGS=( --dataset.training_target_start_feature=exec_start_idx '--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}' --batch_size="${BATCH_SIZE}" + --gradient_accumulation_steps="${GRADIENT_ACCUMULATION_STEPS}" --steps="${STEPS}" --policy.scheduler_warmup_steps="${SCHEDULER_WARMUP_STEPS}" --policy.scheduler_decay_steps="${SCHEDULER_DECAY_STEPS}" diff --git a/src/lerobot/configs/train.py b/src/lerobot/configs/train.py index e3d354691..a9c846b15 100644 --- a/src/lerobot/configs/train.py +++ b/src/lerobot/configs/train.py @@ -99,6 +99,9 @@ class TrainPipelineConfig(HubMixin): # Number of workers for the dataloader. num_workers: int = 4 batch_size: int = 8 + # Number of microbatches accumulated before each optimizer update. The effective global batch is + # batch_size * accelerator.num_processes * gradient_accumulation_steps. + gradient_accumulation_steps: int = 1 prefetch_factor: int = 4 persistent_workers: bool = True steps: int = 100_000 @@ -221,6 +224,8 @@ class TrainPipelineConfig(HubMixin): ) active_cfg = self.trainable_config + if self.gradient_accumulation_steps < 1: + raise ValueError("gradient_accumulation_steps must be at least 1.") if self.rename_map and active_cfg.pretrained_path is None: raise ValueError( "`rename_map` requires a pretrained policy checkpoint. " diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index 7fcd406b6..135839a6c 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -85,6 +85,7 @@ def update_policy( lock=None, sample_weighter=None, log_metrics: bool = True, + track_update_time: bool = True, ) -> tuple[MetricsTracker, dict | None]: """ Performs a single training step to update the policy's weights. @@ -148,13 +149,16 @@ def update_policy( # Use accelerator's backward method accelerator.backward(loss) - # Clip gradients if specified - if grad_clip_norm > 0: - grad_norm = accelerator.clip_grad_norm_(policy.parameters(), grad_clip_norm) - else: - grad_norm = torch.nn.utils.clip_grad_norm_( - policy.parameters(), float("inf"), error_if_nonfinite=False - ) + # Accelerate suppresses gradient synchronization and optimizer updates on intermediate + # microbatches. Clip and report the norm only when the accumulated update is complete. + grad_norm = None + if accelerator.sync_gradients: + if grad_clip_norm > 0: + grad_norm = accelerator.clip_grad_norm_(policy.parameters(), grad_clip_norm) + else: + grad_norm = torch.nn.utils.clip_grad_norm_( + policy.parameters(), float("inf"), error_if_nonfinite=False + ) # Optimizer step with lock if lock is not None else nullcontext(): @@ -163,19 +167,24 @@ def update_policy( optimizer.zero_grad() # Step through pytorch scheduler at every batch instead of epoch - if lr_scheduler is not None: + if lr_scheduler is not None and accelerator.sync_gradients: lr_scheduler.step() # Update internal buffers if policy has update method - if has_method(accelerator.unwrap_model(policy, keep_fp32_wrapper=True), "update"): + if accelerator.sync_gradients and has_method( + accelerator.unwrap_model(policy, keep_fp32_wrapper=True), "update" + ): accelerator.unwrap_model(policy, keep_fp32_wrapper=True).update() - train_metrics.lr = optimizer.param_groups[0]["lr"] - if torch.cuda.is_available(): + if accelerator.sync_gradients: + train_metrics.lr = optimizer.param_groups[0]["lr"] + if torch.cuda.is_available() and accelerator.sync_gradients: train_metrics.gpu_mem_gb = torch.cuda.max_memory_allocated() / (1024**3) train_metrics.accumulate_tensor("loss", loss) - train_metrics.accumulate_tensor("grad_norm", grad_norm) - train_metrics.update_s = time.perf_counter() - start_time + if grad_norm is not None: + train_metrics.accumulate_tensor("grad_norm", grad_norm) + if track_update_time: + train_metrics.update_s = time.perf_counter() - start_time # Synchronize accumulated GPU metrics only when logging. if log_metrics: train_metrics.materialize_tensors() @@ -238,6 +247,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): mixed_precision = {"bfloat16": "bf16", "float16": "fp16", "float32": "no"}.get(policy_dtype) accelerator = Accelerator( step_scheduler_with_optimizer=False, + gradient_accumulation_steps=cfg.gradient_accumulation_steps, mixed_precision=mixed_precision, kwargs_handlers=[ddp_kwargs, ipg_kwargs], cpu=force_cpu, @@ -439,8 +449,11 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): logging.info(f"{dataset.num_frames=} ({format_big_number(dataset.num_frames)})") logging.info(f"{dataset.num_episodes=}") num_processes = accelerator.num_processes - effective_bs = cfg.batch_size * num_processes - logging.info(f"Effective batch size: {cfg.batch_size} x {num_processes} = {effective_bs}") + effective_bs = cfg.batch_size * num_processes * cfg.gradient_accumulation_steps + logging.info( + "Effective batch size: " + f"{cfg.batch_size} x {num_processes} x {cfg.gradient_accumulation_steps} = {effective_bs}" + ) logging.info(f"{num_learnable_params=} ({format_big_number(num_learnable_params)})") logging.info(f"{num_total_params=} ({format_big_number(num_total_params)})") @@ -520,7 +533,12 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): f"batch_size={saved_batch_size}. The data order resumes at the right epoch/offset, " "but per-rank sample-exactness requires the same batch size." ) - sampler_state = compute_sampler_state(step, len(sampler), ckpt_batch_size, ckpt_num_processes) + sampler_state = compute_sampler_state( + step, + len(sampler), + ckpt_batch_size * cfg.gradient_accumulation_steps, + ckpt_num_processes, + ) sampler.load_state_dict(sampler_state) if is_main_process: logging.info( @@ -617,9 +635,9 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): train_metrics["gpu_mem_gb"] = AverageMeter("mem_gb", ":.2f", reduction="max") # Keep global batch size for logging; MetricsTracker handles world size internally. - effective_batch_size = cfg.batch_size * accelerator.num_processes + effective_batch_size = cfg.batch_size * accelerator.num_processes * cfg.gradient_accumulation_steps train_tracker = MetricsTracker( - cfg.batch_size, + cfg.batch_size * cfg.gradient_accumulation_steps, dataset.num_frames, dataset.num_episodes, train_metrics, @@ -641,30 +659,42 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): ) for _ in range(step, cfg.steps): - start_time = time.perf_counter() - batch = next(dl_iter) - for cam_key in dataset.meta.camera_keys: - if cam_key in batch and batch[cam_key].dtype == torch.uint8: - batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0 - if cfg.rename_map: - batch = rename_transition_keys(batch, cfg.rename_map) - batch = preprocessor(batch) - train_tracker.dataloading_s = time.perf_counter() - start_time + update_start_time = time.perf_counter() + dataloading_s = 0.0 + output_dict = None + for microbatch_idx in range(cfg.gradient_accumulation_steps): + start_time = time.perf_counter() + batch = next(dl_iter) + for cam_key in dataset.meta.camera_keys: + if cam_key in batch and batch[cam_key].dtype == torch.uint8: + batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0 + if cfg.rename_map: + batch = rename_transition_keys(batch, cfg.rename_map) + batch = preprocessor(batch) + dataloading_s += time.perf_counter() - start_time - # Synchronize GPU metrics only for updates that will be logged. - log_metrics = cfg.log_freq > 0 and (step + 1) % cfg.log_freq == 0 + # Synchronize GPU metrics only on the final microbatch of logged optimizer updates. + log_metrics = ( + cfg.log_freq > 0 + and (step + 1) % cfg.log_freq == 0 + and microbatch_idx == cfg.gradient_accumulation_steps - 1 + ) - train_tracker, output_dict = update_policy( - train_tracker, - policy, - batch, - optimizer, - cfg.optimizer.grad_clip_norm, - accelerator=accelerator, - lr_scheduler=lr_scheduler, - sample_weighter=sample_weighter, - log_metrics=log_metrics, - ) + with accelerator.accumulate(policy): + train_tracker, output_dict = update_policy( + train_tracker, + policy, + batch, + optimizer, + cfg.optimizer.grad_clip_norm, + accelerator=accelerator, + lr_scheduler=lr_scheduler, + sample_weighter=sample_weighter, + log_metrics=log_metrics, + track_update_time=False, + ) + train_tracker.dataloading_s = dataloading_s + train_tracker.update_s = time.perf_counter() - update_start_time - dataloading_s # Note: eval and checkpoint happens *after* the `step`th training update has completed, so we # increment `step` here. diff --git a/tests/training/test_gradient_accumulation.py b/tests/training/test_gradient_accumulation.py new file mode 100644 index 000000000..3b1fe6e48 --- /dev/null +++ b/tests/training/test_gradient_accumulation.py @@ -0,0 +1,84 @@ +#!/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. + +import pytest +import torch +from accelerate import Accelerator +from torch import nn + +from lerobot.scripts.lerobot_train import update_policy +from lerobot.utils.logging_utils import AverageMeter, MetricsTracker + + +class TinyPolicy(nn.Module): + def __init__(self): + super().__init__() + self.projection = nn.Linear(2, 1, bias=False) + + def forward(self, batch): + loss = self.projection(batch["x"]).square().mean() + return loss, {} + + +def test_gradient_accumulation_steps_optimizer_and_scheduler_once(): + accelerator = Accelerator( + cpu=True, + gradient_accumulation_steps=2, + step_scheduler_with_optimizer=False, + ) + policy = TinyPolicy() + optimizer = torch.optim.SGD(policy.parameters(), lr=0.1) + scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=0.5) + policy, optimizer, scheduler = accelerator.prepare(policy, optimizer, scheduler) + metrics = { + "loss": AverageMeter("loss"), + "grad_norm": AverageMeter("grad_norm"), + "lr": AverageMeter("lr"), + "update_s": AverageMeter("update_s"), + } + tracker = MetricsTracker(1, 2, 1, metrics, accelerator=accelerator) + batch = {"x": torch.ones(1, 2)} + before = policy.projection.weight.detach().clone() + + with accelerator.accumulate(policy): + update_policy( + tracker, + policy, + batch, + optimizer, + grad_clip_norm=0, + accelerator=accelerator, + lr_scheduler=scheduler, + log_metrics=False, + ) + after_first_microbatch = policy.projection.weight.detach().clone() + + with accelerator.accumulate(policy): + update_policy( + tracker, + policy, + batch, + optimizer, + grad_clip_norm=0, + accelerator=accelerator, + lr_scheduler=scheduler, + log_metrics=False, + ) + after_optimizer_step = policy.projection.weight.detach().clone() + + torch.testing.assert_close(after_first_microbatch, before) + assert not torch.equal(after_optimizer_step, after_first_microbatch) + assert optimizer.param_groups[0]["lr"] == pytest.approx(0.05)