Short answer: save your model, optimizer and step counter to one directory, write those files atomically, load them on startup, and run on a platform that copies that directory before the machine disappears and restores it on the next machine. Do that and a preemption costs you minutes of work, not the whole run.
Spot and interruptible GPUs are often a fraction of on-demand prices. The catch is that the provider can take the machine back. Here is the pattern we use at Nodus, and it works anywhere.
1. Put all recovery state in one directory
Everything you need to continue goes in one place: weights, optimizer state, LR scheduler, RNG state, data loader position, current step. Results you want to download (final model, eval reports) go somewhere else. Mixing the two makes checkpoints big and slow.
On Nodus that directory is /nodus/state (also in $NODUS_STATE_DIR), and outputs go to /nodus/outputs.
2. Write atomically
A checkpoint taken while you are halfway through writing a file restores a broken state. Write to a temp file, then rename it over the old one. Rename is atomic on POSIX filesystems.
import json, os
from pathlib import Path
STATE_DIR = Path(os.environ.get("NODUS_STATE_DIR", "/nodus/state"))
STATE = STATE_DIR / "progress.json"
def save(step, model, opt):
STATE_DIR.mkdir(parents=True, exist_ok=True)
tmp = STATE_DIR / "ckpt.pt.tmp"
torch.save({"model": model.state_dict(), "opt": opt.state_dict(), "step": step}, tmp)
os.replace(tmp, STATE_DIR / "ckpt.pt")
3. Resume on startup
Checkpoints restore files, not process memory. Your program starts from the top, so it has to check for saved state:
start = 0
ckpt = STATE_DIR / "ckpt.pt"
if ckpt.exists():
s = torch.load(ckpt)
model.load_state_dict(s["model"]); opt.load_state_dict(s["opt"]); start = s["step"]
print(f"resumed from step {start}")
for step in range(start + 1, total_steps + 1):
...
4. Pick a checkpoint cadence
Too often and you waste GPU time writing files. Too rarely and each preemption throws away a lot of work. A practical rule: aim for at least four checkpoints per expected run, and keep checkpointing under about 10% of runtime. Capacity that is rarely interrupted can be checkpointed less often.
Nodus does this math for you with interval: auto, using how often that capacity actually gets interrupted and how long your saves take.
5. Save when the machine is about to go away
Most providers give a short reclaim notice. Use it. On Nodus your program can subscribe to checkpoint requests over a local socket and acknowledge when files are consistent:
import nodus # pip install nodus-compute; no-ops outside Nodus
nodus.checkpoint.on_request(lambda: save(step, model, opt))
When the capacity gives notice, Nodus sends an urgent request, takes the checkpoint after your ack, and prepares a replacement machine at the same time. The replacement only starts once the old one is provably gone, so two machines never write the same state.
6. Hugging Face Trainer users
Trainer already saves and resumes. Point output_dir at the state directory, set save_steps and save_total_limit=2, and call trainer.train(resume_from_checkpoint=True) when NODUS_RESTORED=1. Or set integration: HFTrainer and Nodus registers the save-on-request callback for you.
Running it
pip install nodus-compute
nodus login
nodus run --gpu H100 --interruptible --checkpoint /nodus/state -d -- python train.py
nodus describe job/<name> shows each attempt, why it ended (for example Preempted) and the latest checkpoint. New accounts get a $30 starter grant, so you can test a preemption-safe run without a card.
FAQ
Does checkpointing restore GPU memory? No. It restores files. Your code reloads its own state.
What if a checkpoint is empty? An empty state directory never replaces an earlier useful checkpoint, so a crash before the first save cannot erase progress.
How many times will it retry? By default up to 8 recoveries (3 for distributed jobs). Two attempts in a row with no progress stops the run with NoProgress.
Does this work for multi-node? Yes, in beta. Each rank writes its shard with torch.distributed.checkpoint and rank 0 commits it.
Full guide: nodus-compute.ai/docs/guides/checkpoints
Top comments (0)