diff --git a/src/relax/transform/fuse_ops.cc b/src/relax/transform/fuse_ops.cc index b0eeba399e90..c6ea63074581 100644 --- a/src/relax/transform/fuse_ops.cc +++ b/src/relax/transform/fuse_ops.cc @@ -1195,7 +1195,8 @@ class CompositeFunctionAnnotator : public ExprMutator { auto all_functions = mod->functions; for (const auto& entry : all_functions) { if (const auto* func = entry.second.as()) { - if (func->GetAttr(attr::kComposite).defined()) { + if (func->GetAttr(attr::kComposite).defined() || + func->GetAttr(attr::kCodegen).defined()) { continue; } auto new_body = VisitExpr(func->body); @@ -1266,6 +1267,13 @@ IRModule FuseOpsByPattern(const tvm::Array& patterns, if (entry.second->IsInstance()) { continue; } + const FunctionNode* function = entry.second.as(); + if (function->GetAttr(attr::kPrimitive).defined() || + function->GetAttr(attr::kComposite).defined() || + function->GetAttr(attr::kCodegen).defined()) { + continue; + } + auto map = PatternBasedPartitioner::Run(pattern->name, pattern->pattern, pattern->annotation_patterns, pattern->check.value_or(nullptr), entry.second, 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 bd434864a081..de356fd5480e 100644 --- a/tests/python/relax/test_transform_fuse_ops_by_pattern.py +++ b/tests/python/relax/test_transform_fuse_ops_by_pattern.py @@ -1046,5 +1046,14 @@ def main( assert "fused_relax_permute_dims_relax_matmul_cublas" in func_names # add is not fused +def test_multple_runs(): + check( + Conv2dReLU_composite_annotated, + [("dnnl.conv2d_relu", conv2d_relu_pat)], + Conv2dReLU_composite_annotated, + annotate_codegen=True, + ) + + if __name__ == "__main__": pytest.main([__file__])