This is an automated email from the ASF dual-hosted git repository.
lunderberg pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 24cd93df8b [Relax] Fix segfault in rewrite_bindings for MatchCast node
(#17226)
24cd93df8b is described below
commit 24cd93df8b70dab4791cd383e542e9f697a3af0b
Author: Eric Lunderberg <[email protected]>
AuthorDate: Thu Aug 1 08:20:35 2024 -0500
[Relax] Fix segfault in rewrite_bindings for MatchCast node (#17226)
Prior to this commit, the `tvm.relax.dpl.rewrite_bindings` utility
would segfault if its input contained a `DataflowBlock` whose first
binding was a `MatchCast`.
The root cause is use of an unintialized `const VarNode* cur_user_;`
when collecting the variable usage. This variable is only initialized
for `VarBinding` nodes, and may be used uninitialized if a `MatchCast`
node is encountered before the first `VarBinding`. This uninitialized
value is later dereferenced during while pattern-matching, causing a
segfault.
This commit provides a default value of `nullptr` for
`MatcherUseDefAnalysis::cur_user_`, preventing the segfault.
---
src/relax/ir/dataflow_block_rewriter.cc | 2 +-
tests/python/relax/test_dataflow_pattern.py | 109 +++++++++++++++++++---------
2 files changed, 75 insertions(+), 36 deletions(-)
diff --git a/src/relax/ir/dataflow_block_rewriter.cc
b/src/relax/ir/dataflow_block_rewriter.cc
index fb08dfe96a..88efad86cf 100644
--- a/src/relax/ir/dataflow_block_rewriter.cc
+++ b/src/relax/ir/dataflow_block_rewriter.cc
@@ -49,7 +49,7 @@ class MatcherUseDefAnalysis : public relax::ExprVisitor {
// caller -> callee table.
std::map<const VarNode*, std::vector<const VarNode*>> caller2callees;
- const VarNode* cur_user_;
+ const VarNode* cur_user_ = nullptr;
void VisitBinding_(const VarBindingNode* binding) override {
// init
diff --git a/tests/python/relax/test_dataflow_pattern.py
b/tests/python/relax/test_dataflow_pattern.py
index f67b0530ca..03a3beb2f2 100644
--- a/tests/python/relax/test_dataflow_pattern.py
+++ b/tests/python/relax/test_dataflow_pattern.py
@@ -1053,9 +1053,17 @@ def test_attention_fake_qkv():
assert ctx.match_dfb(dfb) is None
-def get_qkv_proj_rewriter(
- inp_pat, Q_weight_pat, K_weight_pat, V_weight_pat, matmul1, matmul2,
matmul3
-):
+def get_qkv_proj_rewriter():
+ with PatternContext() as ctx:
+ inp_pat = wildcard()
+ Q_weight_pat = wildcard()
+ K_weight_pat = wildcard()
+ V_weight_pat = wildcard()
+
+ matmul1 = is_op("relax.matmul")(inp_pat, Q_weight_pat)
+ matmul2 = is_op("relax.matmul")(inp_pat, K_weight_pat)
+ matmul3 = is_op("relax.matmul")(inp_pat, V_weight_pat)
+
def qkv_proj_rewriter(matchings, _):
inp = matchings[inp_pat]
Q_weight = matchings[Q_weight_pat]
@@ -1071,7 +1079,7 @@ def get_qkv_proj_rewriter(
return {matchings[matmul1]: Q, matchings[matmul2]: K,
matchings[matmul3]: V}
- return qkv_proj_rewriter
+ return ctx, qkv_proj_rewriter
def test_combine_matmul_twice():
@@ -1123,21 +1131,63 @@ def test_combine_matmul_twice():
R.output(out)
return out
- with PatternContext() as ctx:
- inp_pat = wildcard()
- Q_weight_pat = wildcard()
- K_weight_pat = wildcard()
- V_weight_pat = wildcard()
+ ctx, rewriter = get_qkv_proj_rewriter()
+ rewritten = rewrite_bindings(ctx, rewriter, qkv_x2)
+ tvm.ir.assert_structural_equal(rewritten, expected)
- matmul1 = is_op("relax.matmul")(inp_pat, Q_weight_pat)
- matmul2 = is_op("relax.matmul")(inp_pat, K_weight_pat)
- matmul3 = is_op("relax.matmul")(inp_pat, V_weight_pat)
- rewriter = get_qkv_proj_rewriter(
- inp_pat, Q_weight_pat, K_weight_pat, V_weight_pat, matmul1,
matmul2, matmul3
- )
- rewritten = rewrite_bindings(ctx, rewriter, qkv_x2)
- tvm.ir.assert_structural_equal(rewritten, expected)
+def test_dataflow_may_start_with_match_cast():
+ """Inputs to rewrite_bindings may contain R.match_cast
+
+ This is a regression test. In previous implementations, applying
+ `rewrite_bindings` when `R.match_cast` is the first binding of a
+ `R.dataflow` block would cause a segfault.
+
+ """
+
+ @R.function(private=True)
+ def before(
+ x_untyped: R.Tensor,
+ w0_untyped: R.Tensor,
+ w1_untyped: R.Tensor,
+ w2_untyped: R.Tensor,
+ ):
+ with R.dataflow():
+ x = R.match_cast(x_untyped, R.Tensor((2, 1024, 640), "float32"))
+ w0 = R.match_cast(w0_untyped, R.Tensor((640, 640), "float32"))
+ w1 = R.match_cast(w1_untyped, R.Tensor((640, 640), "float32"))
+ w2 = R.match_cast(w2_untyped, R.Tensor((640, 640), "float32"))
+ out_0 = R.matmul(x, w0)
+ out_1 = R.matmul(x, w1)
+ out_2 = R.matmul(x, w2)
+ out = (out_0, out_1, out_2)
+ R.output(out)
+ return out
+
+ @R.function(private=True)
+ def expected(
+ x_untyped: R.Tensor,
+ w0_untyped: R.Tensor,
+ w1_untyped: R.Tensor,
+ w2_untyped: R.Tensor,
+ ):
+ with R.dataflow():
+ x = R.match_cast(x_untyped, R.Tensor((2, 1024, 640), "float32"))
+ w0 = R.match_cast(w0_untyped, R.Tensor((640, 640), "float32"))
+ w1 = R.match_cast(w1_untyped, R.Tensor((640, 640), "float32"))
+ w2 = R.match_cast(w2_untyped, R.Tensor((640, 640), "float32"))
+ w_concat = R.concat((w0, w1, w2), axis=1)
+ out_concat = R.matmul(x, w_concat)
+ out_0 = R.strided_slice(out_concat, axes=[2], begin=[0], end=[640])
+ out_1 = R.strided_slice(out_concat, axes=[2], begin=[640],
end=[1280])
+ out_2 = R.strided_slice(out_concat, axes=[2], begin=[1280],
end=[1920])
+ out = (out_0, out_1, out_2)
+ R.output(out)
+ return out
+
+ ctx, rewriter = get_qkv_proj_rewriter()
+ rewritten = rewrite_bindings(ctx, rewriter, before)
+ tvm.ir.assert_structural_equal(rewritten, expected)
def test_combine_matmul_emit_order():
@@ -1181,27 +1231,16 @@ def test_combine_matmul_emit_order():
R.output(out)
return out
- with PatternContext() as ctx:
- inp_pat = wildcard()
- Q_weight_pat = wildcard()
- K_weight_pat = wildcard()
- V_weight_pat = wildcard()
+ ctx, rewriter = get_qkv_proj_rewriter()
- matmul1 = is_op("relax.matmul")(inp_pat, Q_weight_pat)
- matmul2 = is_op("relax.matmul")(inp_pat, K_weight_pat)
- matmul3 = is_op("relax.matmul")(inp_pat, V_weight_pat)
+ rewritten = rewrite_bindings(ctx, rewriter, main)
+ tvm.ir.assert_structural_equal(rewritten, expected)
- rewriter = get_qkv_proj_rewriter(
- inp_pat, Q_weight_pat, K_weight_pat, V_weight_pat, matmul1,
matmul2, matmul3
- )
- rewritten = rewrite_bindings(ctx, rewriter, main)
- tvm.ir.assert_structural_equal(rewritten, expected)
-
- # make sure it builds
- mod = tvm.IRModule()
- mod["main"] = rewritten
+ # make sure it builds
+ mod = tvm.IRModule()
+ mod["main"] = rewritten
- rx.build(mod, target="llvm")
+ rx.build(mod, target="llvm")
def test_combine_transposed_matmul_twice():