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]

Reply via email to