skip to content

TorchScript and Deployment

Getting a model out of Python: TorchScript tracing versus scripting, torch.export and ONNX, torch.compile for speed, then serving via TorchServe or a mobile runtime. Interviewers probe where tracing silently drops your control flow.

on this pageshow

questions

6

How do you load a GPU-trained PyTorch checkpoint on a CPU-only inference box?

level: juniorimportance: must knowfreq 55%

answer

  1. a checkpoint remembers its device
  2. one argument to torch.load
  3. load to CPU, then move once
  4. dict of tensors beats a pickled object
  5. unpickling is code execution

basics

~20 s

Pass 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 s

A 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 lines
python
import 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

for a junior

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.

for a middle

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.

for a senior

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.

for a principal

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

context

open as a page

How does torch.compile differ from torch.jit.script for shipping a PyTorch model?

level: middleimportance: must knowfreq 60%

basics

~20 s

torch.compile is a just-in-time accelerator that stays inside Python: TorchDynamo captures graphs from bytecode, a backend generates kernels, and unsupported code simply breaks the graph and runs eagerly. It emits no portable artifact, so it speeds a process up rather than getting the model out of Python.

open as a page

In PyTorch, when does torch.jit.trace silently produce a wrong TorchScript model?

level: middleimportance: must knowfreq 70%

basics

~20 s

torch.jit.trace records only the operations one example input actually executed, so a branch or loop that depends on tensor values is frozen as whatever ran that day. torch.jit.script compiles the Python source instead and keeps the control flow.

open as a page

What does torch.export.export() return, and how is it stricter than tracing?

level: middleimportance: should knowfreq 50%

basics

~20 s

torch.export.export() returns an ExportedProgram: one whole-graph ATen representation plus a signature describing which inputs are parameters, buffers and user arguments, and symbolic shape constraints. Unlike tracing, it refuses to silently specialize — unsupported dynamism raises instead.

open as a page

Your ONNX export of a PyTorch model returns different numbers — how do you debug it?

level: seniorimportance: should knowfreq 45%

basics

~20 s

Separate a real bug from floating-point noise first by comparing with a tolerance, not equality. Then check the usual causes in order: the module was not in evaluation mode at export, shapes were baked because no dynamic axes were declared, and unsupported operators changed semantics. Bisect by exporting submodules.

open as a page

How do you choose between ONNX Runtime, Triton, and ExecuTorch for a PyTorch model?

level: principalimportance: should knowfreq 32%

basics

~20 s

Let the target decide. On-device means ExecuTorch; a heterogeneous CPU or accelerator fleet with simple pre-processing favours ONNX Runtime; a GPU service with Python-shaped pre- and post-processing favours a Triton PyTorch or Python backend. TorchServe is retired and should not anchor a new design.

open as a page