hiyufan opened a new pull request, #20376:
URL: https://github.com/apache/tvm/pull/20376

   ### Problem
   
   `any` and `prod` were each half-wired: one overload missing, and the 
overload that existed returning the wrong dtype (and, for `any`, the wrong 
*value*).
   
   | expression | input | torch | frontend on `main` |
   | --- | --- | --- | --- |
   | `x.any()` | any | bool | `Unsupported function types ['any.default']` |
   | `x.any(1)` | int32 `[[0, 0, 5], [0, 0, 0]]` | bool `[True, False]` | int32 
**`[5, 0]`** |
   | `x.any(0)` | float32 | bool | float32 (the column maxima) |
   | `x.prod(0)` / `x.prod(1, keepdim=True)` | any | | `Unsupported function 
types ['prod.dim_int']` |
   | `x.prod()` | int32 `[2²⁰, 2²⁰, 2]` | int64 `2199023255552` | int32 
(overflows) |
   | `x.prod()` | bool | int64 | bool |
   
   - `_any` computed `max(x)` and only cast back to bool for a *bool* input. 
For every other dtype it returned the maximum in the input dtype — the largest 
value, not "is any element non-zero". torch.any is always bool.
   - `_prod` had no int64 accumulation for bool / integer inputs, which `_sum` 
in the same file already implements for the identical torch rule.
   
   Found by the same differential sweep as #20375; these were 27 of the 90 
raising programs, plus the silent wrong-dtype rows above.
   
   ### Fix
   
   - `_any`: `x != 0` for a non-bool input, cast the mask to int8 (relax's 
`max` does not take bool), reduce with `max`, cast back to bool. `any.default` 
dispatches to it; `any.dim` / `any.dims` unchanged in dispatch. Empty `dim` 
list means every axis.
   - `_prod`: same int64 rule as `_sum` for bool / integer inputs, honours an 
explicit `dtype=`, empty `dim` list means every axis. `prod.dim_int` dispatches 
to it.
   
   ### Verification
   
   Sweep: `any` goes from 14 raising programs to 0, `prod` from 13 to 1 (the 
0-d `()` input — a separate issue across all reductions, out of scope). 
`test_frontend_from_exported_program.py` and `test_frontend_from_fx.py`: 
failure sets identical before and after apart from the new tests 
(`test_prod[float32|bool]` and `test_dtypes[*]` fail identically on `main` in 
my environment — a TVMScript parser issue with parametrised variables, not this 
change). `ruff check` / `ruff format --check` (v0.12.3) clean.
   
   ### Tests
   
   - `test_prod_dim_and_integer_accumulation` — IR-level: `int32 x.prod(1)` 
emits `astype int64` then `R.prod(…, axis=[1])`.
   - `test_prod_values` — `prod()`, `prod(0)`, `prod(1, keepdim=True)` over 
bool, int32 (overflowing input), int64, float32; dtype and values.
   - `test_any_returns_bool_for_every_dtype` — IR-level for `int32 x.any(1)` 
(`not_equal → astype int8 → max → astype bool`); numeric `any()`, `any(1)`, 
`any(0, keepdim=True)` over bool, int32, int64, float32 with an input whose 
second row is all zeros so both truth values appear.
   
   All six fail against the previous head:
   
   ```text
   test_prod_dim_and_integer_accumulation   Unsupported function types 
['prod.dim_int']
   test_prod_values[bool]                   assert 'bool' == 'int64'
   test_prod_values[int32]                  assert 'int32' == 'int64'
   test_any_returns_bool_for_every_dtype    StructuralEqual check failed (… 
dtype)
   ```
   
   Independent of #20372 / #20373 / #20374 / #20375 (branched from `main`).
   
   ---
   
   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]

Reply via email to