How do you load a GPU-trained PyTorch checkpoint on a CPU-only inference box?
answer
- a checkpoint remembers its device
- one argument to torch.load
- load to CPU, then move once
- dict of tensors beats a pickled object
- unpickling is code execution
basics
~20 sPass map_location to torch.load, for example torch.load(path, map_location="cpu"), so saved tensors are restored onto the CPU instead of the CUDA device they were saved from. Then load_state_dict into a model and move the model once with .to(device).
solid answer
~50 sA checkpoint records the device each tensor lived on. Without `map_location`, `torch.load` tries to restore them onto that same CUDA device and fails on a CPU-only machine with an error mentioning that `torch.cuda.is_available() is False`. Pass `map_location="cpu"` — or `map_location=device` where `device` came from `torch.device("cuda" if torch.cuda.is_available() else "cpu")` — and the tensors deserialize onto the target you actually have. Then `model.load_state_dict(state)` and `model.to(device)`; inputs must be on the same device as the parameters or the forward pass raises. Two related habits matter for deployment: save a `state_dict` rather than a pickled module, because that keeps the checkpoint independent of your module's import path, and know that since torch 2.6 `torch.load` defaults to `weights_only=True`, which refuses arbitrary pickled objects. That default is a security improvement — an untrusted checkpoint can otherwise execute code on load — and old whole-module checkpoints may need `torch.serialization.add_safe_globals` or a re-save as a `state_dict`.
code
python · 13 linesimport torch
model = torch.nn.Linear(8, 8)
torch.save(model.state_dict(), "w.pt") # save a state_dict, not the module
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
state = torch.load("w.pt", map_location="cpu") # weights_only=True by default
model.load_state_dict(state) # strict: keys must match
model.to(device).eval()
x = torch.randn(2, 8, device=device) # inputs on the same device
with torch.no_grad():
print(model(x).shape)go deeper
Be able to write the four lines from memory: torch.load with map_location, load_state_dict, model.to(device), eval mode. Say plainly that the checkpoint records its original device and map_location overrides it.
Explain why state_dict beats pickling the module, what strict key matching protects you from, and the module. and _orig_mod. prefixes that come from distributed and compiled training wrappers.
Show that you treat checkpoint loading as an integrity step: verify outputs against a reference input after load, refuse strict=False as a fix, and reason about weights_only and untrusted model files as a real supply-chain risk.
Own the artifact contract — what a checkpoint contains, how it is versioned and signed, which loads are allowed to disable weights_only — so that model files moving between training, staging and third parties are governed rather than trusted by habit.
## Why a checkpoint is device-flavoured When you `torch.save` a tensor, PyTorch stores its data and the device it lived on. Restoring it means re-materializing the storage on that device. On a CPU-only inference box there is no `cuda:0` to restore to, so `torch.load` raises rather than guessing — the message tells you to use `map_location`, and it names `torch.cuda.is_available() is False` as the reason. ## map_location `map_location` tells the deserializer where storages should land. The common forms: - `map_location="cpu"` — everything to CPU. Safest default for a load-then-move flow. - `map_location=torch.device("cuda:0")` — remap onto a specific device, useful when the checkpoint came from `cuda:3` and this box has one GPU. - `map_location=lambda storage, loc: storage` — the old idiom meaning "leave it where the current process would put it", i.e. CPU. The standard deployment pattern is: load to CPU, `load_state_dict`, then `model.to(device)` once at startup. That works identically whether the target has a GPU or not, and it is the version you should be able to write on a whiteboard. ## Devices must agree at run time Parameters and inputs must be on the same device or the forward pass raises "Expected all tensors to be on the same device". So a serving wrapper computes `device` once, moves the model there once, and moves each incoming batch there per request. Do not scatter `.cuda()` calls through model code — a hardcoded `.cuda()` makes the module unloadable on a CPU box no matter what `map_location` you passed. ## state_dict versus the whole module `torch.save(model, path)` pickles the module object, which stores a reference to your class by import path. Load it on a machine where that module has moved, been renamed or has a different signature and it breaks. `torch.save(model.state_dict(), path)` stores only a dict of tensor names to tensors; you reconstruct the architecture in code and call `load_state_dict`. That is the format to use for anything that outlives the training script. `load_state_dict` is strict by default: unexpected or missing keys raise. Common causes are a `module.` prefix from `DataParallel`/DDP-wrapped training, and an `_orig_mod.` prefix from a `torch.compile` wrapper. The fix is to strip the prefix when saving or loading rather than to pass `strict=False`, which silently leaves layers at their random initialization — a genuinely dangerous escape hatch. ## The weights_only default Since torch 2.6, `torch.load` defaults to `weights_only=True`. Unpickling is code execution: a maliciously crafted checkpoint can run arbitrary Python the moment you load it, which matters the instant you accept model files from a hub, a customer or a shared bucket. `weights_only=True` restricts deserialization to tensors and a small allowlist of safe types. The practical consequence is that old checkpoints containing non-tensor objects — a whole pickled module, a custom optimizer config class, a numpy scalar wrapper — now fail to load. The right responses, in order: re-save as a plain `state_dict`; or register the specific classes you trust with `torch.serialization.add_safe_globals`. Setting `weights_only=False` is only defensible for a file you produced and control, and saying that out loud is what separates a good answer from a shrug. ## What to check once it loads Put the model in evaluation mode before serving and confirm parameter dtypes are what you expect — a checkpoint saved from mixed-precision training may hold half-precision tensors that a CPU box will run slowly or refuse for some operators. And run one known input through both the training-time model and the freshly loaded one, comparing outputs with a tolerance. A checkpoint that loads without error is not the same as a checkpoint that loaded correctly, and `strict=False` plus a prefix mismatch produces exactly that failure: a service that starts fine and returns garbage.
- load_state_dict raises about unexpected keys beginning with module. — what happened?The checkpoint was saved from a model wrapped in DataParallel or DistributedDataParallel, which nests the real model under an attribute named module, so every key gains that prefix. Save from the unwrapped model (model.module.state_dict()) during training, or strip the prefix from the keys when loading. Reaching for strict=False instead hides the mismatch and leaves those layers randomly initialized.
- Why did torch.load's weights_only default change to True, and what breaks?Because unpickling executes code: loading an untrusted checkpoint could run arbitrary Python. weights_only=True restricts deserialization to tensors and a small allowlist. What breaks are checkpoints containing non-tensor objects — a whole pickled module or a custom config class. Re-save them as plain state_dicts, or allowlist the specific classes with torch.serialization.add_safe_globals. Flipping the flag back off is only reasonable for files you produced yourself.
- The model loads and runs on CPU but is far slower than expected. What would you check?Parameter dtype and threading. A checkpoint from mixed-precision training can hold half-precision weights, which most CPU kernels handle poorly or not at all — cast to float32 for CPU serving. Then check that inference runs under no-grad so no autograd graph is built, that the model is in evaluation mode, and that the process is not fighting other workers for cores, since PyTorch's CPU kernels default to using many threads.
saying these in an interview costs you the question
- torch.load figures out the right device automatically
- map_location determines where the model lives after loading
- weights_only=False is harmless, it's just a compatibility flag
- strict=False is the normal fix for key mismatches
- A checkpoint trained on GPU can only run on GPU