(self, module_state_dict, zero_optimizer_state, save_frozen_param)
| 4828 | return full_state_dict |
| 4829 | |
| 4830 | def _common_checkpoint_state(self, module_state_dict, zero_optimizer_state, save_frozen_param): |
| 4831 | return dict(module=module_state_dict, |
| 4832 | buffer_names=self._get_buffer_names(), |
| 4833 | optimizer=self.optimizer.state_dict() if self.optimizer and not zero_optimizer_state else None, |
| 4834 | param_shapes=self._get_zero_param_shapes() if self.optimizer and zero_optimizer_state else None, |
| 4835 | frozen_param_shapes=self._get_zero_frozen_param_attributes(self._get_param_shape_func) |
| 4836 | if save_frozen_param else None, |
| 4837 | shared_params=self._get_shared_params() if self.optimizer and zero_optimizer_state else None, |
| 4838 | frozen_param_fragments=self._get_zero_frozen_param_attributes(self._get_param_fragment_func) |
| 4839 | if save_frozen_param else None, |
| 4840 | lr_scheduler=self.lr_scheduler.state_dict() if self.lr_scheduler is not None else None, |
| 4841 | data_sampler=self.training_dataloader.data_sampler.state_dict() if |
| 4842 | (self.training_dataloader is not None and self.curriculum_learning_enabled()) else None, |
| 4843 | random_ltd=self.random_ltd_scheduler.state_dict() if self.random_ltd_enabled() else None, |
| 4844 | sparse_tensor_module_names=self.sparse_tensor_module_names, |
| 4845 | skipped_steps=self.skipped_steps, |
| 4846 | global_steps=self.global_steps, |
| 4847 | global_samples=self.global_samples, |
| 4848 | dp_world_size=self.seq_dp_world_size, |
| 4849 | mp_world_size=self.mp_world_size, |
| 4850 | ds_config=self.config, |
| 4851 | ds_version=version) |
| 4852 | |
| 4853 | def _save_moe_checkpoint(self, save_dir, tag, client_state={}, exclude_frozen_parameters=False): |
| 4854 | save_path = self._get_ckpt_name(save_dir, tag) |
no test coverage detected