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]
