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()

Reply via email to