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] 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` |