MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / _is_checkpointable

Method _is_checkpointable

deepspeed/runtime/pipe/module.py:664–683  ·  view source on GitHub ↗
(self, funcs)

Source from the content-addressed store, hash-verified

662 self._synchronize_tied_weights()
663
664 def _is_checkpointable(self, funcs):
665
666 if self.activation_checkpoint_func is not checkpointing.non_reentrant_checkpoint:
667 # This hook excludes the embedding layer
668 # because only non_reentrant_checkpoint can accept inputs with requires_grad=False
669 # otherwise, the backward of the embedding layer won't receive gradients.
670 if self.__class__.__name__ in ('GPTModelPipe', 'GPT2ModelPipe'):
671 # For GPT models, checkpoint both transformer layers and any additional
672 # layers specified in checkpointable_layers (if provided)
673 return all('ParallelTransformerLayerPipe' in f.__class__.__name__ or (
674 self.checkpointable_layers is not None and f.__class__.__name__ in self.checkpointable_layers)
675 for f in funcs)
676
677 if self.checkpointable_layers is not None:
678 # For non-GPT models, only checkpoint layers specified in checkpointable_layers
679 return all(f.__class__.__name__ in self.checkpointable_layers for f in funcs)
680
681 # Default behavior: checkpoint any layer that has parameters
682 params = [f.parameters() for f in funcs if isinstance(f, torch.nn.Module)]
683 return any(len(list(p)) > 0 for p in params)
684
685 def get_additional_losses(self):
686 """ Returns model specific additional losses for reporting

Calls 1

parametersMethod · 0.45

Tested by 1