Description
ExpandMatmulOfSum rewrites a valid matmul(x, add(a, b)) when one add operand is broadcast along the MatMul reduction axis. The sum has the required extent, but one distributed MatMul does not, so the pass raises while constructing the replacement.
Environment
- TVM
0.26.dev0, executed at e269315c90e3a061c9e1c77b370ce883b1b223f4
- Ubuntu 24.04 x86-64, LLVM 18.1.3 CPU target
- Current upstream
main checked at d55759ec1ebed554d9782a52e43e5730b8a96656
src/relax/transform/expand_matmul_of_sum.cc is byte-identical at both revisions (SHA-256 a94db602a0b5768478dae2fb04ad5998a9a4111e2669339a18236a5e3848b946)
- I did not have a binary built from current
main, so the runtime reproduction below is from the executed revision
Minimal reproduction
from tvm import relax
from tvm.script import ir as I
from tvm.script import relax as R
@I.ir_module
class Module:
@R.function
def main(
x: R.Tensor((3,), "float32"),
a: R.Tensor((3, 2), "float32"),
b: R.Tensor((1, 2), "float32"),
):
weight = R.add(a, b)
return R.matmul(x, weight)
assert relax.analysis.check_well_formed(Module, check_ty=True)
relax.transform.ExpandMatmulOfSum()(Module)
Expected behavior
The pass should leave this expression unchanged or produce a valid equivalent expression.
Actual behavior
The source is well formed and executes on LLVM, but the pass tries to construct R.matmul(x, b) and raises:
ValueError: Matmul requires the reduction length of the operands to be equal ... 3 and 1 are not equal.
A control with b shaped (3, 2) is transformed successfully and matches the source across six deterministic input regimes.
Likely cause
The rewriter distributes MatMul over Add without checking whether each Add operand independently has the required reduction extent. Broadcasting makes the Add legal while making one distributed MatMul illegal.
Duplicate check
I searched open and closed TVM issues and PRs for ExpandMatmulOfSum, reduction-axis broadcasting, and the exception. I did not find the same cause. #20198 and #20331 concern explicit out_dtype, not broadcast legality.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage
Description
ExpandMatmulOfSumrewrites a validmatmul(x, add(a, b))when one add operand is broadcast along the MatMul reduction axis. The sum has the required extent, but one distributed MatMul does not, so the pass raises while constructing the replacement.Environment
0.26.dev0, executed ate269315c90e3a061c9e1c77b370ce883b1b223f4mainchecked atd55759ec1ebed554d9782a52e43e5730b8a96656src/relax/transform/expand_matmul_of_sum.ccis byte-identical at both revisions (SHA-256a94db602a0b5768478dae2fb04ad5998a9a4111e2669339a18236a5e3848b946)main, so the runtime reproduction below is from the executed revisionMinimal reproduction
Expected behavior
The pass should leave this expression unchanged or produce a valid equivalent expression.
Actual behavior
The source is well formed and executes on LLVM, but the pass tries to construct
R.matmul(x, b)and raises:A control with
bshaped(3, 2)is transformed successfully and matches the source across six deterministic input regimes.Likely cause
The rewriter distributes MatMul over Add without checking whether each Add operand independently has the required reduction extent. Broadcasting makes the Add legal while making one distributed MatMul illegal.
Duplicate check
I searched open and closed TVM issues and PRs for
ExpandMatmulOfSum, reduction-axis broadcasting, and the exception. I did not find the same cause. #20198 and #20331 concern explicitout_dtype, not broadcast legality.I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage