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

   ### Problem
   
   When a binary op has a Python scalar operand, `_binary_op` builds the 
constant in the tensor's own dtype:
   
   ```python
   return lhs, relax.const(rhs, lhs.ty.dtype)
   ```
   
   For a float scalar against an integer or bool tensor that truncates the 
scalar before the op runs. torch promotes the other way — a Python scalar takes 
part in type promotion at a lower priority than a tensor and widens it only 
when its category is higher — so the two disagree on both dtype and values:
   
   | expression | tensor dtype | torch | frontend on `main` |
   | --- | --- | --- | --- |
   | `x * 0.5` | int64 `[1, 2, 3]` | float32 `[0.5, 1.0, 1.5]` | int64 **`[0, 
0, 0]`** |
   | `x + 0.5` | int64 | float32 `[1.5, 2.5, 3.5]` | int64 `[1, 2, 3]` |
   | `0.5 - x` | int64 | float32 `[-0.5, -1.5, -2.5]` | int64 `[-1, -2, -3]` |
   | `x < 1.5` | int64 | bool `[T, F, F]` | bool `[F, F, F]` |
   | `x * 2.5` | bool `[T, F, T]` | float32 `[2.5, 0, 2.5]` | bool `[T, F, T]` |
   | `x + 1` | bool | int64 `[2, 1, 2]` | `InternalError` (relax add on bool) |
   | `x ** 0.5` | int64 | float32 | `InternalError` (power only applies to 
float) |
   
   Nothing raises in the common rows, so `x * 0.5` on an integer tensor 
silently imports as a tensor of zeros. Tensor-tensor promotion (`int_tensor * 
float_tensor`) was already handled by `_promote_common_dtype` and is correct; 
only the scalar path was wrong.
   
   The rule torch applies is `torch.result_type(tensor, scalar)`:
   
   | tensor | bool scalar | int scalar | float scalar |
   | --- | --- | --- | --- |
   | bool | bool | **int64** | **float32** |
   | int8 / uint8 / int32 / int64 | tensor dtype | tensor dtype | **float32** |
   | float16 / bfloat16 / float32 / float64 | tensor dtype | tensor dtype | 
tensor dtype |
   
   ### Fix
   
   `_scalar_result_dtype` asks `torch.result_type` for the promoted dtype; the 
scalar branch of `promote_binary_op_args` casts the tensor when it has to widen 
and builds the constant in that dtype. The two `Constant`-vs-scalar dispatch 
branches pre-cast the scalar the same wrong way, and now go through the same 
path.
   
   ```python
   target = self._scalar_result_dtype(tensor.ty.dtype, scalar) or 
tensor.ty.dtype
   if str(tensor.ty.dtype) != str(target):
       tensor = self.block_builder.emit(relax.op.astype(tensor, target))
   return tensor, relax.const(scalar, target)
   ```
   
   This is the shared `_binary_op`, so it covers 
`add`/`sub`/`mul`/`pow`/`remainder`/the six 
comparisons/`maximum`/`minimum`/`atan2`/`logaddexp` in both `from_fx` and 
`from_exported_program`.
   
   ### Verification
   
   A sweep of 18 binary ops × 8 tensor dtypes (bool, uint8, int8, int32, int64, 
float16, float32, float64) × 6 Python scalars (`True`, `2`, `-3`, `0.5`, `1.5`, 
`-0.5`), each built with `relax.build(llvm)` and compared with torch on 
**result dtype and values** — a shape or dtype-only check would not see `[0, 0, 
0]`:
   
   | | matched | wrong dtype or values | raised |
   | --- | --- | --- | --- |
   | `main` | 477 | 149 | 203 |
   | this PR | **645** | 52 | 132 |
   
   864 programs, 829 accepted by torch. Diffing the two failure lists per case: 
**no case fails with this change that did not fail before it; 168 repaired.**
   
   What is left is not this change:
   
   - **157 are the division family** (`x / 2`, `x // 2`, `torch.div(x, 2, 
rounding_mode=…)`). True division of two integers has to give a float even for 
an *int* scalar — a rule on top of `result_type` — and `div.Tensor_mode` builds 
its constant with no dtype at all, which is why `x // 2` currently raises on 
every float tensor. That is a separate PR on top of this one.
   - 12 are relax rejecting arithmetic or comparison on a bool tensor with a 
bool scalar (`bool_tensor + True`), 11 are a uint8 tensor against a negative 
scalar (torch wraps modulo 256, `relax.const(-3, "uint8")` raises), 6 are `x ** 
True` on integer tensors. All pre-existing and out of scope.
   
   `test_linspace`'s expected IR encoded the truncation: torch's `linspace` 
decomposition splits the range at `lt.Scalar(arange, 4.5)`, which the frontend 
emitted as `R.less(i, R.const(4, "int64"))`. It now promotes the index to 
float32 and compares against `4.5`; the expected module is updated. (For 
`linspace(0, 1, 9)` both branches of the `where` compute the same values, which 
is why the truncated split never showed up numerically.)
   
   `tests/python/relax/test_frontend_from_exported_program.py` and 
`test_frontend_from_fx.py`: the failure *sets* are identical before and after 
apart from the new tests (9 and 15 pre-existing failures in my environment, all 
`test_dtypes` / `test_prod` and friends). `ruff check` and `ruff format 
--check` are clean on both touched files.
   
   ### Tests
   
   - `test_binary_python_scalar_promotes_the_tensor` — IR-level: `int64 + 0.5` 
emits `astype` plus `R.const(0.5, "float32")`.
   - `test_binary_python_scalar_promotion_values` — `add`/`mul`/`lt`/`ge`/`eq`, 
both operand orders, over the eight `(dtype, scalar)` rows of the table above, 
asserting the result dtype and the values.
   - `test_binary_python_scalar_promotion_sub_pow_remainder` — the three ops 
torch does not define on a bool tensor, on their own case list.
   
   22 of the 44 parametrised cases fail against the previous head with the 
messages above (`result dtype int64, torch gives torch.float32`); the other 22 
are the rows where the tensor's dtype wins, and pin that nothing there moved.
   
   ---
   
   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