[Relax][Frontend][Torch] `_squeeze` does not validate `dim`: out-of-bounds tuple dims are silently converted to `squeeze(None)` via `from_fx` (native PyTorch raises `IndexError`)
Expected behavior
tvm.relax.frontend.torch._squeeze
(python/tvm/relax/frontend/torch/base_fx_graph_translator.py:2523-2543) reads the
dim/dims argument of a PyTorch squeeze call, drops out-of-range axes from a
list/tuple, and falls back to dim=None when the filtered list is empty:
if isinstance(dim, list | tuple) and len(dim) > 0:
shape = self.shape_of(x)
valid_dims = []
for d in dim:
axis = d if d >= 0 else len(shape) + d
if axis < len(shape):
valid_dims.append(d)
# If no valid dims, use None to squeeze all size-1 dimensions
dim = valid_dims if valid_dims else None
return self.block_builder.emit(relax.op.squeeze(x, dim))An out-of-range dim should be rejected with a clear, frontend-level error, matching
native PyTorch, which raises IndexError: Dimension out of range for the same call. It
should not be silently dropped and reinterpreted as a different operation.
Actual behavior
When dim is a list/tuple whose elements are all positive out-of-range, valid_dims
becomes empty, dim is reset to None, and the call is silently converted into
squeeze(None) (remove all size-1 dimensions). torch.fx.symbolic_trace records an
out-of-range tuple dim without validating it, so the legacy from_fx path imports the
model and emits a differently-shaped tensor where native PyTorch raises IndexError:
shape=(2, 3) squeeze((5,)) torch=IndexError tvm=OK (2, 3)
shape=(2, 1, 3) squeeze((5,)) torch=IndexError tvm=OK (2, 3) <- size-1 dim silently removed
shape=(1, 2, 1) squeeze((5,)) torch=IndexError tvm=OK (2,) <- both size-1 dims removed(2, 1, 3) is the clearest case: the module is declared to keep its shape, and the
frontend instead returns a tensor with the size-1 axis removed.
The other argument forms are rejected, but only by the C++ side of relax.op.squeeze,
with an opaque op-level error rather than a frontend one — so the forms behave
inconsistently:
shape=(2, 3) squeeze(5) torch=IndexError tvm=InternalError
shape=(2, 3) squeeze(-5) torch=IndexError tvm=InternalError
shape=(2, 3) squeeze((-5,)) torch=IndexError tvm=InternalErrortvm.error.InternalError: In Op(relax.squeeze), the input axis 5 is out of range.
The input tensor has 2 dimensions, so axis should be in range [-2, 2).The negative out-of-range case is kept by the filter (len(shape) + d is always
< len(shape) for d < 0), so it reaches the op; the positive out-of-range case is
dropped by the filter and never reaches it. Both should be a frontend error.
Additional context:
- Scope — Only the legacy
from_fxpath is affected.torch.exportrejects an out-of-rangedim(scalar or tuple) at trace time withIndexError, sofrom_exported_programnever reaches_squeezewith a bad dim. Because the model is invalid in native PyTorch anyway (it would crash on any input), this is a robustness/validation gap rather than a wrong result on a valid model — but a model with a latent, never-exercised baddimsilently produces a differently-shaped module instead of surfacing the error. - In-bounds dims are unaffected — control cases (
squeeze((0,2))on(1,2,1),squeeze((1,))on(2,1,3)) match PyTorch exactly. - Misleading comment — the filter comment says "filter out axes where dimension is not 1", but the code only bounds-checks axes; it never inspects the size. The comment does not describe what the code does, and the fallback it feeds does not match PyTorch semantics.
- Where to validate — this converter is shared by
from_fxandfrom_exported_program(bothsqueeze,squeeze.dimandsqueeze.dimsdispatch to it), so a single fix covers all of them and makes the error message consistent with the range check thatrelax.op.squeezealready performs.
Environment
- OS: Linux
- TVM: main branch (
60a9871a9, re-verified 2026-09-12; also observed on390af87345) - Python: 3.11
- torch: 2.10.0+cu128
- target:
llvm
Steps to reproduce
"""Repro: TVM torch frontend silently converts out-of-bounds tuple dims to squeeze(None)."""
import numpy as np
import torch
import torch.nn as nn
from torch.fx import symbolic_trace
import tvm
from tvm import relax
from tvm.relax.frontend.torch import from_fx
class M(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, t):
return t.squeeze(self.dim)
def run(shape, dim):
m = M(dim).eval()
x = np.arange(int(np.prod(shape))).reshape(shape).astype(np.float32) + 1
xt = torch.tensor(x)
try: # native PyTorch ground truth
ref = m(xt).numpy().shape
ref_txt = f"OK {ref}"
except Exception as e:
ref_txt = f"{type(e).__name__}"
try: # TVM legacy from_fx path
gm = symbolic_trace(m)
mod = from_fx(gm, [((*shape,), "float32")])
ex = relax.build(mod, target="llvm")
vm = relax.VirtualMachine(ex, tvm.cpu())
out = vm["main"](x)
arr = out[0] if hasattr(out, "__len__") and len(out) else out
tv_txt = f"OK {tuple(np.asarray(arr.numpy()).shape)}"
except Exception as e:
tv_txt = f"{type(e).__name__}"
print(f" shape={shape} dim={dim} torch={ref_txt} tvm={tv_txt}")
if __name__ == "__main__":
print("tvm:", tvm.__version__, "| torch:", torch.__version__)
print("\n# Out-of-bounds dims in a tuple silently become 'squeeze(None)':")
run((2, 3), (5,))
run((2, 1, 3), (5,))
run((1, 2, 1), (5,))
print("\n# Same out-of-bounds dim as a scalar int raises an opaque op-level error:")
run((2, 3), 5)
run((2, 3), -5)
run((2, 3), (-5,))
print("\n# In-bounds dims are unaffected (control):")
run((1, 2, 1), (0, 2))
run((2, 1, 3), (1,))Actual output:
tvm: 0.24.dev0 | torch: 2.10.0+cu128
# Out-of-bounds dims in a tuple silently become 'squeeze(None)':
shape=(2, 3) dim=(5,) torch=IndexError tvm=OK (2, 3)
shape=(2, 1, 3) dim=(5,) torch=IndexError tvm=OK (2, 3)
shape=(1, 2, 1) dim=(5,) torch=IndexError tvm=OK (2,)
# Same out-of-bounds dim as a scalar int raises an opaque op-level error:
shape=(2, 3) dim=5 torch=IndexError tvm=InternalError
shape=(2, 3) dim=-5 torch=IndexError tvm=InternalError
shape=(2, 3) dim=(-5,) torch=IndexError tvm=InternalError
# In-bounds dims are unaffected (control):
shape=(1, 2, 1) dim=(0, 2) torch=OK (2,) tvm=OK (2,)
shape=(2, 1, 3) dim=(1,) torch=OK (2, 3) tvm=OK (2, 3)Triage
- needs-triage
- bug
- relax
- frontend/torch
Source: apache/tvm