MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / _load_optimizer

Function _load_optimizer

accessory/util/misc.py:474–486  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

472
473
474 def _load_optimizer():
475 # if dp_rank == 0:
476 consilidated_optimizer_checkpoint_path = os.path.join(
477 args.resume,
478 f"consolidated.{mp_rank:02d}-of-{mp_world_size:02d}.optimizer.pth",
479 )
480 full_osd = torch.load(consilidated_optimizer_checkpoint_path)['optimizer']
481 # else:
482 # full_osd = None
483
484 sharded_osd = FSDP.shard_full_optim_state_dict(full_osd, model, optim=optimizer)
485 optimizer.load_state_dict(sharded_osd)
486 print(f"load optimizer from {consilidated_optimizer_checkpoint_path}")
487 _load_optimizer()
488
489 def _load_other():

Callers 1

resume_stage2Function · 0.85

Calls 2

printFunction · 0.85
load_state_dictMethod · 0.45

Tested by

no test coverage detected