This is an automated email from the ASF dual-hosted git repository.

tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/main by this push:
     new 06a8102d4d [Fix][Relax] Preserve tensor-derived symbols during fusion 
(#20139)
06a8102d4d is described below

commit 06a8102d4d3adeaff903ddf404c9a455c0d1b0d4
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Thu Aug 27 22:47:52 2026 -0700

    [Fix][Relax] Preserve tensor-derived symbols during fusion (#20139)
    
    `FuseOps` may lift a symbolic primitive argument into an independent
    scalar parameter before the complete fused-function boundary is known.
    This can sever its relationship with tensor or shape types, making the
    result argument-order dependent and potentially producing an invalid or
    redundant ABI. This PR classifies symbolic primitive arguments after
    collecting the full fused-function boundary, preserving their
    relationship with tensor and shape types regardless of argument order.
---
 src/relax/transform/fuse_ops.cc               |  85 ++++-
 tests/python/relax/test_transform_fuse_ops.py | 472 ++++++++++++++++++++++++++
 2 files changed, 547 insertions(+), 10 deletions(-)

diff --git a/src/relax/transform/fuse_ops.cc b/src/relax/transform/fuse_ops.cc
index d2f9870a38..5ebe36373b 100644
--- a/src/relax/transform/fuse_ops.cc
+++ b/src/relax/transform/fuse_ops.cc
@@ -522,7 +522,11 @@ class FunctionCreator : public ExprMutator {
       params_.insert(params_.begin() + param_idx, item_params.begin(), 
item_params.end());
     }
 
-    // Step 3. Visit each binding and collect outputs one by one.
+    // Step 3. Now that the complete function boundary is known, inline 
primitive arguments whose
+    // symbols participate in the boundary's tensor/shape types.
+    InlineShapeDependentPrimArgs();
+
+    // Step 4. Visit each binding and collect outputs one by one.
     ffi::Array<Expr> outputs(output_vars_.size(), Expr());
     for (const Binding& binding : bindings_) {
       // Special handing for TupleGetItem.
@@ -554,7 +558,7 @@ class FunctionCreator : public ExprMutator {
       }
     }
 
-    // Step 4. Finish constructing the new block.
+    // Step 5. Finish constructing the new block.
     BindingBlock new_block = builder_->EndBlock();
     if (outputs.empty()) {
       // If the result is not used outside
@@ -606,6 +610,53 @@ class FunctionCreator : public ExprMutator {
     return std::nullopt;
   }
 
+  Expr ResolveOuterBinding(const Expr& expr) {
+    Expr bound_value = expr;
+    std::unordered_set<const VarNode*> visited;
+    while (const auto* current_var = bound_value.as<VarNode>()) {
+      if (!visited.insert(current_var).second) break;
+      auto it = outer_bindings_.find(ffi::GetRef<Var>(current_var));
+      if (it == outer_bindings_.end()) break;
+      bound_value = (*it).second;
+    }
+    return bound_value;
+  }
+
+  void InlineShapeDependentPrimArgs() {
+    ffi::Array<Type> boundary_types = params_.Map([](const Var& param) { 
return GetType(param); });
+    for (const VarNode* output_var : output_vars_) {
+      boundary_types.push_back(GetType(ffi::GetRef<Var>(output_var)));
+    }
+
+    std::unordered_set<tirx::Var> boundary_shape_vars;
+    for (const tirx::Var& var : TIRVarsInType(TupleType(boundary_types))) {
+      boundary_shape_vars.insert(var);
+    }
+
+    for (const Expr& argument : deferred_prim_args_) {
+      auto it = std::find_if(arguments_.begin(), arguments_.end(),
+                             [&](const Expr& candidate) { return 
candidate.same_as(argument); });
+      TVM_FFI_ICHECK(it != arguments_.end());
+
+      Expr bound_value = ResolveOuterBinding(argument);
+      bool inline_argument =
+          bound_value.same_as(argument) && IsShapeDependentPrimExpr(argument, 
boundary_shape_vars);
+      bool inline_bound_value = !bound_value.same_as(argument) &&
+                                IsShapeDependentPrimExpr(bound_value, 
boundary_shape_vars);
+      if (!inline_argument && !inline_bound_value) continue;
+
+      if (inline_bound_value) {
+        const auto* argument_var = argument.as<VarNode>();
+        TVM_FFI_ICHECK(argument_var);
+        inlined_bindings_[argument_var] = bound_value;
+      }
+
+      size_t param_idx = it - arguments_.begin();
+      arguments_.erase(arguments_.begin() + param_idx);
+      params_.erase(params_.begin() + param_idx);
+    }
+  }
+
   /*!
    * \brief Check whether the input expression is defined within this 
function. If not, create a new
    * parameter for the expression.
@@ -624,14 +675,7 @@ class FunctionCreator : public ExprMutator {
     const auto* var = expr.as<VarNode>();
     if (var != nullptr && defined_vars_.count(var) == 0) {
       Var bound_var = ffi::GetRef<Var>(var);
-      Expr bound_value = bound_var;
-      std::unordered_set<const VarNode*> visited;
-      while (const auto* current_var = bound_value.as<VarNode>()) {
-        if (!visited.insert(current_var).second) break;
-        auto it = outer_bindings_.find(ffi::GetRef<Var>(current_var));
-        if (it == outer_bindings_.end()) break;
-        bound_value = (*it).second;
-      }
+      Expr bound_value = ResolveOuterBinding(bound_var);
       if (!bound_value.same_as(bound_var) && 
IsInlinableConstants(bound_value)) {
         inlined_bindings_[var] = bound_value;
         return;
@@ -646,6 +690,9 @@ class FunctionCreator : public ExprMutator {
         Var param(std::move(name), GetType(expr));
         arguments_.push_back(expr);
         params_.push_back(param);
+        if (IsSymbolicPrimExpr(expr)) {
+          deferred_prim_args_.push_back(expr);
+        }
       }
 
       // Mark the tuple parameter is partially referenced in the beginning.
@@ -675,6 +722,22 @@ class FunctionCreator : public ExprMutator {
 
   // Check if the expression is constant PrimExpr or ShapeExpr or tuple of 
them that can be
   // inlined in the composite functions and excluded from args/params.
+  bool IsSymbolicPrimExpr(const Expr& expr) {
+    if (expr.as<CallNode>()) return false;
+    if (auto prim_value = expr.as<PrimExpr>()) {
+      return !tvm::tirx::UndefinedVars(prim_value.value()).empty();
+    }
+    return false;
+  }
+
+  bool IsShapeDependentPrimExpr(const Expr& expr,
+                                const std::unordered_set<tirx::Var>& 
referenced_shape_vars) {
+    if (!IsSymbolicPrimExpr(expr)) return false;
+    ffi::Array<tirx::Var> undefined_vars = 
tvm::tirx::UndefinedVars(expr.as_or_throw<PrimExpr>());
+    return std::all_of(undefined_vars.begin(), undefined_vars.end(),
+                       [&](const tirx::Var& var) { return 
referenced_shape_vars.count(var); });
+  }
+
   bool IsInlinableConstants(const Expr& expr) {
     if (const auto* tuple = expr.as<TupleNode>()) {
       return std::all_of(tuple->fields.begin(), tuple->fields.end(),
@@ -695,6 +758,8 @@ class FunctionCreator : public ExprMutator {
   std::unordered_set<const VarNode*> defined_vars_;
   /*! \brief Caller variables replaced by statically inlinable bound values. */
   std::unordered_map<const VarNode*, Expr> inlined_bindings_;
+  /*! \brief Symbolic primitive arguments classified after the complete 
boundary is known. */
+  ffi::Array<Expr> deferred_prim_args_;
   /*! \brief The number of parameters reserved for constants */
   int n_param_for_const_ = 0;
   /*! \brief The output vars */
diff --git a/tests/python/relax/test_transform_fuse_ops.py 
b/tests/python/relax/test_transform_fuse_ops.py
index ed3b02d911..bee21e0686 100644
--- a/tests/python/relax/test_transform_fuse_ops.py
+++ b/tests/python/relax/test_transform_fuse_ops.py
@@ -1352,6 +1352,478 @@ def test_symbolic_shape_aware_fuse_2():
     _check(Before, Expected)
 
 
+def test_symbolic_prim_arg_after_tensor_arg():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def add_one(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n), "float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n), "float32")
+            for i in range(n):
+                with T.sblock("add_one"):
+                    vi = T.axis.spatial(n, i)
+                    out[0, vi] = x[0, vi] + T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def exp(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n), "float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n), "float32")
+            for i in range(n):
+                with T.sblock("exp"):
+                    vi = T.axis.spatial(n, i)
+                    out[0, vi] = T.exp(x[0, vi])
+
+        @R.function
+        def main(
+            x: R.Tensor((1, "n"), dtype="float32"),
+        ) -> R.Tensor((1, "n"), dtype="float32"):
+            n = T.int64()
+            cls = Before
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.add_one,
+                    (x, n),
+                    out_ty=R.Tensor((1, n), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.exp,
+                    (lv, n),
+                    out_ty=R.Tensor((1, n), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+    assert relax.analysis.check_well_formed(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 1
+    assert not isinstance(fused.params[0].ty, tvm.ir.PrimType)
+    assert fused.ret_ty.shape is not None
+
+    mod = relax.transform.FuseTIR()(mod)
+    assert relax.analysis.check_well_formed(mod)
+    fused_tir = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert tvm.tirx.analysis.verify_well_formed(fused_tir)
+
+
+def test_symbolic_prim_arg_before_tensor_arg():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def add_one(n: T.int64, x_handle: T.handle, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n), "float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n), "float32")
+            for i in range(n):
+                with T.sblock("add_one"):
+                    vi = T.axis.spatial(n, i)
+                    out[0, vi] = x[0, vi] + T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def exp(n: T.int64, x_handle: T.handle, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n), "float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n), "float32")
+            for i in range(n):
+                with T.sblock("exp"):
+                    vi = T.axis.spatial(n, i)
+                    out[0, vi] = T.exp(x[0, vi])
+
+        @R.function
+        def main(
+            x: R.Tensor((1, "n"), dtype="float32"),
+        ) -> R.Tensor((1, "n"), dtype="float32"):
+            n = T.int64()
+            cls = Before
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.add_one,
+                    (n, x),
+                    out_ty=R.Tensor((1, n), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.exp,
+                    (n, lv),
+                    out_ty=R.Tensor((1, n), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+    assert relax.analysis.check_well_formed(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 1
+    assert not isinstance(fused.params[0].ty, tvm.ir.PrimType)
+    assert fused.ret_ty.shape is not None
+
+    mod = relax.transform.FuseTIR()(mod)
+    assert relax.analysis.check_well_formed(mod)
+    fused_tir = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert tvm.tirx.analysis.verify_well_formed(fused_tir)
+
+
+def test_symbolic_prim_arg_reused_from_derived_tensor_shape():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def add_one(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(
+                x_handle,
+                (T.int64(1), (n - T.int64(1)) // T.int64(4) + T.int64(1)),
+                "float32",
+            )
+            out = T.match_buffer(
+                out_handle,
+                (T.int64(1), (n - T.int64(1)) // T.int64(4) + T.int64(1)),
+                "float32",
+            )
+            for i in range((n - T.int64(1)) // T.int64(4) + T.int64(1)):
+                with T.sblock("add_one"):
+                    vi = T.axis.spatial((n - T.int64(1)) // T.int64(4) + 
T.int64(1), i)
+                    out[0, vi] = x[0, vi] + T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def exp(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(
+                x_handle,
+                (T.int64(1), (n - T.int64(1)) // T.int64(4) + T.int64(1)),
+                "float32",
+            )
+            out = T.match_buffer(
+                out_handle,
+                (T.int64(1), (n - T.int64(1)) // T.int64(4) + T.int64(1)),
+                "float32",
+            )
+            for i in range((n - T.int64(1)) // T.int64(4) + T.int64(1)):
+                with T.sblock("exp"):
+                    vi = T.axis.spatial((n - T.int64(1)) // T.int64(4) + 
T.int64(1), i)
+                    out[0, vi] = T.exp(x[0, vi])
+
+        @R.function
+        def main(
+            source: R.Tensor(("n",), dtype="float32"),
+            x: R.Tensor((1, "(n - 1) // 4 + 1"), dtype="float32"),
+        ) -> R.Tensor((1, "(n - 1) // 4 + 1"), dtype="float32"):
+            n = T.int64()
+            cls = Before
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.add_one,
+                    (x, n),
+                    out_ty=R.Tensor((1, (n - 1) // 4 + 1), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.exp,
+                    (lv, n),
+                    out_ty=R.Tensor((1, (n - 1) // 4 + 1), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 2
+    assert isinstance(fused.params[1].ty, relax.ShapeType)
+    assert all(not isinstance(param.ty, tvm.ir.PrimType) for param in 
fused.params)
+
+    mod = relax.transform.FuseTIR()(mod)
+    fused_tir = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert tvm.tirx.analysis.verify_well_formed(fused_tir)
+
+
+def test_symbolic_prim_arg_not_bound_by_derived_tensor_shape():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def add_one(x_handle: T.handle, n: T.int64, m: T.int64, out_handle: 
T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n + T.int64(1)), 
"float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n + T.int64(1)), 
"float32")
+            for i in range(n + T.int64(1)):
+                with T.sblock("add_one"):
+                    vi = T.axis.spatial(n + T.int64(1), i)
+                    out[0, vi] = x[0, vi] + T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def exp(x_handle: T.handle, n: T.int64, m: T.int64, out_handle: 
T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(1), n + T.int64(1)), 
"float32")
+            out = T.match_buffer(out_handle, (T.int64(1), n + T.int64(1)), 
"float32")
+            for i in range(n + T.int64(1)):
+                with T.sblock("exp"):
+                    vi = T.axis.spatial(n + T.int64(1), i)
+                    out[0, vi] = T.exp(x[0, vi])
+
+        @R.function
+        def main(
+            shape: R.Shape(["n", "m"]),
+            x: R.Tensor((1, "n + 1"), dtype="float32"),
+        ) -> R.Tensor((1, "n + 1"), dtype="float32"):
+            n = T.int64()
+            m = T.int64()
+            cls = Before
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.add_one,
+                    (x, n, m),
+                    out_ty=R.Tensor((1, n + 1), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.exp,
+                    (lv, n, m),
+                    out_ty=R.Tensor((1, n + 1), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 3
+    assert sum(isinstance(param.ty, relax.ShapeType) for param in 
fused.params) == 1
+    assert sum(isinstance(param.ty, tvm.ir.PrimType) for param in 
fused.params) == 1
+
+
+def test_primitive_call_arg_not_inlined():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def add_scalar(x_handle: T.handle, value: T.int64, out_handle: 
T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(4),), "int64")
+            out = T.match_buffer(out_handle, (T.int64(4),), "int64")
+            for i in range(4):
+                with T.sblock("add_scalar"):
+                    vi = T.axis.spatial(4, i)
+                    out[vi] = x[vi] + value
+
+        @T.prim_func(private=True, s_tir=True)
+        def double(x_handle: T.handle, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (T.int64(4),), "int64")
+            out = T.match_buffer(out_handle, (T.int64(4),), "int64")
+            for i in range(4):
+                with T.sblock("double"):
+                    vi = T.axis.spatial(4, i)
+                    out[vi] = x[vi] * T.int64(2)
+
+        @R.function
+        def main(x: R.Tensor((4,), dtype="int64")):
+            cls = Before
+            with R.dataflow():
+                value: R.Prim("int64") = R.call_pure_packed("get_scalar", 
ty_args=R.Prim("int64"))
+                lv = R.call_tir(
+                    cls.add_scalar,
+                    (x, value),
+                    out_ty=R.Tensor((4,), dtype="int64"),
+                )
+                gv = R.call_tir(
+                    cls.double,
+                    (lv,),
+                    out_ty=R.Tensor((4,), dtype="int64"),
+                )
+                R.output(gv, value)
+            return gv, value
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+    assert relax.analysis.check_well_formed(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 2
+    assert sum(isinstance(param.ty, tvm.ir.PrimType) for param in 
fused.params) == 1
+
+    mod = relax.transform.FuseTIR()(mod)
+    assert relax.analysis.check_well_formed(mod)
+    fused_tir = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert tvm.tirx.analysis.verify_well_formed(fused_tir)
+
+
+def test_primitive_call_arg_used_by_output_shape_not_inlined():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def make(n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            out = T.match_buffer(out_handle, (n,), "float32")
+            for i in range(n):
+                with T.sblock("make"):
+                    vi = T.axis.spatial(n, i)
+                    out[vi] = T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def double(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (n,), "float32")
+            out = T.match_buffer(out_handle, (n,), "float32")
+            for i in range(n):
+                with T.sblock("double"):
+                    vi = T.axis.spatial(n, i)
+                    out[vi] = x[vi] * T.float32(2)
+
+        @R.function(pure=False)
+        def main():
+            cls = Before
+            n: R.Prim("int64") = R.call_packed("get_extent", 
ty_args=R.Prim("int64"))
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.make,
+                    (n,),
+                    out_ty=R.Tensor((n,), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.double,
+                    (lv, n),
+                    out_ty=R.Tensor((n,), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+    assert relax.analysis.check_well_formed(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 1
+    assert isinstance(fused.params[0].ty, tvm.ir.PrimType)
+
+    packed_calls = []
+
+    def collect_packed_calls(expr):
+        if (
+            isinstance(expr, relax.Call)
+            and isinstance(expr.op, relax.ExternFunc)
+            and expr.op.global_symbol == "get_extent"
+        ):
+            packed_calls.append(expr)
+
+    relax.analysis.post_order_visit(mod["main"], collect_packed_calls)
+    assert len(packed_calls) == 1
+
+    packed_calls.clear()
+    relax.analysis.post_order_visit(
+        fused,
+        collect_packed_calls,
+    )
+    assert not packed_calls
+
+
+def test_symbolic_prim_arg_used_only_by_output_shape():
+    @I.ir_module(s_tir=True)
+    class Before:
+        @T.prim_func(private=True, s_tir=True)
+        def make(n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            out = T.match_buffer(out_handle, (n,), "float32")
+            for i in range(n):
+                with T.sblock("make"):
+                    vi = T.axis.spatial(n, i)
+                    out[vi] = T.float32(1)
+
+        @T.prim_func(private=True, s_tir=True)
+        def double(x_handle: T.handle, n: T.int64, out_handle: T.handle):
+            T.func_attr({"op_pattern": 0, "tirx.noalias": True})
+            x = T.match_buffer(x_handle, (n,), "float32")
+            out = T.match_buffer(out_handle, (n,), "float32")
+            for i in range(n):
+                with T.sblock("double"):
+                    vi = T.axis.spatial(n, i)
+                    out[vi] = x[vi] * T.float32(2)
+
+        @R.function
+        def main(
+            source: R.Tensor(("n",), dtype="float32"),
+        ) -> R.Tensor(("n",), dtype="float32"):
+            n = T.int64()
+            cls = Before
+            with R.dataflow():
+                lv = R.call_tir(
+                    cls.make,
+                    (n,),
+                    out_ty=R.Tensor((n,), dtype="float32"),
+                )
+                gv = R.call_tir(
+                    cls.double,
+                    (lv, n),
+                    out_ty=R.Tensor((n,), dtype="float32"),
+                )
+                R.output(gv)
+            return gv
+
+    mod = relax.transform.AnnotateTIROpPattern()(Before)
+    mod = relax.transform.FuseOps()(mod)
+    assert relax.analysis.check_well_formed(mod)
+
+    fused = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert len(fused.params) == 1
+    assert isinstance(fused.params[0].ty, relax.ShapeType)
+    assert fused.ret_ty.shape is not None
+
+    mod = relax.transform.FuseTIR()(mod)
+    assert relax.analysis.check_well_formed(mod)
+    fused_tir = next(
+        mod[global_var]
+        for global_var in mod.get_global_vars()
+        if global_var.name_hint.startswith("fused_")
+    )
+    assert tvm.tirx.analysis.verify_well_formed(fused_tir)
+
+
 def test_shape_expr_arg():
     @I.ir_module(s_tir=True)
     class Before:

Reply via email to