| 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 | |