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:
@@ -116,6 +116,8 @@ class GradientCheckpointingCallback(TrainCallback):
|
||||
del module._original_forward
|
||||
|
||||
def on_train_begin(self, context: TrainContext):
|
||||
if not self.modules:
|
||||
return
|
||||
context.model.apply(self._enable)
|
||||
logger.info("Gradient checkpointing enabled")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user