Skip to content

[Fix][Relax][Frontend][Torch] Validate dim in the squeeze converter - #20322

Open
siyiweigeHEW wants to merge 1 commit into
apache:mainfrom
siyiweigeHEW:fix/relax-torch-squeeze-dim-validation
Open

[Fix][Relax][Frontend][Torch] Validate dim in the squeeze converter#20322
siyiweigeHEW wants to merge 1 commit into
apache:mainfrom
siyiweigeHEW:fix/relax-torch-squeeze-dim-validation

Conversation

@siyiweigeHEW

Copy link
Copy Markdown
Contributor

Fixes: #20321

Summary

The Relax PyTorch frontend's _squeeze converter
(python/tvm/relax/frontend/torch/base_fx_graph_translator.py) reads the dim/dims
argument of a squeeze call, drops out-of-range axes from a list/tuple, and falls
back to dim=None when the filtered list turns out 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

For an all-out-of-range tuple such as squeeze((5,)), valid_dims is empty, dim
becomes None, and the call is silently reinterpreted as squeeze(None) — remove
every size-1 dimension. Native PyTorch raises IndexError for the same model, so the
frontend turns an invalid model into a differently-shaped one instead of reporting the
bad axis:

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

The other argument forms already fail, but only on the C++ side of relax.op.squeeze,
with an opaque op-level error rather than a frontend one:

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).

This PR replaces the filter with an explicit range check that rejects the out-of-range
axis with a clear ValueError, on every argument form.

Root cause

The filter's only observable effect was to drop axes with d >= rank (for d < 0,
len(shape) + d < len(shape) is always true, so negative axes were never dropped). So:

  • axes with d >= rank were silently dropped, and dropping all of them flipped the
    meaning of the call to "squeeze everything";
  • axes with d < -rank were kept and handed to relax.op.squeeze, which rejected them
    with an InternalError;
  • scalar dim never went through the filter at all, so it always took the
    InternalError path.

Both outcomes are wrong for the frontend: torch rejects the model outright, and the
converter is the last place that can report which axis is bad. The comment above the
filter ("filter out axes where dimension is not 1") also does not describe what the code
does — it never inspects the dimension size, only the bounds.

Fix

python/tvm/relax/frontend/torch/base_fx_graph_translator.py_squeeze:

# torch rejects an out-of-range dim with IndexError, but only when the model is
# executed. fx.symbolic_trace does not execute it, so the invalid axis reaches this
# converter and has to be rejected here. An out-of-range axis in a list used to be
# filtered out, and a list whose axes were all out of range fell back to
# `dim=None` -- silently squeezing every size-1 dim instead of reporting the axis.
if dim is not None:
    rank = len(self.shape_of(x))
    for d in dim if isinstance(dim, list | tuple) else [dim]:
        if isinstance(d, int) and not -rank <= d < rank:
            raise ValueError(
                f"squeeze dim {d} is out of range "
                f"[-{rank}, {rank - 1}] for an input of rank {rank}"
            )

return self.block_builder.emit(relax.op.squeeze(x, dim))

Notes:

  • the isinstance(d, int) guard means the check only applies to literal integer axes,
    so a non-literal axis is forwarded unchanged — the fix cannot reject anything the old
    code accepted other than genuinely out-of-range integer axes;
  • dim is no longer rewritten, so a tuple is passed through as-is (relax.op.squeeze
    accepts both tuples and lists) and an empty tuple keeps its existing no-op meaning;
  • _squeeze is shared by from_fx and from_exported_program, and is dispatched from
    squeeze, squeeze.dim and squeeze.dims, so a single change covers all of them and
    makes the message consistent with the range check relax.op.squeeze already performs.

Validation

In-tree regression test (added)

test_squeeze_out_of_range_dim in tests/python/relax/test_frontend_from_fx.py:

  • rejects dim ∈ {(5,), 3, -4, (-4,), (0, 5)} on a (1, 2, 1) input with
    ValueError: squeeze dim <d> is out of range ... — this covers a fully
    out-of-range tuple, a mixed tuple, and the scalar forms;
  • keeps the in-range tuple form working: squeeze((0, 2)) on (1, 2, 1) still
    lowers to R.squeeze(inp_0, axis=[0, 2]) (asserted structurally via
    verify_model).

