fix: skip gradient checkpointing log when no modules configured

- GradientCheckpointingCallback.on_train_begin returns early on empty module list
- previously logged "Gradient checkpointing enabled" even when checkpointing was inactive, misleading profiling
This commit is contained in:
2026-08-31 14:24:51 +08:00
parent 962c10c52b
commit 0546331637
+2
View File
@@ -116,6 +116,8 @@ class GradientCheckpointingCallback(TrainCallback):
del module._original_forward del module._original_forward
def on_train_begin(self, context: TrainContext): def on_train_begin(self, context: TrainContext):
if not self.modules:
return
context.model.apply(self._enable) context.model.apply(self._enable)
logger.info("Gradient checkpointing enabled") logger.info("Gradient checkpointing enabled")