sepcnt opened a new issue, #20340:
URL: https://github.com/apache/tvm/issues/20340
### Expected behavior
The default legalization of tirx.q_multiply_shift and
tirx.q_multiply_shift_per_axis should:
- Accept runtime multipliers and shifts, using the general lowering when the
constant-only optimization is
inapplicable.
- Return x for q_multiply_shift(x, 1 << 30, 31, 1).
- Accept the documented integer 0/1 flags in the per-axis operation.
### Actual behavior
Three related defects occur in src/target/intrin_rule.cc:
1. Runtime operands: get_int_value asserts that its argument is an IntImm or
Broadcast(IntImm). Runtime y, or runtime
s with y == 1 << 30, triggers an internal error instead of falling back
to QMultiplyShift.
2. Zero exponent: With y == 1 << 30 and s == 1, the fast path constructs a
rounding factor using 1 << -1, causing an
internal shift-range error.
3. Integer flags: The per-axis legalizer forwards is_lshift_required
directly to Select, whose condition must be
boolean. Passing integer 0 or 1 triggers an internal type check.
The fast path also assumes q == 31 without checking it.
These defects were reported independently by downstream TileLang fuzzing:
- Runtime multiplier/shift: https://github.com/tile-ai/tilelang/issues/3014
- Zero-exponent boundary: https://github.com/tile-ai/tilelang/issues/3003
- Integer per-axis flags: https://github.com/tile-ai/tilelang/issues/3006
The reproducer below confirms the issues directly in TVM, without TileLang
or GPU execution.
### Environment
- Apache TVM checkout: cc0f9f07c17c8118a781fdca55f7fe45f7de916a
- Windows, Python 3.12.13
- Built from source with clang-cl 22.1.8
- CUDA and LLVM disabled
- Pinned tvm-ffi built from source
### Steps to reproduce
```python
import tvm
from tvm import tirx
x = tirx.Var("x", "int32")
y = tirx.Var("y", "int32")
s = tirx.Var("s", "int32")
cases = [
("runtime multiplier", "tirx.q_multiply_shift", [x, y, 31, 1]),
("runtime shift", "tirx.q_multiply_shift", [x, 1 << 30, 31, s]),
("zero exponent", "tirx.q_multiply_shift", [x, 1 << 30, 31, 1]),
(
"integer flag",
"tirx.q_multiply_shift_per_axis",
[x, y, 0, 1, 31, 0, 1],
),
]
for name, op_name, args in cases:
call = tirx.call_intrin("int32", op_name, *args)
legalize = tvm.ir.Op.get(op_name).get_attr("default.FLegalize")
try:
print(name, legalize(call))
except tvm.error.InternalError as err:
print(name, err)
```
A localized fix can make constant extraction non-asserting, guard the fast
path appropriately, handle the zero
exponent explicitly, and normalize per-axis flags to boolean conditions.
### Triage
- needs-triage
- type: bug
--
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]