The GitHub Actions job "Lint" on tvm.git/main has succeeded. Run started by GitHub user tlopex (triggered by tlopex).
Head commit for run: b3f51d1f28e602376c04c7c3d57b4f58e170954e / Chen Yufan <[email protected]> [Fix][Relax][Frontend][Torch] Validate flatten dims in `from_fx` (#20245) 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: ```python 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: ```python 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_dims` — `start_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. Co-authored-by: Claude <[email protected]> Report URL: https://github.com/apache/tvm/actions/runs/34184160273 With regards, GitHub Actions via GitBox --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