The test fails without the fix with Failed: DID NOT RAISE <class 'ValueError'> on the
headline (5,) case, and passes with it.

No test is added for from_exported_program: torch.export rejects an out-of-range
dim at trace time with IndexError, so that path never reaches _squeeze with a bad
axis.

Behaviour after the fix

shape=(2, 3)    squeeze((5,))  torch=IndexError  tvm=ValueError: squeeze dim 5 is out of range [-2, 1] for an input of rank 2
shape=(2, 1, 3) squeeze((5,))  torch=IndexError  tvm=ValueError: squeeze dim 5 is out of range [-3, 2] for an input of rank 3
shape=(1, 2, 1) squeeze((5,))  torch=IndexError  tvm=ValueError: squeeze dim 5 is out of range [-3, 2] for an input of rank 3
shape=(2, 3)    squeeze(5)     torch=IndexError  tvm=ValueError: squeeze dim 5 is out of range [-2, 1] for an input of rank 2
shape=(2, 3)    squeeze(-5)    torch=IndexError  tvm=ValueError: squeeze dim -5 is out of range [-2, 1] for an input of rank 2
shape=(2, 3)    squeeze((-5,)) torch=IndexError  tvm=ValueError: squeeze dim -5 is out of range [-2, 1] for an input of rank 2

In-range behaviour is unchanged; the control cases from the existing test_squeeze and
from a direct differential against native PyTorch all still match:

shape=(1, 2, 1, 3) squeeze((0, 2))  torch=OK (2, 3)        tvm=OK (2, 3)
shape=(2, 1, 3)    squeeze((0,))    torch=OK (2, 1, 3)     tvm=OK (2, 1, 3)
shape=(1, 2, 1, 3) squeeze((-1,))   torch=OK (1, 2, 1, 3)  tvm=OK (1, 2, 1, 3)
shape=(2, 3)       squeeze(())      torch=OK (2, 3)        tvm=OK (2, 3)
shape=(2, 1, 3)    squeeze((0, 1))  torch=OK (2, 3)        tvm=OK (2, 3)

Full-suite run

tests/python/relax/test_frontend_from_fx.py and
tests/python/relax/test_frontend_from_exported_program.py run in full (373 passed,
16 failed, 2 skipped). The 16 failures are the same pre-existing, unrelated ones seen on
the unmodified tree (test_extended_unary_ops, test_interpolate, test_select_slice,
test_masked_select, test_to_copy, test_index_put, test_eye, test_cross_entropy,
the test_dynamic_shape* family, test_sym_size_int, test_stochastic_depth); the
failure set is unchanged, so there are no regressions.

Files changed

  • python/tvm/relax/frontend/torch/base_fx_graph_translator.py — replace the
    axis-dropping filter in _squeeze with an out-of-range check.
  • tests/python/relax/test_frontend_from_fx.py — add
    test_squeeze_out_of_range_dim.

`_squeeze` dropped out-of-range axes from a list/tuple `dim` and fell back to
`dim=None` when the filtered list turned out empty. For an all-out-of-range
tuple such as `squeeze((5,))` this silently reinterpreted the call as
`squeeze(None)`, removing every size-1 dimension, where native PyTorch raises
`IndexError`:

    shape=(2, 1, 3) squeeze((5,))  torch=IndexError  tvm=OK (2, 3)

The other argument forms were rejected, but only by the C++ side of
`relax.op.squeeze`, with an opaque op-level error:

    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

The filter's only observable effect was to drop axes with `d >= rank`, so
replace it with an explicit out-of-range check on all argument forms. The
converter is shared by `from_fx` and `from_exported_program` and is dispatched
from `squeeze`, `squeeze.dim` and `squeeze.dims`, so one check covers them all
and reports the offending axis instead of leaving it to the op.

Add a regression test.

Reported in apache#20321.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

1 participant