hiyufan opened a new pull request, #20374: URL: https://github.com/apache/tvm/pull/20374
### Problem With no `dtype` argument, `torch.cumsum` / `torch.cumprod` accumulate every integral and bool input in **int64**. The converters passed `dtype=None` through to `relax.op.cumsum` / `cumprod`, so the running sum kept the input dtype — and wrapped: | input | torch | frontend on `main` | | --- | --- | --- | | `uint8 [200, 100, 50].cumsum(1)` | int64 `[200, 300, 350]` | uint8 **`[200, 44, 94]`** | | `int8 [100, 100, 50].cumsum(1)` | int64 `[100, 200, 250]` | int8 `[100, -56, -6]` | | `int32 [2³⁰, 2³⁰, 5].cumsum(1)` | int64 `[…, 2147483653]` | int32 `[…, -2147483643]` | | `uint8 [200, 100, 50].cumprod(1)` | int64 `[200, 20000, 1000000]` | uint8 `[200, 32, 64]` | | `bool [T, T, F].cumsum(0)` | int64 | `InternalError` (relax cumsum on bool) | Not a dtype-label difference: the values are wrong as soon as the sum leaves the narrow range, which for uint8/int8 is the second element. ### Fix `_cumulative_dtype` returns the explicit `dtype=` if given, otherwise `int64` for an integral or bool input and `None` (keep) for floats — torch's rule. Both converters use it. ### Verification `cumsum` and `cumprod` along both axes plus `cumsum(dtype=float32)`, over bool, uint8, int8, int16, int32, int64, float16, float32 and float64 inputs chosen to overflow the narrow types, each built with `relax.build(llvm)` and compared with torch on dtype and values: | | matched | wrong | raised | | --- | --- | --- | --- | | `main` | 25 | 16 | 4 | | this PR | **45** | 0 | 0 | The explicit-`dtype` variant and every float input were already correct and are unchanged. `test_frontend_from_exported_program.py` and `test_frontend_from_fx.py`: failure sets identical before and after apart from the new tests. `ruff check` / `ruff format --check` (v0.12.3, the version CI pins) clean. ### Tests - `test_cumsum_integer_input_accumulates_in_int64` — IR-level: a uint8 input emits `R.cumsum(x, axis=1, dtype="int64")`. - `test_cumsum_cumprod_integer_values` — `cumsum`, `cumprod` and `cumsum(dtype=float32)` over bool, uint8, int8, int32 and int64, asserting dtype and values. Against the previous head the bool, uint8, int8 and int32 rows fail (`assert 'uint8' == 'int64'` and the IR mismatch); int64 passes there and pins that it is untouched. Independent of #20372 / #20373 (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]
