Skip to content
PyTorch

Save and Load

Checkpoint models, optimizer state, and weights.

#save#load#checkpoint

Code

pytorch
import torch

model = torch.nn.Linear(10, 2)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

# Save state_dict (recommended)
torch.save(model.state_dict(), "model.pth")
model.load_state_dict(torch.load("model.pth"))
model.eval()

# Full checkpoint with optimizer and epoch
checkpoint = {
    "epoch": 10,
    "model_state": model.state_dict(),
    "optim_state": optimizer.state_dict(),
}
torch.save(checkpoint, "ckpt.pth")

ck = torch.load("ckpt.pth")
model.load_state_dict(ck["model_state"])
optimizer.load_state_dict(ck["optim_state"])

# Save on GPU and load on CPU
torch.save(model.state_dict(), "gpu.pth")
model.load_state_dict(torch.load("gpu.pth", map_location="cpu"))