Save training checkpoint Arguments: save_dir: Required. Directory for saving the checkpoint tag: Optional. Checkpoint tag used as a unique identifier for the checkpoint, global step is used if not provided. Tag name must be the same across all ranks.
(self, save_dir, tag=None, client_state={}, save_latest=True, exclude_frozen_parameters=False)
| 4982 | logger.warning(msg) |
| 4983 | |
| 4984 | def save_checkpoint(self, save_dir, tag=None, client_state={}, save_latest=True, exclude_frozen_parameters=False): |
| 4985 | """Save training checkpoint |
| 4986 | |
| 4987 | Arguments: |
| 4988 | save_dir: Required. Directory for saving the checkpoint |
| 4989 | tag: Optional. Checkpoint tag used as a unique identifier for the checkpoint, global step is |
| 4990 | used if not provided. Tag name must be the same across all ranks. |
| 4991 | client_state: Optional. State dictionary used for saving required training states in the client code. |
| 4992 | save_latest: Optional. Save a file 'latest' pointing to the latest saved checkpoint. |
| 4993 | exclude_frozen_parameters: Optional. Exclude frozen parameters from checkpointed state. |
| 4994 | Important: all processes must call this method and not just the process with rank 0. It is |
| 4995 | because each process needs to save its master weights and scheduler+optimizer states. This |
| 4996 | method will hang waiting to synchronize with other processes if it's called just for the |
| 4997 | process with rank 0. |
| 4998 | |
| 4999 | """ |
| 5000 | if not save_dir: |
| 5001 | raise ValueError(f"save_dir must be a non-empty string, got {save_dir!r}") |
| 5002 | |
| 5003 | if self._optimizer_has_ckpt_event_prologue(): |
| 5004 | # Custom preparation for checkpoint save, if applicable |
| 5005 | self.optimizer.checkpoint_event_prologue() |
| 5006 | |
| 5007 | rank = self.local_rank if self.use_node_local_storage() else self.global_rank |
| 5008 | |
| 5009 | # This is to make sure the checkpoint names are created without collision |
| 5010 | # There seems to be issue creating them in parallel |
| 5011 | |
| 5012 | # Ensure save_dir directory exists |
| 5013 | if rank == 0: |
| 5014 | self.checkpoint_engine.makedirs(save_dir, exist_ok=True) |
| 5015 | dist.barrier() |
| 5016 | |
| 5017 | if tag is None: |
| 5018 | tag = f"global_step{self.global_steps}" |
| 5019 | |
| 5020 | # Ensure tag is a string |
| 5021 | tag = str(tag) |
| 5022 | commit_info = CheckpointCommitInfo(tag=tag, save_dir=save_dir, save_latest=save_latest) |
| 5023 | |
| 5024 | self.checkpoint_engine.create(commit_info) |
| 5025 | |
| 5026 | # Ensure checkpoint tag is consistent across ranks |
| 5027 | self._checkpoint_tag_validation(tag) |
| 5028 | |
| 5029 | if self.has_moe_layers: |
| 5030 | self.save_non_zero_checkpoint = False |
| 5031 | self._create_checkpoint_file(save_dir, tag, False) |
| 5032 | self._save_moe_checkpoint(save_dir, |
| 5033 | tag, |
| 5034 | client_state=client_state, |
| 5035 | exclude_frozen_parameters=exclude_frozen_parameters) |
| 5036 | |
| 5037 | # We distribute the task of saving layer checkpoint files among |
| 5038 | # data parallel instances, so all procs should call _save_checkpoint. |
| 5039 | # All procs then call module_state_dict(), but only procs of data |
| 5040 | # parallel rank 0 save the general model params. |
| 5041 | if not self.has_moe_layers: |