Skip to content
Open
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
31 changes: 29 additions & 2 deletions src/relax/transform/fuse_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -870,6 +870,7 @@ class OperatorFusor : public ExprMutator {

BindingBlock VisitBindingBlock_(const DataflowBlockNode* block) final {
group2func_.clear();
CollectGroupsWithCall(block->bindings);

// Step 1. Collect the bindings for each grouped function.
CollectFuncBindings(block->bindings);
Expand Down Expand Up @@ -904,7 +905,7 @@ class OperatorFusor : public ExprMutator {
// Case 1. If the binding is the only binding in its group, recurse into it and emit the
// transformed binding as usual.
Group* group = GetGroupFromBinding(binding);
if (group->num_nodes == 1 && group->attrs.empty()) {
if (!NeedsGroupedFunction(group)) {
VisitBinding(binding);
continue;
}
Expand Down Expand Up @@ -987,6 +988,30 @@ class OperatorFusor : public ExprMutator {
return builder_->EndBlock();
}

/*! \brief Record the groups that contain at least one call, e.g. a call to `relax.call_tir`. */
void CollectGroupsWithCall(const ffi::Array<Binding>& bindings) {
groups_with_call_.clear();
for (const Binding& binding : bindings) {
const auto* var_binding = binding.as<VarBindingNode>();
if (var_binding && var_binding->value->IsInstance<CallNode>()) {
groups_with_call_.insert(GetGroupFromBinding(binding));
}
}
}

/*!
* \brief Check whether a grouped function should be created for the group.
* \note A group with a single binding needs no function. Neither does a group without any call,
* e.g. a chain of TupleGetItem: there is nothing to fuse, and the resulting function would
* contain no PrimFunc call for FuseTIR to lower.
*/
bool NeedsGroupedFunction(Group* group) const {
if (!group->attrs.empty()) {
return true;
}
return group->num_nodes > 1 && groups_with_call_.count(group);
}

/*!
* \brief Collect the bindings for each grouped function and update the information of the grouped
* function
Expand All @@ -997,7 +1022,7 @@ class OperatorFusor : public ExprMutator {
for (const Binding& binding : bindings) {
// If the binding is the only binding in its group, there is no need to create a new function.
Group* group = GetGroupFromBinding(binding);
if (group->num_nodes == 1 && group->attrs.empty()) {
if (!NeedsGroupedFunction(group)) {
continue;
}
// Add the binding to the grouped function it's in, and update the function information
Expand Down Expand Up @@ -1130,6 +1155,8 @@ class OperatorFusor : public ExprMutator {
support::Arena arena_;
/*! \brief The group assignment map. */
GroupMap obj2group_;
/*! \brief The groups in the current binding block that contain at least one call. */
std::unordered_set<Group*> groups_with_call_;
/*! \brief Internal function information map. */
std::unordered_map<Group*, FunctionCreator> group2func_;
/*! \brief Bindings visible while rewriting the current Relax function. */
Expand Down
21 changes: 21 additions & 0 deletions tests/python/relax/test_transform_fuse_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,6 +377,27 @@ def expected(dim: int):
_check(before(dim), expected(dim))


def test_tuple_get_item_chain_without_call():
"""Chained TupleGetItem without any call has nothing to fuse, so it stays in main.

Fusing it used to create a function without any call_tir, which FuseTIR rejects.
"""

@I.ir_module
class Module:
@R.function
def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"):
with R.dataflow():
inner = (x, x)
outer = (inner, x)
a = outer[0]
b = a[0]
R.output(b)
return b

_check(Module, Module)


def test_tuple_intermediate():
def before():
bb = relax.BlockBuilder()
Expand Down