mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-27 11:46:04 +00:00
feat(training): support gradient accumulation
This commit is contained in:
@@ -113,6 +113,19 @@ accelerate launch --num_processes=2 $(which lerobot-train) \
|
|||||||
--policy.type=act
|
--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
|
## 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
|
DDP replicates the full model on every GPU, so a model that doesn't fit on one GPU won't fit under
|
||||||
|
|||||||
@@ -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
|
`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
|
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
|
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
|
```bash
|
||||||
TARGET_SAMPLES=5120000 \
|
TARGET_SAMPLES=5120000 \
|
||||||
NUM_PROCESSES=4 \
|
NUM_PROCESSES=4 \
|
||||||
BATCH_SIZE=12 \
|
BATCH_SIZE=8 \
|
||||||
|
GRADIENT_ACCUMULATION_STEPS=2 \
|
||||||
VARIANT=visual-memory \
|
VARIANT=visual-memory \
|
||||||
RUN_EVAL=false \
|
RUN_EVAL=false \
|
||||||
bash examples/robomme/smolvla_visual_memory_ablation.sh
|
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
|
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
|
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.
|
80,000-step recipe is not sample-matched to its global-batch-64 memory recipe.
|
||||||
|
|||||||
@@ -6,10 +6,11 @@ set -euo pipefail
|
|||||||
|
|
||||||
BATCH_SIZE="${BATCH_SIZE:-4}"
|
BATCH_SIZE="${BATCH_SIZE:-4}"
|
||||||
NUM_PROCESSES="${NUM_PROCESSES:-1}"
|
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
|
# 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.
|
# execution-frame epoch; set STEPS explicitly to use a different optimizer-step budget.
|
||||||
TARGET_SAMPLES="${TARGET_SAMPLES:-476857}"
|
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))}"
|
STEPS="${STEPS:-$(((TARGET_SAMPLES + EFFECTIVE_BATCH_SIZE - 1) / EFFECTIVE_BATCH_SIZE))}"
|
||||||
SCHEDULER_WARMUP_STEPS="${SCHEDULER_WARMUP_STEPS:-$(((STEPS + 29) / 30))}"
|
SCHEDULER_WARMUP_STEPS="${SCHEDULER_WARMUP_STEPS:-$(((STEPS + 29) / 30))}"
|
||||||
SCHEDULER_DECAY_STEPS="${SCHEDULER_DECAY_STEPS:-${STEPS}}"
|
SCHEDULER_DECAY_STEPS="${SCHEDULER_DECAY_STEPS:-${STEPS}}"
|
||||||
@@ -58,6 +59,7 @@ COMMON_TRAIN_ARGS=(
|
|||||||
--dataset.training_target_start_feature=exec_start_idx
|
--dataset.training_target_start_feature=exec_start_idx
|
||||||
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
||||||
--batch_size="${BATCH_SIZE}"
|
--batch_size="${BATCH_SIZE}"
|
||||||
|
--gradient_accumulation_steps="${GRADIENT_ACCUMULATION_STEPS}"
|
||||||
--steps="${STEPS}"
|
--steps="${STEPS}"
|
||||||
--policy.scheduler_warmup_steps="${SCHEDULER_WARMUP_STEPS}"
|
--policy.scheduler_warmup_steps="${SCHEDULER_WARMUP_STEPS}"
|
||||||
--policy.scheduler_decay_steps="${SCHEDULER_DECAY_STEPS}"
|
--policy.scheduler_decay_steps="${SCHEDULER_DECAY_STEPS}"
|
||||||
|
|||||||
@@ -99,6 +99,9 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
# Number of workers for the dataloader.
|
# Number of workers for the dataloader.
|
||||||
num_workers: int = 4
|
num_workers: int = 4
|
||||||
batch_size: int = 8
|
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
|
prefetch_factor: int = 4
|
||||||
persistent_workers: bool = True
|
persistent_workers: bool = True
|
||||||
steps: int = 100_000
|
steps: int = 100_000
|
||||||
@@ -221,6 +224,8 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
active_cfg = self.trainable_config
|
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:
|
if self.rename_map and active_cfg.pretrained_path is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"`rename_map` requires a pretrained policy checkpoint. "
|
"`rename_map` requires a pretrained policy checkpoint. "
|
||||||
|
|||||||
@@ -85,6 +85,7 @@ def update_policy(
|
|||||||
lock=None,
|
lock=None,
|
||||||
sample_weighter=None,
|
sample_weighter=None,
|
||||||
log_metrics: bool = True,
|
log_metrics: bool = True,
|
||||||
|
track_update_time: bool = True,
|
||||||
) -> tuple[MetricsTracker, dict | None]:
|
) -> tuple[MetricsTracker, dict | None]:
|
||||||
"""
|
"""
|
||||||
Performs a single training step to update the policy's weights.
|
Performs a single training step to update the policy's weights.
|
||||||
@@ -148,13 +149,16 @@ def update_policy(
|
|||||||
# Use accelerator's backward method
|
# Use accelerator's backward method
|
||||||
accelerator.backward(loss)
|
accelerator.backward(loss)
|
||||||
|
|
||||||
# Clip gradients if specified
|
# Accelerate suppresses gradient synchronization and optimizer updates on intermediate
|
||||||
if grad_clip_norm > 0:
|
# microbatches. Clip and report the norm only when the accumulated update is complete.
|
||||||
grad_norm = accelerator.clip_grad_norm_(policy.parameters(), grad_clip_norm)
|
grad_norm = None
|
||||||
else:
|
if accelerator.sync_gradients:
|
||||||
grad_norm = torch.nn.utils.clip_grad_norm_(
|
if grad_clip_norm > 0:
|
||||||
policy.parameters(), float("inf"), error_if_nonfinite=False
|
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
|
# Optimizer step
|
||||||
with lock if lock is not None else nullcontext():
|
with lock if lock is not None else nullcontext():
|
||||||
@@ -163,19 +167,24 @@ def update_policy(
|
|||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
|
|
||||||
# Step through pytorch scheduler at every batch instead of epoch
|
# 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()
|
lr_scheduler.step()
|
||||||
|
|
||||||
# Update internal buffers if policy has update method
|
# 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()
|
accelerator.unwrap_model(policy, keep_fp32_wrapper=True).update()
|
||||||
|
|
||||||
train_metrics.lr = optimizer.param_groups[0]["lr"]
|
if accelerator.sync_gradients:
|
||||||
if torch.cuda.is_available():
|
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.gpu_mem_gb = torch.cuda.max_memory_allocated() / (1024**3)
|
||||||
train_metrics.accumulate_tensor("loss", loss)
|
train_metrics.accumulate_tensor("loss", loss)
|
||||||
train_metrics.accumulate_tensor("grad_norm", grad_norm)
|
if grad_norm is not None:
|
||||||
train_metrics.update_s = time.perf_counter() - start_time
|
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.
|
# Synchronize accumulated GPU metrics only when logging.
|
||||||
if log_metrics:
|
if log_metrics:
|
||||||
train_metrics.materialize_tensors()
|
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)
|
mixed_precision = {"bfloat16": "bf16", "float16": "fp16", "float32": "no"}.get(policy_dtype)
|
||||||
accelerator = Accelerator(
|
accelerator = Accelerator(
|
||||||
step_scheduler_with_optimizer=False,
|
step_scheduler_with_optimizer=False,
|
||||||
|
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
|
||||||
mixed_precision=mixed_precision,
|
mixed_precision=mixed_precision,
|
||||||
kwargs_handlers=[ddp_kwargs, ipg_kwargs],
|
kwargs_handlers=[ddp_kwargs, ipg_kwargs],
|
||||||
cpu=force_cpu,
|
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_frames=} ({format_big_number(dataset.num_frames)})")
|
||||||
logging.info(f"{dataset.num_episodes=}")
|
logging.info(f"{dataset.num_episodes=}")
|
||||||
num_processes = accelerator.num_processes
|
num_processes = accelerator.num_processes
|
||||||
effective_bs = cfg.batch_size * num_processes
|
effective_bs = cfg.batch_size * num_processes * cfg.gradient_accumulation_steps
|
||||||
logging.info(f"Effective batch size: {cfg.batch_size} x {num_processes} = {effective_bs}")
|
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_learnable_params=} ({format_big_number(num_learnable_params)})")
|
||||||
logging.info(f"{num_total_params=} ({format_big_number(num_total_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, "
|
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."
|
"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)
|
sampler.load_state_dict(sampler_state)
|
||||||
if is_main_process:
|
if is_main_process:
|
||||||
logging.info(
|
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")
|
train_metrics["gpu_mem_gb"] = AverageMeter("mem_gb", ":.2f", reduction="max")
|
||||||
|
|
||||||
# Keep global batch size for logging; MetricsTracker handles world size internally.
|
# 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(
|
train_tracker = MetricsTracker(
|
||||||
cfg.batch_size,
|
cfg.batch_size * cfg.gradient_accumulation_steps,
|
||||||
dataset.num_frames,
|
dataset.num_frames,
|
||||||
dataset.num_episodes,
|
dataset.num_episodes,
|
||||||
train_metrics,
|
train_metrics,
|
||||||
@@ -641,30 +659,42 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
)
|
)
|
||||||
|
|
||||||
for _ in range(step, cfg.steps):
|
for _ in range(step, cfg.steps):
|
||||||
start_time = time.perf_counter()
|
update_start_time = time.perf_counter()
|
||||||
batch = next(dl_iter)
|
dataloading_s = 0.0
|
||||||
for cam_key in dataset.meta.camera_keys:
|
output_dict = None
|
||||||
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
for microbatch_idx in range(cfg.gradient_accumulation_steps):
|
||||||
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
start_time = time.perf_counter()
|
||||||
if cfg.rename_map:
|
batch = next(dl_iter)
|
||||||
batch = rename_transition_keys(batch, cfg.rename_map)
|
for cam_key in dataset.meta.camera_keys:
|
||||||
batch = preprocessor(batch)
|
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
||||||
train_tracker.dataloading_s = time.perf_counter() - start_time
|
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.
|
# 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
|
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(
|
with accelerator.accumulate(policy):
|
||||||
train_tracker,
|
train_tracker, output_dict = update_policy(
|
||||||
policy,
|
train_tracker,
|
||||||
batch,
|
policy,
|
||||||
optimizer,
|
batch,
|
||||||
cfg.optimizer.grad_clip_norm,
|
optimizer,
|
||||||
accelerator=accelerator,
|
cfg.optimizer.grad_clip_norm,
|
||||||
lr_scheduler=lr_scheduler,
|
accelerator=accelerator,
|
||||||
sample_weighter=sample_weighter,
|
lr_scheduler=lr_scheduler,
|
||||||
log_metrics=log_metrics,
|
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
|
# Note: eval and checkpoint happens *after* the `step`th training update has completed, so we
|
||||||
# increment `step` here.
|
# increment `step` here.
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user