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
|
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")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user