Why does a data-dependent if in forward() break torch.export for on-device use?
answer
- one static graph, no Python on device
- the condition needs data that is not there yet
- fails loudly, or freezes one branch silently
- torch.cond keeps both branches
- or hoist the decision into the app
basics
~20 sExport captures one static graph, so a branch whose condition is a tensor value either fails with a data-dependent guard error or gets frozen to whichever side the example input took. The device runtime has no Python to re-evaluate it.
solid answer
~50 s`torch.export` traces the model once with example inputs and must emit a single complete graph, because the on-device runtime is a C++ interpreter with no Python left to fall back on. A branch like `if x.sum() > 0:` asks for a concrete boolean from a value that is only symbolic at export time. Two outcomes follow: on a symbolic tensor value, export raises a data-dependent guard error; if the condition happens to reduce to a Python constant during tracing, it silently specializes and the exported graph contains only the branch your example input took — which is the dangerous case, because nothing fails until production inputs hit the other path. The fixes are `torch.cond(pred, true_fn, false_fn, operands)`, which emits both branches as a real runtime conditional, rewriting the branch as arithmetic such as `torch.where`, or hoisting the decision into the calling app.
code
python · 20 linesimport torch
from torch.export import export
class Branchy(torch.nn.Module):
def forward(self, x):
if x.sum() > 0: # data-dependent: not exportable
return x * 2
return x - 1
class Exportable(torch.nn.Module):
def forward(self, x):
return torch.cond(
x.sum() > 0,
lambda t: t * 2,
lambda t: t - 1,
(x,),
)
ep = export(Exportable().eval(), (torch.randn(4),))
print(ep.graph_module.code)go deeper
Recall that a model destined for a phone must become one fixed graph, and that an if statement depending on tensor values is the classic thing that stops that from working.
Be able to explain both outcomes — a data-dependent guard error, or silent specialization to the branch the example input took — and name torch.cond as the construct that preserves a real runtime branch.
Demonstrate the diagnosis: read the export error, distinguish semantic branches from incidental debug or training branches, and insist on parity tests that exercise every path before the artifact ships to users.
Set the authoring rules that keep models exportable in the first place — no data-dependent control flow in forward(), decisions hoisted into application code, exportability checked in CI on every model change rather than discovered at release time.
## What export is trying to do `torch.export.export` produces one static graph that describes the model for every input it will ever see. It does that by running your `forward()` once with example inputs while substituting symbolic tensors for real data. Any Python construct that needs to *look at the data* to decide what to do next is therefore a problem: at trace time the data is not there, and at run time the Python is not there. This is the sharpest difference between export and `torch.compile`. `torch.compile` may hit something it cannot capture, insert a **graph break**, and let plain Python execute that region before resuming compiled code. Export has no such escape hatch, and neither does the ExecuTorch runtime on the device — it walks a serialized instruction stream. Full-graph capture is mandatory. ## The two failure modes Write this in a module: ``` def forward(self, x): if x.sum() > 0: return self.a(x) return self.b(x) ``` `x.sum() > 0` is a zero-dimensional tensor. Python's `if` calls `__bool__` on it, which forces a concrete answer. What happens next depends on how the value was produced: 1. **Loud failure.** If the condition depends on a value the exporter is tracking symbolically, export raises a data-dependent error — it cannot guard on an expression whose truth value is unknown. This is the good outcome: you find out on your build machine. 2. **Silent specialization.** If the condition collapses to a plain Python `bool` during tracing — for instance because it depends on a shape that got specialized, on a configuration attribute, or on a constant — export happily records only the taken branch. The other branch simply does not exist in the graph. The exported model runs, is fast, and is wrong for half your inputs. The second is the classic on-device bug, because the eager model in your notebook and the `.pte` on the phone disagree only for inputs your example did not cover. It is also why output-parity testing across a *varied* input set, not one sample, belongs in the export CI job. ## The same trap in loops `for i in range(x.shape[0])` behaves the same way. If the batch dimension is specialized, the loop is unrolled to exactly that many iterations and baked into the graph; feed a different batch size and the guard rejects it (or, worse, the model quietly processes the wrong number of elements if you built the loop over a Python integer you also passed in). `while` loops with tensor conditions fail outright. ## Fix 1 — `torch.cond` `torch.cond(pred, true_fn, false_fn, operands)` is the supported way to keep a genuine runtime branch: ``` return torch.cond(x.sum() > 0, self.a, self.b, (x,)) ``` It is a higher-order operator: both branches are traced and both end up in the graph, and the predicate is evaluated at run time. The constraints are real and worth stating in an interview: both branches must take the same operands, return the same number of outputs with matching dtypes and compatible shapes, must not mutate their inputs or close over tensors outside `operands`, and each must itself be exportable. Backend delegation is also less complete inside conditionals than in straight-line code, so a `torch.cond` region may execute on the runtime's own kernels rather than the accelerator. ## Fix 2 — make it arithmetic Many "branches" are really elementwise selection. `torch.where(cond, a, b)` produces a data-dependent *result* with no control flow at all. The cost is that both sides are always computed, so this is right for cheap expressions and wrong for skipping an expensive subnetwork. ## Fix 3 — hoist the decision out of the model Often the cleanest answer, and the one senior candidates reach for: if the branch decides between two genuinely different networks, export **two** methods (or two `.pte` files) and let the application code pick. The branch condition becomes app logic in Kotlin or Swift, where branching is free and debuggable, and each exported graph stays straight-line and fully delegatable. Preprocessing decisions — orientation, language, model size — usually belong here rather than inside `forward()`. ## Practical shape of the answer When an interviewer hands you "this model won't export", the useful sequence is: read the error to see whether it is a data-dependent guard or an unsupported operator; find the offending `if`/`while` in `forward()`; decide whether the branch is *semantic* (needs `torch.cond` or two exported methods) or *incidental* (a debug path, a training-only branch, an assertion) and simply delete it for the export path; then re-export and diff outputs against eager on inputs that exercise **both** sides of the original branch.
- Why is silent specialization more dangerous than an export error?An error stops your build; specialization ships. The exported graph contains only the branch the example input happened to take, so the model loads, runs fast, and returns confident wrong answers for every input that should have gone the other way. Nothing raises on device, which is why export CI must compare outputs against eager on inputs covering both branches.
- What constraints does torch.cond place on its two branch functions?They must accept the same operands, return the same number of outputs with matching dtypes and compatible shapes, avoid mutating their inputs, avoid closing over tensors not passed in operands, and each must itself be exportable. Both branches are traced into the graph, so both pay their capture cost even though only one executes per call.
- When would you export two separate methods instead of using torch.cond?When the branch selects between genuinely different networks rather than two cheap expressions. Two straight-line graphs delegate to a backend more completely than one containing a conditional region, each can be memory-planned tightly, and the decision becomes ordinary application code in Kotlin or Swift where it is easy to log and test.
- How does a Python for loop over x.shape[0] behave under export?If that dimension is specialized to the example input, the loop is fully unrolled into the graph at exactly that trip count, so the artifact only accepts that size and the graph grows with it. Declaring the dimension dynamic does not rescue a Python loop; you need a vectorized formulation or a supported loop construct instead.
saying these in an interview costs you the question
- Expects export to fall back to Python like a graph break
- Assumes both branches always end up in the graph
- Thinks torch.where gives real control flow rather than elementwise selection
- Believes the mobile runtime can evaluate the Python condition
- Tests export parity on a single input that exercises one branch