(self, checkpoints_path, tag, mp_placeholder=None)
| 398 | return ckpt_files |
| 399 | |
| 400 | def _get_ckpt_name(self, checkpoints_path, tag, mp_placeholder=None): |
| 401 | if mp_placeholder is not None: |
| 402 | mp_rank_str = mp_placeholder |
| 403 | else: |
| 404 | mp_rank = 0 if self.mpu is None else self.mpu.get_model_parallel_rank() |
| 405 | mp_rank_str = "{:02d}".format(mp_rank) |
| 406 | |
| 407 | ckpt_name = os.path.join( |
| 408 | checkpoints_path, |
| 409 | "mp_rank_" + mp_rank_str + "_model_states.pt", |
| 410 | ) |
| 411 | return ckpt_name |
| 412 | |
| 413 | def _load_checkpoint(self, load_dir, load_module_strict=True, tag=None): |
| 414 | is_pipe_parallel = isinstance(self.module, PipelineModule) |
no test coverage detected