| 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 |