siyiweigeHEW opened a new issue, #20321:
URL: https://github.com/apache/tvm/issues/20321
### 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:
```python
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=InternalError
```
```
tvm.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_fx` path is affected. `torch.export`
rejects an
out-of-range `dim` (scalar or tuple) at trace time with `IndexError`, so
`from_exported_program` never reaches `_squeeze` with 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 bad `dim` silently 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_fx` and
`from_exported_program` (both `squeeze`, `squeeze.dim` and `squeeze.dims`
dispatch to
it), so a single fix covers all of them and makes the error message
consistent with the
range check that `relax.op.squeeze` already performs.
### Environment
- OS: Linux
- TVM: main branch (`60a9871a9`, re-verified 2026-09-12; also observed on
`390af87345`)
- Python: 3.11
- torch: 2.10.0+cu128
- target: `llvm`
### Steps to reproduce
```python
"""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
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]