Skip to content

[Bug][Relax] AdjustMatmulOrder changes results for a rank-one middle operand #20492

Description

@Yuhx141

Description

AdjustMatmulOrder changes (a @ b) @ c into a @ (b @ c) when b and c are compile-time
parameters. This reassociation is not valid when b is rank one. Relax interprets the vector as a
column in a @ b, and as a row in b @ c. Both modules are well formed and executable, but they
compute different values.

Environment

  • Ubuntu 24.04.4 LTS, x86-64
  • TVM 0.26.dev0
  • Runtime reproduction: e269315c90e3a061c9e1c77b370ce883b1b223f4
  • LLVM 18.1.3, CPU bytecode VM

Current upstream main was 2393db61bbd0100f048b1a6d3c2519fb70914564 when checked on
2026-09-30. src/relax/transform/adjust_matmul_order.cc has SHA-256
a407a71e13eb07d13207d4d6c43e904811cb9195af044090326a24d14df8d34e at both revisions.
I did not have a binary build for that exact main revision, so the runtime reproduction below is
from the stated fixed revision; the current implementation is byte-identical.

Minimal reproducer

import numpy as np
import tvm
from tvm import relax
from tvm.script import ir as I
from tvm.script import relax as R


@I.ir_module
class Before:
    @R.function
    def main(
        a: R.Tensor([2, 2], "float32"),
        b: R.Tensor([2], "float32"),
        c: R.Tensor([2, 2], "float32"),
    ):
        R.func_attr({"num_input": 1})
        inner = R.matmul(a, b)
        return R.matmul(inner, c)


After = relax.transform.AdjustMatmulOrder()(Before)
assert relax.analysis.check_well_formed(Before, check_ty=True)
assert relax.analysis.check_well_formed(After, check_ty=True)

a = np.array([[1, 2], [3, 4]], dtype="float32")
b = np.array([5, 6], dtype="float32")
c = np.array([[7, 8], [9, 10]], dtype="float32")


def run(mod):
    executable = relax.build(
        mod,
        target=tvm.target.Target("llvm", host="llvm"),
        relax_pipeline="default",
        exec_mode="bytecode",
    )
    vm = relax.VirtualMachine(executable, tvm.cpu())
    return vm["main"](*(tvm.runtime.tensor(x) for x in (a, b, c))).numpy()


print(run(Before))  # [470. 526.]
print(run(After))   # [289. 667.]

The transformed module contains:

gv = R.matmul(b, c)
gv1 = R.matmul(a, gv)
return gv1

Expected behavior

The pass should leave this expression unchanged, or otherwise preserve rank-one MatMul orientation
semantics.

Actual behavior

Both source and target pass check_well_formed(..., check_ty=True) and run in the LLVM Relax VM.
The pass changes the output from [470, 526] to [289, 667], with maximum absolute difference
181.

The mismatch reproduces for square extents 2 through 7, with maximum absolute differences
181, 336, 2200, 9600, 32340, and 90944. Paired controls change only b from rank one
to rank two; all six controls are rewritten and have exact source/target equality.

Suspected cause

The compile-time grouping branch in adjust_matmul_order.cc lines 245-250 immediately returns
matmul(a, matmul(b, c)) when b and c are compile-time values. The later shape logic at lines
256-271 knows that a rank-one middle operand has different row/column orientation depending on the
association, but it is reached only after the early branch has returned.

Duplicate check

I searched existing TVM issues and pull requests for AdjustMatmulOrder and rank-one/vector cases.
I found related reports about batched shapes, permutations, and dtype handling, but none with the
same cause: reassociation changes rank-one MatMul orientation.

I can send a fix PR and help with follow-up testing if this diagnosis looks right.

Triage

  • needs-triage
  • type: bug

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions