From 08721f6d3125d96d6fc99d52ecfaede3224006fd Mon Sep 17 00:00:00 2001 From: 0z5a Date: Tue, 1 Sep 2026 03:05:52 +0000 Subject: [PATCH 1/2] fix: save checkpoints after optimizer steps - add a post-step callback hook for checkpoint saves - preserve updated model, optimizer, and scheduler state - cover checkpoint ordering with a regression test --- astrai/trainer/train_callback.py | 7 +++++-- astrai/trainer/trainer.py | 2 ++ tests/trainer/test_callbacks.py | 35 ++++++++++++++++++++++++++++++++ 3 files changed, 42 insertions(+), 2 deletions(-) diff --git a/astrai/trainer/train_callback.py b/astrai/trainer/train_callback.py index 3149280..3a2db73 100644 --- a/astrai/trainer/train_callback.py +++ b/astrai/trainer/train_callback.py @@ -55,7 +55,10 @@ class TrainCallback(Protocol): """Called at the end of each batch.""" def on_optimizer_step(self, context: TrainContext): - """Called on every optimizer step (sync step only).""" + """Called immediately before every optimizer step (sync step only).""" + + def on_after_optimizer_step(self, context: TrainContext): + """Called after the optimizer and scheduler step (sync step only).""" def on_error(self, context: TrainContext): """Called when an error occurs during training.""" @@ -170,7 +173,7 @@ class CheckpointCallback(TrainCallback): ) context.checkpoint.save(save_path) - def on_batch_end(self, context: TrainContext): + def on_after_optimizer_step(self, context: TrainContext): if context.optimizer_step - self.last_ckpt_step >= self.interval: self._save_checkpoint(context) diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index cb1d5eb..a4542a6 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -101,6 +101,8 @@ class Trainer: if context.scheduler: context.scheduler.step() + self._call_callbacks("on_after_optimizer_step", context) + self._call_callbacks("on_epoch_end", context) if context.stop_requested: diff --git a/tests/trainer/test_callbacks.py b/tests/trainer/test_callbacks.py index d819c37..de283b1 100644 --- a/tests/trainer/test_callbacks.py +++ b/tests/trainer/test_callbacks.py @@ -1,8 +1,12 @@ +from pathlib import Path + import torch from astrai.model.components.decoder_block import DecoderBlock +from astrai.serialization import Checkpoint from astrai.trainer.train_callback import GradientCheckpointingCallback, TrainCallback from astrai.trainer.trainer import Trainer +from tests.helpers import RandomTokenDataset def test_gradient_checkpointing_enable_disable(test_model): @@ -135,3 +139,34 @@ def test_callback_integration( assert "on_train_begin" in callback_calls assert "on_batch_end" in callback_calls assert "on_epoch_end" in callback_calls + + +def test_checkpoint_captures_completed_optimizer_step( + base_test_env, train_config_factory, device +): + """Checkpoint state must include the update represented by its step number.""" + model = base_test_env["model"] + initial_state = { + name: tensor.detach().cpu().clone() + for name, tensor in model.state_dict().items() + } + train_config = train_config_factory( + model_fn=lambda: model, + dataset=RandomTokenDataset(length=2), + test_dir=base_test_env["test_dir"], + device=device, + batch_per_device=2, + ckpt_interval=1, + ) + + Trainer(train_config).train() + + checkpoint = Checkpoint.load( + str(Path(base_test_env["test_dir"]) / "epoch_0_step_1") + ) + assert any( + not torch.equal(checkpoint.state_dict[name].cpu(), initial_tensor) + for name, initial_tensor in initial_state.items() + ) + assert checkpoint.extra["optimizer"]["state"] + assert checkpoint.extra["scheduler"]["last_epoch"] == 1 From 0e7dafad8e8f04856c0d8db61a34484073906fb5 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 1 Sep 2026 14:22:54 +0800 Subject: [PATCH 2/2] refactor: rename optimizer step callback hooks to before and after - rename on_optimizer_step to before_optimizer_step across the callback protocol, built-in callbacks, and trainer call site - rename on_after_optimizer_step to after_optimizer_step for the symmetric post-step hook - document the hook pair and the checkpoint save location in developer and training guides --- astrai/trainer/train_callback.py | 12 ++++++------ astrai/trainer/trainer.py | 4 ++-- docs/developer/architecture.md | 11 ++++++----- docs/developer/internals.md | 8 +++++--- docs/guides/training.md | 8 +++++--- 5 files changed, 24 insertions(+), 19 deletions(-) diff --git a/astrai/trainer/train_callback.py b/astrai/trainer/train_callback.py index 3a2db73..a391837 100644 --- a/astrai/trainer/train_callback.py +++ b/astrai/trainer/train_callback.py @@ -54,10 +54,10 @@ class TrainCallback(Protocol): def on_batch_end(self, context: TrainContext): """Called at the end of each batch.""" - def on_optimizer_step(self, context: TrainContext): + def before_optimizer_step(self, context: TrainContext): """Called immediately before every optimizer step (sync step only).""" - def on_after_optimizer_step(self, context: TrainContext): + def after_optimizer_step(self, context: TrainContext): """Called after the optimizer and scheduler step (sync step only).""" def on_error(self, context: TrainContext): @@ -85,7 +85,7 @@ class GradientClippingCallback(TrainCallback): def __init__(self, max_grad_norm: float): self.max_grad_norm = max_grad_norm - def on_optimizer_step(self, context: TrainContext): + def before_optimizer_step(self, context: TrainContext): context.grad_norm = context.executor.clip_grad_norm( context.model, self.max_grad_norm ) @@ -173,7 +173,7 @@ class CheckpointCallback(TrainCallback): ) context.checkpoint.save(save_path) - def on_after_optimizer_step(self, context: TrainContext): + def after_optimizer_step(self, context: TrainContext): if context.optimizer_step - self.last_ckpt_step >= self.interval: self._save_checkpoint(context) @@ -219,7 +219,7 @@ class ProgressBarCallback(TrainCallback): ) @only_on_rank(0) - def on_optimizer_step(self, context: TrainContext): + def before_optimizer_step(self, context: TrainContext): postfix = { "step": f"{context.optimizer_step:d}", "loss": f"{context.loss:.4f}", @@ -346,7 +346,7 @@ class MetricCallback(TrainCallback): for log in self.log_cache: f.write(json.dumps(log) + "\n") - def on_optimizer_step(self, context): + def before_optimizer_step(self, context): context.grad_snr_tracker.update(context.model) if ( diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index a4542a6..09fcb3a 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -93,7 +93,7 @@ class Trainer: self._call_callbacks("on_batch_end", context) if executor.sync_gradients: - self._call_callbacks("on_optimizer_step", context) + self._call_callbacks("before_optimizer_step", context) context.optimizer.step() context.strategy.on_optimizer_step() context.optimizer.zero_grad() @@ -101,7 +101,7 @@ class Trainer: if context.scheduler: context.scheduler.step() - self._call_callbacks("on_after_optimizer_step", context) + self._call_callbacks("after_optimizer_step", context) self._call_callbacks("on_epoch_end", context) diff --git a/docs/developer/architecture.md b/docs/developer/architecture.md index d5a7186..7cf96d8 100644 --- a/docs/developer/architecture.md +++ b/docs/developer/architecture.md @@ -746,13 +746,14 @@ classDiagram +on_epoch_end(context) +on_batch_begin(context) +on_batch_end(context) - +on_optimizer_step(context) + +before_optimizer_step(context) + +after_optimizer_step(context) +on_error(context) } class GradientClippingCallback { +Optional[float] max_grad_norm - +on_optimizer_step(context) + +before_optimizer_step(context) } class GradientCheckpointingCallback { @@ -767,7 +768,7 @@ classDiagram +bool weight_only +Callable save_extra_fn -_save_checkpoint(context) - +on_batch_end(context) + +after_optimizer_step(context) +on_train_end(context) +on_error(context) +save_extra(context) dict @@ -779,7 +780,7 @@ classDiagram +IO file +tqdm progress_bar +on_epoch_begin(context) - +on_optimizer_step(context) + +before_optimizer_step(context) +on_epoch_end(context) } @@ -788,7 +789,7 @@ classDiagram +int save_interval +List[str] metrics +int val_step - +on_optimizer_step(context) + +before_optimizer_step(context) +on_epoch_end(context) +on_train_end(context) +on_error(context) diff --git a/docs/developer/internals.md b/docs/developer/internals.md index 262c89f..71379dd 100644 --- a/docs/developer/internals.md +++ b/docs/developer/internals.md @@ -118,12 +118,13 @@ on_train_begin on_batch_end if executor.sync_gradients: - on_optimizer_step + before_optimizer_step optimizer.step() strategy.on_optimizer_step() optimizer.zero_grad() if scheduler: scheduler.step() + after_optimizer_step on_epoch_end on_train_end ``` @@ -139,8 +140,9 @@ Strategy metrics are detached and converted to Python `float` values before the | `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_batch_begin` | Every batch | — | -| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | -| `on_batch_end` | Every batch | `CheckpointCallback` | +| `before_optimizer_step` | Every accumulation window, before `optimizer.step()` | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | +| `on_batch_end` | Every batch | — | +| `after_optimizer_step` | Every accumulation window, after `optimizer.step()` and `scheduler.step()` | `CheckpointCallback` | | `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | diff --git a/docs/guides/training.md b/docs/guides/training.md index fd7efad..6e4a64a 100644 --- a/docs/guides/training.md +++ b/docs/guides/training.md @@ -70,12 +70,13 @@ on_train_begin on_batch_end if executor.sync_gradients: - on_optimizer_step + before_optimizer_step optimizer.step() strategy.on_optimizer_step() optimizer.zero_grad() if scheduler: scheduler.step() + after_optimizer_step on_epoch_end on_train_end ``` @@ -87,8 +88,9 @@ on_train_end | `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_batch_begin` | Every batch | — | -| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | -| `on_batch_end` | Every batch | `CheckpointCallback` | +| `before_optimizer_step` | Every accumulation window, before `optimizer.step()` | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | +| `on_batch_end` | Every batch | — | +| `after_optimizer_step` | Every accumulation window, after `optimizer.step()` and `scheduler.step()` | `CheckpointCallback` | | `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |