(self, finish=False)
| 338 | ), f"Load model failed, unexpected keys {load_result.unexpected_keys.__str__()}" |
| 339 | |
| 340 | def save_model(self, finish=False): |
| 341 | if self.config.save_model and (finish or self._last_global_step_saved != self.global_step): |
| 342 | if finish: |
| 343 | model_fname = os.path.join(self.config.exp_dir, "finish.pt") |
| 344 | else: |
| 345 | model_fname = os.path.join(self.config.exp_dir, f"global_step{self.global_step}.pt") |
| 346 | |
| 347 | if self.use_deepspeed or self.use_ddp: |
| 348 | distributed_save_path = os.path.join(self.config.exp_dir, "saved_model") |
| 349 | self.trainer.model.save_checkpoint(distributed_save_path) |
| 350 | torch.distributed.barrier() |
| 351 | if dist.get_rank() == 0: |
| 352 | trainable_states = zero_to_fp32.get_fp32_state_dict_from_zero_checkpoint(distributed_save_path) |
| 353 | prefix_length = len("module.model.") |
| 354 | trainable_states = {k[prefix_length:]: v for k, v in trainable_states.items()} |
| 355 | torch.save(trainable_states, model_fname) |
| 356 | else: |
| 357 | trainable_states = { |
| 358 | param_name: param_weight.cpu() |
| 359 | for param_name, param_weight in self.model.state_dict().items() |
| 360 | if param_name in self.trainable_param_names |
| 361 | } |
| 362 | torch.save(trainable_states, model_fname) |
| 363 | |
| 364 | self._last_global_step_saved = self.global_step |
| 365 | |
| 366 | def on_before_optimizer_step(self, optimizer, optimizer_idx): |
| 367 | if self.config.fishmask_mode is not None: |
no outgoing calls
no test coverage detected