Views
No views yet
checkpoint.pt - Model weights + training stateoptimizer_rank0.pt - AdamW optimizer state (GPU 0)optimizer_rank1.pt - AdamW optimizer state (GPU 1)training_state.json - Step counter, LR, etc.1import torch
2
3checkpoint = torch.load("checkpoint.pt", map_location="cpu")
4model.load_state_dict(checkpoint["model"])
5
6# Load optimizer for your GPU rank (0 or 1)
7rank = torch.distributed.get_rank()
8optimizer_state = torch.load(f"optimizer_rank{rank}.pt", map_location="cpu")
9optimizer.load_state_dict(optimizer_state)
10
11# Resume from step 954../final/ instead.