Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions src/relax/transform/fuse_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1238,10 +1238,14 @@ class CompositeFunctionAnnotator : public ExprMutator {

Expr VisitExpr_(const FunctionNode* func_node) final {
Function f_inner = Downcast<Function>(ExprMutator::VisitExpr_(func_node));
auto composite_name = func_node->GetAttr<String>(attr::kComposite);

if (!func_node->GetAttr<String>(attr::kComposite)) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Are non-composite functions visited?

auto new_func = Downcast<Function>(VisitExpr(func));
here we only visit composite functions

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The PatternBasedPartitioner only visits non-composite functions, produces a composite function for each pattern match, and updates the non-composite function to call the newly-generated composite function. Afterwards, the call to CompositeFunctionAnnotator is called. This visits only non-composite functions, finds any relax-to-relax function calls, and asserts that the callee is composite.

The callee will be composite for every function call generated by PatternBasedPartitioner, but that doesn't guarantee that all relax-to-relax function calls have a composite callee. If the IRModule contains a relax-to-relax call prior to PatternBasedPartitioner, that callee may be non-composite. This IRModule would be entirely legal, but would trigger the assert in CompositeFunctionAnnotator.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

So, the problem isn't with calls to inner functions as on line 1224, but with calls to other functions within the IRModule.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think the problem is if the callee is not a global var, the callee function will still be visited, so the fix makes sense to me

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Whoops, you're right on that one. It's if there is a inner function in the input IRModule. (Apologies, trying to track too many PRs at one time.)

// 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<Var> param_vars;
Array<Expr> params;
Expand Down
54 changes: 54 additions & 0 deletions tests/python/relax/test_transform_fuse_ops_by_pattern.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"))
Expand Down