Description
ReorderTakeAfterMatmul rewrites a valid vector dot product into a scalar MatMul followed by take(axis=-1). A scalar has no valid axis, 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/reorder_take_after_matmul.cc is byte-identical at both revisions (SHA-256 80324bd216442834d9c181f2466fc6ff8aa5538b93e78251b59f863ce5a6e6ba)
- 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(
lhs: R.Tensor((3,), "float32"),
weight: R.Tensor((3,), "float32"),
indices: R.Tensor((3,), "int64"),
):
selected = R.take(weight, indices, axis=0, mode="clip")
return R.matmul(lhs, selected)
assert relax.analysis.check_well_formed(Module, check_ty=True)
relax.transform.ReorderTakeAfterMatmul()(Module)
Expected behavior
The pass should preserve the result or leave this unsupported rank pattern unchanged.
Actual behavior
The source is well formed, executes on LLVM, and returns a scalar. The pass raises:
ValueError: the input axis -1 is out of range. The input tensor has 0 dimensions
The adjacent rank-two weight control is rewritten successfully and its source and target outputs match.
Likely cause
The rewrite guard accepts a rank-one weight and indices tensor. The replacement MatMul therefore produces a scalar, but the pass constructs take(out_table, indices, matmul_ty->ndim - 1, mode), yielding axis -1 for a rank-zero value.
Duplicate check
I searched open and closed TVM issues and PRs for ReorderTakeAfterMatmul, vector weights, scalar outputs, and axis -1. #20198/#20331 concern out_dtype, and #20201/#20206 concern take mode; I did not find this rank-guard failure.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage
Description
ReorderTakeAfterMatmulrewrites a valid vector dot product into a scalar MatMul followed bytake(axis=-1). A scalar has no valid axis, so the pass raises while constructing the replacement.Environment
0.26.dev0, executed ate269315c90e3a061c9e1c77b370ce883b1b223f4mainchecked atd55759ec1ebed554d9782a52e43e5730b8a96656src/relax/transform/reorder_take_after_matmul.ccis byte-identical at both revisions (SHA-25680324bd216442834d9c181f2466fc6ff8aa5538b93e78251b59f863ce5a6e6ba)main, so the runtime reproduction below is from the executed revisionMinimal reproduction
Expected behavior
The pass should preserve the result or leave this unsupported rank pattern unchanged.
Actual behavior
The source is well formed, executes on LLVM, and returns a scalar. The pass raises:
The adjacent rank-two weight control is rewritten successfully and its source and target outputs match.
Likely cause
The rewrite guard accepts a rank-one weight and indices tensor. The replacement MatMul therefore produces a scalar, but the pass constructs
take(out_table, indices, matmul_ty->ndim - 1, mode), yielding axis-1for a rank-zero value.Duplicate check
I searched open and closed TVM issues and PRs for
ReorderTakeAfterMatmul, vector weights, scalar outputs, and axis-1. #20198/#20331 concernout_dtype, and #20201/#20206 concern take mode; I did not find this rank-guard failure.I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage