()
| 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(): |
no test coverage detected