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]