(self, module_state_dict, zero_optimizer_state, save_frozen_param)
| 4703 | return full_state_dict |
| 4704 | |
| 4705 | def _common_checkpoint_state(self, module_state_dict, zero_optimizer_state, save_frozen_param): |
| 4706 | return dict(module=module_state_dict, |
| 4707 | buffer_names=self._get_buffer_names(), |
| 4708 | optimizer=self.optimizer.state_dict() if self.optimizer and not zero_optimizer_state else None, |
| 4709 | param_shapes=self._get_zero_param_shapes() if self.optimizer and zero_optimizer_state else None, |
| 4710 | frozen_param_shapes=self._get_zero_frozen_param_attributes(self._get_param_shape_func) |
| 4711 | if save_frozen_param else None, |
| 4712 | shared_params=self._get_shared_params() if self.optimizer and zero_optimizer_state else None, |
| 4713 | frozen_param_fragments=self._get_zero_frozen_param_attributes(self._get_param_fragment_func) |
| 4714 | if save_frozen_param else None, |
| 4715 | lr_scheduler=self.lr_scheduler.state_dict() if self.lr_scheduler is not None else None, |
| 4716 | data_sampler=self.training_dataloader.data_sampler.state_dict() if |
| 4717 | (self.training_dataloader is not None and self.curriculum_learning_enabled()) else None, |
| 4718 | random_ltd=self.random_ltd_scheduler.state_dict() if self.random_ltd_enabled() else None, |
| 4719 | sparse_tensor_module_names=self.sparse_tensor_module_names, |
| 4720 | skipped_steps=self.skipped_steps, |
| 4721 | global_steps=self.global_steps, |
| 4722 | global_samples=self.global_samples, |
| 4723 | dp_world_size=self.seq_dp_world_size, |
| 4724 | mp_world_size=self.mp_world_size, |
| 4725 | ds_config=self.config, |
| 4726 | ds_version=version) |
| 4727 | |
| 4728 | def _save_moe_checkpoint(self, save_dir, tag, client_state={}, exclude_frozen_parameters=False): |
| 4729 | save_path = self._get_ckpt_name(save_dir, tag) |
no test coverage detected