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]

Reply via email to