From 76c1708018d2bdcf4d5a207100ff8682be6cf09d Mon Sep 17 00:00:00 2001 From: Eric Lunderberg Date: Mon, 5 Feb 2024 22:07:41 +0000 Subject: [PATCH] [Transform][Bugfix] Handle non-composite lambda functions in FuseOps Prior to this commit, calling `FuseOpsByPattern` with `annotate_codegen=True` would cause an error when encountering a lambda function. This was caused by the `CompositeFunctionAnnotator` asserting that all `relax::Function` encountered must have the `kComposite` attribute. While this is true for all lambda functions produced by `FuseOpsByPattern`, the user may have defined other lambda functions as well. This commit updates `CompositeFunctionAnnotator` to ignore lambda functions that do not have a `kComposite` attribute. --- src/relax/transform/fuse_ops.cc | 8 ++- .../test_transform_fuse_ops_by_pattern.py | 54 +++++++++++++++++++ 2 files changed, 60 insertions(+), 2 deletions(-) diff --git a/src/relax/transform/fuse_ops.cc b/src/relax/transform/fuse_ops.cc index 0dbee3667061..5ead71f3b396 100644 --- a/src/relax/transform/fuse_ops.cc +++ b/src/relax/transform/fuse_ops.cc @@ -1238,10 +1238,14 @@ class CompositeFunctionAnnotator : public ExprMutator { Expr VisitExpr_(const FunctionNode* func_node) final { Function f_inner = Downcast(ExprMutator::VisitExpr_(func_node)); - auto composite_name = func_node->GetAttr(attr::kComposite); + + if (!func_node->GetAttr(attr::kComposite)) { + // This lambda function doesn't have `attr::kComposite`, so it + // was not produced by FuseOps. + return std::move(f_inner); + } f_inner = WithoutAttr(std::move(f_inner), tvm::relax::attr::kPrimitive); - ICHECK(composite_name); Array param_vars; Array params; diff --git a/tests/python/relax/test_transform_fuse_ops_by_pattern.py b/tests/python/relax/test_transform_fuse_ops_by_pattern.py index de356fd5480e..99ca117d65b6 100644 --- a/tests/python/relax/test_transform_fuse_ops_by_pattern.py +++ b/tests/python/relax/test_transform_fuse_ops_by_pattern.py @@ -530,6 +530,60 @@ def test_annotate_codegen(): ) +@pytest.mark.parametrize("annotate_codegen", [True, False]) +def test_no_op_if_no_patterns_match(annotate_codegen): + """If no matches occur, FuseOpsByPattern is a no-op""" + check( + Conv2dReLU, + [], + Conv2dReLU, + annotate_codegen=annotate_codegen, + ) + + +@pytest.mark.parametrize("annotate_codegen", [True, False]) +def test_unmatched_calls_may_include_lambda_functions(annotate_codegen): + """If no matches occur, FuseOpsByPattern is a no-op + + This is a regression test. Previous implementations of + CompositeFunctionAnnotator assumed that all lambda functions + resulted from FuseOps, and would contain the `kComposite` + attribute. + """ + + @tvm.script.ir_module + class Module: + @R.function + def main( + data: R.Tensor((1, 64, 56, 56), "float32"), + weight1: R.Tensor((64, 64, 3, 3), "float32"), + ): + with R.dataflow(): + conv1 = R.nn.relu(R.nn.conv2d(data, weight1, padding=(1, 1))) + R.output(conv1) + + return conv1 + + @R.function + def unrelated_function(A: R.Tensor([16, 16], dtype="float16")): + @R.function + def inner_func(B: R.Tensor([16, 16], dtype="float16")): + with R.dataflow(): + C = R.multiply(B, R.const(2, "float16")) + R.output(C) + return C + + D = inner_func(A) + return D + + check( + Module, + [], + Module, + annotate_codegen=annotate_codegen, + ) + + def test_compare_with_merge_composite_path(): x = relax.Var("x", relax.TensorStructInfo([10, 10], "float32")) y = relax.Var("y", relax.TensorStructInfo([10, 10], "float32"))