This is an automated email from the ASF dual-hosted git repository.
masahi pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/unity by this push:
new 44c849f144 [Unity] Reset match state when backtracking (#14984)
44c849f144 is described below
commit 44c849f144aab4f1f4759fa7faf8abcb8547983f
Author: Lite Ye <[email protected]>
AuthorDate: Tue May 30 16:25:49 2023 -0400
[Unity] Reset match state when backtracking (#14984)
Reset match state when backtracking
---
src/relax/ir/dataflow_matcher.cc | 1 +
tests/python/relax/test_dataflow_pattern.py | 110 ++++++++++++++++++++++++++--
2 files changed, 106 insertions(+), 5 deletions(-)
diff --git a/src/relax/ir/dataflow_matcher.cc b/src/relax/ir/dataflow_matcher.cc
index 81cccec86d..2d06ce1fb9 100644
--- a/src/relax/ir/dataflow_matcher.cc
+++ b/src/relax/ir/dataflow_matcher.cc
@@ -700,6 +700,7 @@ static std::optional<MatchState> MatchTree(
return new_match;
}
// Recursive matching has failed, backtrack.
+ new_match = current_match;
continue;
}
}
diff --git a/tests/python/relax/test_dataflow_pattern.py
b/tests/python/relax/test_dataflow_pattern.py
index e32314cf39..cbb19c6743 100644
--- a/tests/python/relax/test_dataflow_pattern.py
+++ b/tests/python/relax/test_dataflow_pattern.py
@@ -16,13 +16,14 @@
# under the License.
import pytest
-import tvm.testing
-from tvm import relay
-from tvm.relax.dpl import *
+import tvm.testing
+from tvm import relax as rx
+from tvm import relay, tir
from tvm.relax.analysis import get_var2val
-from tvm import relax as rx, tir
-from tvm.script import relax as R, tir as T
+from tvm.relax.dpl import *
+from tvm.script import relax as R
+from tvm.script import tir as T
@tvm.script.ir_module
@@ -1182,5 +1183,104 @@ def test_combine_matmul_emit_order():
rx.build(mod, target="llvm")
+def test_combine_transposed_matmul_twice():
+ @R.function
+ def main(
+ x1: R.Tensor((2, 1024, 640), "float32"),
+ x2: R.Tensor((2, 1024, 640), "float32"),
+ w0: R.Tensor((640, 640), "float32"),
+ w1: R.Tensor((640, 640), "float32"),
+ w2: R.Tensor((640, 640), "float32"),
+ w3: R.Tensor((640, 640), "float32"),
+ ) -> R.Tensor:
+ with R.dataflow():
+ w0_t = R.permute_dims(w0, axes=None)
+ lv0 = R.matmul(x1, w0_t)
+ w1_t = R.permute_dims(w1, axes=None)
+ lv1 = R.matmul(x1, w1_t)
+ w2_t = R.permute_dims(w2, axes=None)
+ lv2 = R.matmul(x2, w2_t)
+ w3_t = R.permute_dims(w3, axes=None)
+ lv3 = R.matmul(x2, w3_t)
+ out = (lv0, lv1, lv2, lv3)
+ R.output(out)
+ return out
+
+ @R.function
+ def expected(
+ x1: R.Tensor((2, 1024, 640), dtype="float32"),
+ x2: R.Tensor((2, 1024, 640), dtype="float32"),
+ w0: R.Tensor((640, 640), dtype="float32"),
+ w1: R.Tensor((640, 640), dtype="float32"),
+ w2: R.Tensor((640, 640), dtype="float32"),
+ w3: R.Tensor((640, 640), dtype="float32"),
+ ) -> R.Tensor:
+ with R.dataflow():
+ lv: R.Tensor((1280, 640), dtype="float32") = R.concat((w0, w1),
axis=0)
+ lv1: R.Tensor((640, 1280), dtype="float32") = R.permute_dims(lv,
axes=None)
+ lv2: R.Tensor((2, 1024, 1280), dtype="float32") = R.matmul(x1,
lv1, out_dtype="void")
+ lv3: R.Tuple(
+ R.Tensor((2, 640, 1280), dtype="float32"),
+ R.Tensor((2, 384, 1280), dtype="float32"),
+ R.Tensor((2, 0, 1280), dtype="float32"),
+ ) = R.split(lv2, indices_or_sections=[640, 1280], axis=1)
+ lv0: R.Tensor((2, 640, 1280), dtype="float32") = lv3[0]
+ lv1_1: R.Tensor((2, 384, 1280), dtype="float32") = lv3[1]
+ lv_1: R.Tensor((1280, 640), dtype="float32") = R.concat((w2, w3),
axis=0)
+ lv1_2: R.Tensor((640, 1280), dtype="float32") =
R.permute_dims(lv_1, axes=None)
+ lv2_1: R.Tensor((2, 1024, 1280), dtype="float32") = R.matmul(
+ x2, lv1_2, out_dtype="void"
+ )
+ lv3_1: R.Tuple(
+ R.Tensor((2, 640, 1280), dtype="float32"),
+ R.Tensor((2, 384, 1280), dtype="float32"),
+ R.Tensor((2, 0, 1280), dtype="float32"),
+ ) = R.split(lv2_1, indices_or_sections=[640, 1280], axis=1)
+ lv2_1_1: R.Tensor((2, 640, 1280), dtype="float32") = lv3_1[0]
+ lv3_1_1: R.Tensor((2, 384, 1280), dtype="float32") = lv3_1[1]
+ out: R.Tuple(
+ R.Tensor((2, 640, 1280), dtype="float32"),
+ R.Tensor((2, 384, 1280), dtype="float32"),
+ R.Tensor((2, 640, 1280), dtype="float32"),
+ R.Tensor((2, 384, 1280), dtype="float32"),
+ ) = (lv0, lv1_1, lv2_1_1, lv3_1_1)
+ R.output(out)
+ return out
+
+ with PatternContext() as ctx:
+ inp_pat = wildcard()
+ w1_pat = wildcard()
+ w2_pat = wildcard()
+ matmul1 = is_op("relax.matmul")(inp_pat,
is_op("relax.permute_dims")(w1_pat))
+ matmul2 = is_op("relax.matmul")(inp_pat,
is_op("relax.permute_dims")(w2_pat))
+
+ def rewriter(matchings, _):
+ inp = matchings[inp_pat]
+ w1 = matchings[w1_pat]
+ w2 = matchings[w2_pat]
+
+ concat = R.concat([w1, w2], axis=0)
+ matmul = R.matmul(inp, R.permute_dims(concat))
+ sections = [w1.struct_info.shape[0], w1.struct_info.shape[0] +
w2.struct_info.shape[0]]
+
+ chunks = R.split(matmul, sections, 1)
+
+ return {
+ matchings[matmul1]: chunks[0],
+ matchings[matmul2]: chunks[1],
+ }
+
+ rewritten = rewrite_bindings(ctx, rewriter, main)
+ print(rewritten.script())
+ tvm.ir.assert_structural_equal(rewritten, expected)
+
+ # make sure it builds
+ mod = tvm.IRModule()
+ mod["main"] = rewritten
+ mod = rx.transform.LegalizeOps()(mod)
+
+ rx.build(mod, target="llvm")
+
+
if __name__ == "__main__":
tvm.testing.main()