hiyufan opened a new pull request, #20377:
URL: https://github.com/apache/tvm/pull/20377
### Problem
`_reshape` skips an identity reshape with
```python
if current_shape is not None and list(current_shape) == list(dims):
return x
```
On a symbolic dimension `==` builds a `PrimExpr` instead of answering, and
Python then asks it for a truth value:
```text
ValueError: Cannot use and / or / not operator to Expr, hint: use
tvm.tirx.all / tvm.tirx.any, ...
File "python/tvm/relax/frontend/torch/base_fx_graph_translator.py", line
2559, in _reshape
File "python/tvm/ir/expr.py", line 347, in __bool__
```
It only surfaces when the ranks match, because list equality compares
lengths first — which is why `x.reshape(x.shape[0], -1)` on a rank-3 input
imports and `x.reshape(x.shape[0], 0, x.shape[0])` raises:
| expression | input | torch | frontend on `main` |
| --- | --- | --- | --- |
| `x.reshape(x.shape[0], 0, x.shape[0])` | `(batch, 0, 4)` | `(batch, 0,
batch)` | `ValueError` |
| `x.reshape(0, x.shape[1], x.shape[1])` | `(0, batch, 4)` | `(0, batch,
batch)` | `ValueError` |
| `x.reshape(x.shape[0], -1)` | `(batch, 0, 4)` | `(batch, 0)` | ok (rank
differs) |
This is the "40 pre-existing `ValueError` cases" I reported as out of scope
in #20255 — I had assumed they were below the frontend. They are this line.
### Fix
`_same_dims` compares dimension by dimension: static dims as integers,
symbolic dims with `tvm_ffi.structural_equal`, a static-vs-symbolic pair as
different. A genuine identity with a symbolic batch is still recognised
(`x.reshape(x.shape[0], 2, 4)` on `(batch, 2, 4)` emits no reshape, asserted in
the test). An expression written differently on the two sides (`s*2` vs `2*s`)
is treated as different, which costs a no-op reshape and never a wrong shape.
### Verification
The reshape sweep from #20255 (938 `reshape`/`view`/`flatten`/`unflatten`
targets over inputs mixing symbolic and zero dims, output shape compared with
`torch.export`'s, symbol names canonicalised):
| | matched | mismatched | of which raising |
| --- | --- | --- | --- |
| `main` | 892 | 46 | 40 |
| this PR | **932** | 6 | **0** |
Per-case diff: **no case fails with this change that did not fail before it;
40 repaired.** The 6 left are the sweep's own renderer printing `s*2` and `2*s`
differently (verified in #20255's thread), not the frontend.
`test_frontend_from_exported_program.py` and `test_frontend_from_fx.py`:
failure sets identical before and after apart from the new test. `ruff check` /
`ruff format --check` (v0.12.3) clean.
### Test
`test_reshape_symbolic_target_same_rank` — the first row above builds and
matches torch under a dynamic batch, and the identity case emits no reshape.
Fails against the previous head with the `ValueError` above.
Independent of #20372–#20376 (branched from `main`). Touches the same
`_reshape` as #20255 (merged), one line below it.
---
This change was prepared with AI assistance (Claude). I have reviewed and
verified it, and can speak to it in review.
--
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]