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
This commit is contained in:
2026-09-01 14:22:54 +08:00
parent 08721f6d31
commit 0e7dafad8e
5 changed files with 24 additions and 19 deletions
+6 -6
View File
@@ -54,10 +54,10 @@ class TrainCallback(Protocol):
def on_batch_end(self, context: TrainContext): def on_batch_end(self, context: TrainContext):
"""Called at the end of each batch.""" """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).""" """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).""" """Called after the optimizer and scheduler step (sync step only)."""
def on_error(self, context: TrainContext): def on_error(self, context: TrainContext):
@@ -85,7 +85,7 @@ class GradientClippingCallback(TrainCallback):
def __init__(self, max_grad_norm: float): def __init__(self, max_grad_norm: float):
self.max_grad_norm = max_grad_norm 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.grad_norm = context.executor.clip_grad_norm(
context.model, self.max_grad_norm context.model, self.max_grad_norm
) )
@@ -173,7 +173,7 @@ class CheckpointCallback(TrainCallback):
) )
context.checkpoint.save(save_path) 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: if context.optimizer_step - self.last_ckpt_step >= self.interval:
self._save_checkpoint(context) self._save_checkpoint(context)
@@ -219,7 +219,7 @@ class ProgressBarCallback(TrainCallback):
) )
@only_on_rank(0) @only_on_rank(0)
def on_optimizer_step(self, context: TrainContext): def before_optimizer_step(self, context: TrainContext):
postfix = { postfix = {
"step": f"{context.optimizer_step:d}", "step": f"{context.optimizer_step:d}",
"loss": f"{context.loss:.4f}", "loss": f"{context.loss:.4f}",
@@ -346,7 +346,7 @@ class MetricCallback(TrainCallback):
for log in self.log_cache: for log in self.log_cache:
f.write(json.dumps(log) + "\n") 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) context.grad_snr_tracker.update(context.model)
if ( if (
+2 -2
View File
@@ -93,7 +93,7 @@ class Trainer:
self._call_callbacks("on_batch_end", context) self._call_callbacks("on_batch_end", context)
if executor.sync_gradients: if executor.sync_gradients:
self._call_callbacks("on_optimizer_step", context) self._call_callbacks("before_optimizer_step", context)
context.optimizer.step() context.optimizer.step()
context.strategy.on_optimizer_step() context.strategy.on_optimizer_step()
context.optimizer.zero_grad() context.optimizer.zero_grad()
@@ -101,7 +101,7 @@ class Trainer:
if context.scheduler: if context.scheduler:
context.scheduler.step() 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) self._call_callbacks("on_epoch_end", context)
+6 -5
View File
@@ -746,13 +746,14 @@ classDiagram
+on_epoch_end(context) +on_epoch_end(context)
+on_batch_begin(context) +on_batch_begin(context)
+on_batch_end(context) +on_batch_end(context)
+on_optimizer_step(context) +before_optimizer_step(context)
+after_optimizer_step(context)
+on_error(context) +on_error(context)
} }
class GradientClippingCallback { class GradientClippingCallback {
+Optional[float] max_grad_norm +Optional[float] max_grad_norm
+on_optimizer_step(context) +before_optimizer_step(context)
} }
class GradientCheckpointingCallback { class GradientCheckpointingCallback {
@@ -767,7 +768,7 @@ classDiagram
+bool weight_only +bool weight_only
+Callable save_extra_fn +Callable save_extra_fn
-_save_checkpoint(context) -_save_checkpoint(context)
+on_batch_end(context) +after_optimizer_step(context)
+on_train_end(context) +on_train_end(context)
+on_error(context) +on_error(context)
+save_extra(context) dict +save_extra(context) dict
@@ -779,7 +780,7 @@ classDiagram
+IO file +IO file
+tqdm progress_bar +tqdm progress_bar
+on_epoch_begin(context) +on_epoch_begin(context)
+on_optimizer_step(context) +before_optimizer_step(context)
+on_epoch_end(context) +on_epoch_end(context)
} }
@@ -788,7 +789,7 @@ classDiagram
+int save_interval +int save_interval
+List[str] metrics +List[str] metrics
+int val_step +int val_step
+on_optimizer_step(context) +before_optimizer_step(context)
+on_epoch_end(context) +on_epoch_end(context)
+on_train_end(context) +on_train_end(context)
+on_error(context) +on_error(context)
+5 -3
View File
@@ -118,12 +118,13 @@ on_train_begin
on_batch_end on_batch_end
if executor.sync_gradients: if executor.sync_gradients:
on_optimizer_step before_optimizer_step
optimizer.step() optimizer.step()
strategy.on_optimizer_step() strategy.on_optimizer_step()
optimizer.zero_grad() optimizer.zero_grad()
if scheduler: if scheduler:
scheduler.step() scheduler.step()
after_optimizer_step
on_epoch_end on_epoch_end
on_train_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_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — | | `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | | `before_optimizer_step` | Every accumulation window, before `optimizer.step()` | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback` | | `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_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
+5 -3
View File
@@ -70,12 +70,13 @@ on_train_begin
on_batch_end on_batch_end
if executor.sync_gradients: if executor.sync_gradients:
on_optimizer_step before_optimizer_step
optimizer.step() optimizer.step()
strategy.on_optimizer_step() strategy.on_optimizer_step()
optimizer.zero_grad() optimizer.zero_grad()
if scheduler: if scheduler:
scheduler.step() scheduler.step()
after_optimizer_step
on_epoch_end on_epoch_end
on_train_end on_train_end
``` ```
@@ -87,8 +88,9 @@ on_train_end
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | | `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — | | `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` | | `before_optimizer_step` | Every accumulation window, before `optimizer.step()` | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback` | | `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_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |