(
checkpoint_dir,
model_no_ddp,
optimizer,
epoch,
args,
best_val_metrics,
filename=None,
)
| 6 | |
| 7 | |
| 8 | def save_checkpoint( |
| 9 | checkpoint_dir, |
| 10 | model_no_ddp, |
| 11 | optimizer, |
| 12 | epoch, |
| 13 | args, |
| 14 | best_val_metrics, |
| 15 | filename=None, |
| 16 | ): |
| 17 | if not is_primary(): |
| 18 | return |
| 19 | if filename is None: |
| 20 | filename = f"checkpoint_{epoch:04d}.pth" |
| 21 | checkpoint_name = os.path.join(checkpoint_dir, filename) |
| 22 | |
| 23 | weight_ckpt = model_no_ddp.state_dict() |
| 24 | sd = { |
| 25 | "model": weight_ckpt, |
| 26 | "optimizer": optimizer.state_dict(), |
| 27 | "epoch": epoch, |
| 28 | "args": args, |
| 29 | "best_val_metrics": best_val_metrics, |
| 30 | } |
| 31 | torch.save(sd, checkpoint_name) |
| 32 | |
| 33 | |
| 34 | def resume_if_possible(checkpoint_dir, model_no_ddp, optimizer): |
no test coverage detected