From 76071589700a2b3675b35eda59c65c8432f6d47c Mon Sep 17 00:00:00 2001 From: wwoosshh <122337168+wwoosshh@users.noreply.github.com> Date: Wed, 7 Oct 2026 16:04:37 +0900 Subject: [PATCH] [Fix][Relax] Skip FuseOps groups that have no call to fuse FuseOps treats TupleGetItem as injective, so two chained TupleGetItem bindings, e.g. `a = outer[0]; b = a[0]`, form a group with two nodes. Only single-node groups were left unfused, so a Primitive function was created for this group. FunctionCreator then moves the first TupleGetItem to the caller to pass the used tuple field directly, which leaves a function with no call_tir and the default name "fused". FuseTIR names the fused PrimFunc after the PrimFuncs it calls, and rejects such a function with `Check failed: (func_info_.global_name != "fused")`. A group without any call has nothing to fuse, so keep its bindings as they are, the same way as a single-node group. Groups created by FuseOpsByPattern carry attributes and are not affected. Fixes #20194 Generated-by: Claude Code (Claude Opus 5.5) --- src/relax/transform/fuse_ops.cc | 31 +++++++++++++++++-- tests/python/relax/test_transform_fuse_ops.py | 21 +++++++++++++ 2 files changed, 50 insertions(+), 2 deletions(-) diff --git a/src/relax/transform/fuse_ops.cc b/src/relax/transform/fuse_ops.cc index 4bdeef9d1bfe..922f9e86b0db 100644 --- a/src/relax/transform/fuse_ops.cc +++ b/src/relax/transform/fuse_ops.cc @@ -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); @@ -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; } @@ -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& bindings) { + groups_with_call_.clear(); + for (const Binding& binding : bindings) { + const auto* var_binding = binding.as(); + if (var_binding && var_binding->value->IsInstance()) { + 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 @@ -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 @@ -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 groups_with_call_; /*! \brief Internal function information map. */ std::unordered_map group2func_; /*! \brief Bindings visible while rewriting the current Relax function. */ diff --git a/tests/python/relax/test_transform_fuse_ops.py b/tests/python/relax/test_transform_fuse_ops.py index 3ca2d32ea82e..2b6829b77d71 100644 --- a/tests/python/relax/test_transform_fuse_ops.py +++ b/tests/python/relax/test_transform_fuse_ops.py @@ -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()