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

Function _save_optimizer

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

Source from the content-addressed store, hash-verified

388 print("model saved")
389
390 def _save_optimizer():
391 # torch.FSDP has a bug that passing dp_group to FSDP.full_optim_state_dict still calls dist.gather within the whole world
392 _world = torch.distributed.GroupMember.WORLD
393 torch.distributed.GroupMember.WORLD = fs_init.get_data_parallel_group()
394
395 consolidated_optim_state_dict = {
396 "optimizer": FSDP.full_optim_state_dict(model, optimizer)
397 }
398 save_path = os.path.join(
399 save_dir,
400 f"consolidated.{mp_rank:02d}-of-{mp_world_size:02d}.optimizer.pth",
401 )
402 if fs_init.get_data_parallel_rank() == 0:
403 torch.save(consolidated_optim_state_dict , save_path)
404
405 torch.distributed.GroupMember.WORLD = _world
406 _save_optimizer()
407 print("optimizer saved")
408

Callers 1

save_checkpointFunction · 0.85

Calls 1

saveMethod · 0.80

Tested by

no test coverage detected