How does torch.compile differ from torch.jit.script for shipping a PyTorch model?
answer
- one accelerates, one exports
- hooks bytecode, keeps Python alive
- unsupported code breaks the graph
- guards fail, you recompile
- the wrapper renames your state_dict keys
basics
~20 storch.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.
solid answer
~50 s`torch.compile(model)` returns a wrapper around the same module. On the first call TorchDynamo intercepts the Python bytecode, extracts FX graphs, attaches *guards* describing the assumptions it made (shapes, dtypes, Python values), and hands each graph to a backend — Inductor by default — which generates fused kernels. Anything Dynamo cannot capture becomes a **graph break**: that region runs in ordinary eager Python and compilation resumes after it. When guards fail, it recompiles. `torch.jit.script` is a different job: it compiles a typed Python subset ahead of time into a serializable TorchScript program that a C++ runtime can load with no interpreter. So the two are not competing optimizations — `torch.compile` is for making training and Python-hosted inference faster and needs the Python process, while scripting or `torch.export` is for producing a file that leaves it. Practical gotchas: the first call pays capture and codegen, `state_dict` keys gain an `_orig_mod.` prefix so save the original module, and `fullgraph=True` turns graph breaks from a silent slowdown into an error.
code
python · 10 linesimport torch
model = torch.nn.Linear(8, 8).eval()
compiled = torch.compile(model) # OptimizedModule wrapper
compiled(torch.randn(16, 8)) # first call: capture + codegen
print(list(compiled.state_dict())[:2]) # ['_orig_mod.weight', '_orig_mod.bias']
print(list(model.state_dict())[:2]) # ['weight', 'bias']
torch.save(compiled._orig_mod.state_dict(), "w.pt") # save the real modulego deeper
Know that torch.compile is a one-line speedup for a model running in Python and that it is not a way to save or export the model. Recognize that the first call is slow because compilation happens then.
Explain the pipeline — Dynamo captures graphs from bytecode with guards, Inductor generates kernels — and contrast it with scripting: JIT and in-process versus ahead-of-time and serializable, tolerant graph breaks versus hard compile errors.
Show operational command: warm up compiled models before serving traffic, diagnose recompilation with TORCH_LOGS, control shape variance with dynamic=True or mark_dynamic, audit with fullgraph=True, and avoid the _orig_mod checkpoint trap.
Frame it as a cost decision: compile time and cache management in every container start, a model-code standard that keeps graphs capturable, and a clear split between the compile path used for training throughput and the export path that produces the artifact you actually ship.
## Two different verbs People put `torch.compile` and `torch.jit.script` side by side because both take a module and return something faster-sounding. They do different jobs. `torch.jit.script` **exports**: it compiles ahead of time into TorchScript, a serialized program a C++ runtime can execute with no Python present. `torch.compile` **accelerates in place**: it is a just-in-time compiler that only ever runs inside a live Python process, and it produces no artifact you can ship. If the interview question is "how do we get this model off the training box", `torch.compile` is not the answer. If it is "our training step is CPU-launch-bound", it usually is. ## The three pieces under torch.compile - **TorchDynamo** hooks CPython's frame evaluation and reads the bytecode of your `forward`. It traces the tensor operations into an FX graph and, crucially, records **guards** — the conditions under which this graph is valid: this input has rank 3, this dtype is float32, this Python flag is `True`, this attribute is that object. - **AOTAutograd** captures the backward graph too, so training gets compiled as well, not just inference. - **A backend**, by default **TorchInductor**, lowers the graph to generated kernels — Triton kernels on GPU, C++/OpenMP on CPU — fusing elementwise chains and cutting kernel launches. ## Graph breaks: the thing that makes it forgiving and slow When Dynamo meets code it cannot capture — a call into an opaque C extension, a data-dependent branch on a tensor value, printing a tensor, an unsupported data structure — it does not fail. It **breaks the graph**: compiles what it had, runs the unsupported region in eager Python, and starts a fresh graph after it. The model stays correct and gets less speedup, which is a much friendlier default than TorchScript's compile error but also means "it compiled" tells you nothing about whether it helped. `fullgraph=True` forbids breaks and raises instead, which is how you audit. `torch._dynamo.explain` and `TORCH_LOGS="graph_breaks"` show where they are. ## Guards and recompilation: the other tax Every compiled graph is protected by guards. Change something a guard covers — a new input shape, a different dtype, a flipped Python flag — and Dynamo compiles a *new* variant. That is fine for a handful of shapes and a disaster for a service with arbitrary sequence lengths: you pay compile time repeatedly and eventually hit the recompilation limit (`torch._dynamo.config.cache_size_limit`), after which that code falls back to eager permanently. The mitigations are `dynamic=True` or marking a dimension with `torch._dynamo.mark_dynamic` so one graph covers a range of sizes, and padding or bucketing inputs so only a few shapes ever occur. `TORCH_LOGS="recompiles"` prints which guard failed and why, and is the first thing to turn on when a compiled service is mysteriously slow. ## Modes and warmup `mode="default"` is a balanced compile; `"reduce-overhead"` adds CUDA graphs to cut launch overhead for small, repetitive shapes; `"max-autotune"` benchmarks kernel variants and takes substantially longer to compile. All of them front-load cost: the first call to a compiled model is seconds to minutes, not milliseconds. Inductor caches generated artifacts on disk (`TORCHINDUCTOR_CACHE_DIR`), so a warm cache shortens later starts — but a fresh container with a cold cache pays in full. Any service using `torch.compile` needs a deliberate warmup pass over representative shapes before it accepts traffic. ## The wrapper is not your module `torch.compile` returns an `OptimizedModule` holding the original as `_orig_mod`. Its `state_dict` keys are prefixed accordingly, so a checkpoint saved from the compiled wrapper will not load cleanly into the plain model later. Save `compiled._orig_mod.state_dict()` — or keep a reference to the original module and save from that. This is a real and frequent production bug. ## How to talk about it The crisp framing: `torch.compile` is an optimization applied to a running Python program, tolerant of anything it cannot handle, and it leaves no artifact behind. `torch.jit.script` and `torch.export` are captures that produce artifacts and are intolerant by design, because a static file cannot fall back to an interpreter that will not be there. A production PyTorch story usually uses both — compile the training job, export the model for serving — and knowing they are not alternatives is the point of the question.
- Can you serialize the result of torch.compile and load it on a machine without PyTorch?No. torch.compile produces in-process compiled code guarded by Python-level assumptions; there is no portable file at the end, and the wrapper is not serializable as a program. To leave Python you need a capture path — torch.export to a .pt2 and an ahead-of-time or on-device consumer, an ONNX export, or the older TorchScript file that libtorch loads. Compile makes a process fast; export makes a model portable.
- A compiled service is slower than eager under production traffic. Where do you look first?Recompilation. Turn on TORCH_LOGS="recompiles" and look at which guard keeps failing — usually an input shape that varies per request. If so, compile with dynamic=True or mark the varying dimension with torch._dynamo.mark_dynamic, and bucket or pad inputs so few distinct shapes reach the model. Then check graph breaks with fullgraph=True: many small compiled regions can cost more than they save.
- What does fullgraph=True actually change?It forbids graph breaks. Instead of compiling what it can and running the rest eagerly, Dynamo raises on the first construct it cannot capture, naming it. Nothing about the generated code changes when capture succeeds — it is an audit switch. Teams use it in a test to assert a model stays fully capturable, then ship with the default so an unexpected construct degrades performance rather than taking the service down.
saying these in an interview costs you the question
- torch.compile writes out a compiled model file you can ship
- Compiling removes the dependency on Python at inference time
- A graph break is an error that aborts compilation
- The first call after torch.compile is as fast as the rest
- Stack torch.compile on top of torch.jit.script for more speed