MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / _get_ckpt_name

Method _get_ckpt_name

deepspeed/inference/engine.py:400–411  ·  view source on GitHub ↗
(self, checkpoints_path, tag, mp_placeholder=None)

Source from the content-addressed store, hash-verified

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)

Callers 1

_get_all_ckpt_namesMethod · 0.95

Calls 1

Tested by

no test coverage detected