Skip to content

[Fix][Relax][Frontend][Torch] Validate flatten dims in from_fx - #20245

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

[Fix][Relax][Frontend][Torch] Validate flatten dims in from_fx#20245
hiyufan wants to merge 1 commit into
apache:mainfrom
hiyufan:fix/relax-torch-flatten-dim-validation

Conversation

@hiyufan

@hiyufan hiyufan commented Aug 31, 2026

Copy link
Copy Markdown

Fixes: #20227

Summary

tvm.relax.frontend.torch.from_fx crashed with an internal
TypeError: reduce() of empty iterable with no initial value when a traced model
contained a flatten whose start_dim comes after its end_dim.

torch rejects such a flatten with a clear RuntimeError, but only when the model is
executed. fx.symbolic_trace does not execute the model, so the invalid node reaches
the frontend as a perfectly traceable graph and has to be rejected there. This PR
validates the dims in _flatten_impl, mirroring what from_onnx now does for
Flatten (#20145).

Only the from_fx / TorchFXImporter path is affected. from_exported_program runs
run_decompositions() by default, which lowers aten.flatten.using_ints to
aten.view, so _flatten_impl is never reached there.

Root cause

_flatten_impl in python/tvm/relax/frontend/torch/base_fx_graph_translator.py
normalized negative dims but never checked their range or their ordering:

start_dim = start_dim if start_dim >= 0 else len(shape) + start_dim
end_dim = end_dim if end_dim >= 0 else len(shape) + end_dim
flattened = reduce(lambda x, y: x * y, [shape[i] for i in range(start_dim, end_dim + 1)])

For start_dim=2, end_dim=1 the range is empty, so functools.reduce is called over
an empty iterable with no initial value. Out-of-range dims (flatten(x, 0, 3) on a
rank-3 input) leaked an IndexError out of shape[i] instead.

Both entry points that reach this helper are affected: torch.flatten (dispatched via
_flatten) and torch.nn.Flatten (via _flatten_module).

Fix

Normalize both dims against max(rank, 1), validate each against [-r, r-1], and
reject start_dim > end_dim with torch's own wording:

dim_post_expr = max(rank, 1)
norm_start_dim = start_dim + dim_post_expr if start_dim < 0 else start_dim
norm_end_dim = end_dim + dim_post_expr if end_dim < 0 else end_dim
...
if norm_start_dim > norm_end_dim:
    raise ValueError("flatten() has invalid args: start_dim cannot come after end_dim")

A 0-d input is handled explicitly: torch normalizes flatten dims against a rank of at
least one, so torch.flatten(scalar) is valid and returns a 1-d tensor holding the
single element. The old code hit the same empty-reduce crash for that input.

Validation

Built from source (CPU-only, USE_LLVM=OFF) on Linux, torch 2.13.0+cpu, Python 3.14.

Behavior on the reported cases, measured before and after the change:

Input shape dims Before After
(2,3,4) (2,1) TypeError: reduce() of empty iterable with no initial value ValueError: flatten() has invalid args: start_dim cannot come after end_dim
(2,3,4) (0,3) IndexError: ShapeExpr index out of range ValueError: flatten end_dim 3 is out of range [-3, 2] for an input of rank 3
(2,3,4) (-4,2) ValueError: Reshape expects the new shape to be convertible from the old shape ValueError: flatten start_dim -4 is out of range [-3, 2] for an input of rank 3
() (0-d) (0,-1) TypeError: reduce() of empty iterable with no initial value converts, output (1,)

The (-4,2) row is worth calling out: an out-of-range negative start_dim did not fail in
_flatten_impl at all. It normalized to -1, so range(-1, 3) silently folded the last
dimension into the product and emitted a reshape to (96,) for a 24-element tensor, which
only failed later inside relax.op.reshape. That is a mis-computation, not just a crash.

Valid dims are unaffected — (1,3,10,10) with (2,-1), and (2,3,4) with (1,2), (0,-1),
(-3,-1) and (2,2), all convert to the same shapes as before.

Regression run of the whole from_fx suite, base commit vs. this branch:

base (HEAD~1) this branch
passed 164 166 (+2 new tests)
failed 16 16
skipped 1 1

The failure sets are identical line for line — the 16 failures (test_dtypes[*], test_round)
are pre-existing on the base commit in this environment and unrelated to this change.

Tests added to tests/python/relax/test_frontend_from_fx.py:

  • test_flatten_invalid_dimsstart_dim > end_dim through both torch.flatten and
    torch.nn.Flatten, plus an out-of-range end_dim
  • test_flatten_scalar_input — 0-d input flattens to shape (1,)
$ pytest tests/python/relax/test_frontend_from_fx.py -k flatten -v
test_flatten PASSED
test_flatten_invalid_dims PASSED
test_flatten_scalar_input PASSED
3 passed, 180 deselected

$ pre-commit run --files python/tvm/relax/frontend/torch/base_fx_graph_translator.py \
                    tests/python/relax/test_frontend_from_fx.py
ruff check ....... Passed
ruff format ...... Passed
(all hooks passed)

Files changed

  • python/tvm/relax/frontend/torch/base_fx_graph_translator.py_flatten_impl:
    normalize and validate dims, handle a 0-d input.
  • tests/python/relax/test_frontend_from_fx.py — regression tests.

Note for reviewers

The 0-d handling (dim_post_expr = max(rank, 1) plus the rank == 0 branch) is
separable from the reported bug. It fixes the same empty-reduce crash for scalar
inputs and keeps the new range check from rejecting a valid torch.flatten(scalar),
but if you would rather keep this PR to exactly the reported case I am happy to drop
that hunk and its test.


This change was developed with AI assistance. The build, the before/after comparison
and the test runs reported above were all executed locally against this branch as
submitted.

Fixes: apache#20227

`_flatten_impl` normalized negative flatten dims but never checked their
range or their ordering, so `from_fx` crashed with an internal
`TypeError: reduce() of empty iterable with no initial value` on a traced
model whose `start_dim` came after its `end_dim`. Out-of-range dims leaked
an `IndexError`, and an out-of-range negative `start_dim` silently computed
a wrong shape that only failed later inside `relax.op.reshape`.

Normalize both dims against `max(rank, 1)`, validate each against
`[-r, r-1]`, and reject `start_dim > end_dim` with a message matching
torch's own. A 0-d input is now handled explicitly, matching
`torch.flatten` on a scalar.

Co-authored-by: Claude <noreply@anthropic.com>
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