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

tqchen 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 fa8691eb2b [REFACTOR][TIR] Remove IRTransform in favour of 
tvm_ffi.structural_map (#20304)
fa8691eb2b is described below

commit fa8691eb2b41bcabca724174c2ac7dfabfaa4124
Author: Tianqi Chen <[email protected]>
AuthorDate: Thu Sep 10 08:07:44 2026 -0400

    [REFACTOR][TIR] Remove IRTransform in favour of tvm_ffi.structural_map 
(#20304)
    
    Replace the legacy TIRx IRTransform API with the typed
    tvm_ffi.structural_map interface while preserving the existing postorder
    replacement behavior.
    
    - Remove the C++ declaration, implementation, registration, and Python
    wrapper for IRTransform.
    - Migrate all remaining postorder-only tile-primitive tests to typed
    structural_map callbacks.
    - Remove the dedicated mixed preorder/postorder legacy API test; no
    production caller uses both phases in one pass.
    - Validate the changed area with 128 passing tests and compare the
    full-suite failure set against the parent tree.
---
 include/tvm/tirx/stmt_functor.h                    | 18 -------
 python/tvm/tirx/stmt_functor.py                    | 28 ----------
 src/tirx/ir/stmt_functor.cc                        | 57 +-------------------
 .../test_tir_stmt_functor_ir_transform.py          | 63 ----------------------
 .../operator/tile_primitive/trn/test_binary_trn.py |  9 +---
 .../tile_primitive/trn/test_compose_op_trn.py      |  9 +---
 .../operator/tile_primitive/trn/test_copy_trn.py   | 10 ++--
 .../operator/tile_primitive/trn/test_gemm_trn.py   |  9 +---
 .../tile_primitive/trn/test_reduction_trn.py       |  9 +---
 .../operator/tile_primitive/trn/test_select_trn.py | 10 ++--
 .../operator/tile_primitive/trn/test_unary_trn.py  |  9 +---
 11 files changed, 17 insertions(+), 214 deletions(-)

diff --git a/include/tvm/tirx/stmt_functor.h b/include/tvm/tirx/stmt_functor.h
index 78ebec8434..54cf329a0a 100644
--- a/include/tvm/tirx/stmt_functor.h
+++ b/include/tvm/tirx/stmt_functor.h
@@ -366,24 +366,6 @@ class TVM_DLL StmtExprMutator : public ExprMutator, public 
StmtMutator {
   Expr VisitExpr_(const BufferRegionNode* op) override;
 };
 
-/*!
- * \brief recursively visit the ir nodes in post DFS order, and transform it
- *
- * \param stmt The ir to be transformed.
- * \param preorder The function called in before recursive mutation
- *          If preorder returns None, then the transform will proceed to 
recursive call.
- *          If preorder returns a not None Stmt/Expr, the transformer will 
simply return it and
- *          won't do further recursion.
- * \param postorder The function called after recursive mutation.
- *          The recursive mutation result is passed to postorder for further 
mutation.
- * \param only_enable List of String.
- *          If it is null, all IRNode will call preorder/postorder
- *          If it is not null, preorder/postorder will only be called
- *          when the IRNode's type key is in the list.
- */
-TVM_DLL Stmt IRTransform(Stmt stmt, const ffi::Function& preorder, const 
ffi::Function& postorder,
-                         ffi::Optional<ffi::Array<ffi::String>> only_enable = 
std::nullopt);
-
 /*!
  * \brief Recursively visit a statement or expression in post DFS order, 
applying fvisit.
  * Each node is guaranteed to be visited only once.
diff --git a/python/tvm/tirx/stmt_functor.py b/python/tvm/tirx/stmt_functor.py
index 83405d7eea..eec1c184f9 100644
--- a/python/tvm/tirx/stmt_functor.py
+++ b/python/tvm/tirx/stmt_functor.py
@@ -980,34 +980,6 @@ class StmtExprMutator(StmtMutator, ExprMutator):
         return ExprMutator.visit_expr(self, expr)
 
 
-def ir_transform(stmt, preorder, postorder, only_enable=None):
-    """Recursively visit and transform ir nodes in post DFS order.
-
-    Parameters
-    ----------
-    stmt : tvm.tirx.Stmt
-        The input to be transformed.
-
-    preorder: function
-        The function called in before recursive mutation
-        If preorder returns None, then the transform will proceed to recursive 
call.
-        If preorder returns a not None tvm.tirx.Stmt/Expr, the transformer 
will simply return it and
-        won't do further recursion.
-
-    postorder : function
-        The function called after recursive mutation.
-
-    only_enable : Optional[List[str]]
-        List of types that we only enable.
-
-    Returns
-    -------
-    result : tvm.tirx.Stmt
-        The result.
-    """
-    return _ffi_api.IRTransform(stmt, preorder, postorder, only_enable)  # 
type: ignore
-
-
 def post_order_visit(node, fvisit):
     """Recursively visit a statement or expression in post DFS order, applying 
fvisit.
        Each node is guaranteed to be visited only once.
diff --git a/src/tirx/ir/stmt_functor.cc b/src/tirx/ir/stmt_functor.cc
index ee0161d42d..251f9d217e 100644
--- a/src/tirx/ir/stmt_functor.cc
+++ b/src/tirx/ir/stmt_functor.cc
@@ -743,7 +743,7 @@ Stmt StmtMutator::VisitStmt_(const 
tirx::TilePrimitiveCallNode* op) {
   }
 }
 
-// Implementations of IRTransform, PostOrderVisit and Substitute
+// Implementations of PostOrderVisit and Substitute
 class IRApplyVisit : public StmtExprVisitor {
  public:
   explicit IRApplyVisit(std::function<void(const ffi::ObjectRef&)> f) : f_(f) 
{}
@@ -780,60 +780,6 @@ void PostOrderVisit(const ffi::ObjectRef& node, 
std::function<void(const ffi::Ob
   }
 }
 
-class IRTransformer final : public StmtExprMutator {
- public:
-  IRTransformer(const ffi::Function& f_preorder, const ffi::Function& 
f_postorder,
-                const std::unordered_set<uint32_t>& only_enable)
-      : f_preorder_(f_preorder), f_postorder_(f_postorder), 
only_enable_(only_enable) {}
-
-  Stmt VisitStmt(const Stmt& stmt) final {
-    return MutateInternal<Stmt>(stmt, [this](const Stmt& s) { return 
this->BaseVisitStmt(s); });
-  }
-  Expr VisitExpr(const Expr& expr) final {
-    return MutateInternal<Expr>(expr, [this](const Expr& e) { return 
this->BaseVisitExpr(e); });
-  }
-
- private:
-  // NOTE: redirect to parent's call
-  // This is used to get around limitation of gcc-4.8
-  Stmt BaseVisitStmt(const Stmt& s) { return StmtMutator::VisitStmt(s); }
-  Expr BaseVisitExpr(const Expr& e) { return ExprMutator::VisitExpr(e); }
-
-  template <typename T, typename F>
-  T MutateInternal(const T& node, F fmutate) {
-    if (only_enable_.size() && !only_enable_.count(node->type_index())) {
-      return fmutate(node);
-    }
-    if (f_preorder_ != nullptr) {
-      T pre = f_preorder_(node).template cast<T>();
-      if (pre.defined()) return pre;
-    }
-    T new_node = fmutate(node);
-    if (f_postorder_ != nullptr) {
-      T post = f_postorder_(new_node).template cast<T>();
-      if (post.defined()) return post;
-    }
-    return new_node;
-  }
-  // The functions
-  const ffi::Function& f_preorder_;
-  const ffi::Function& f_postorder_;
-  // type indices enabled.
-  const std::unordered_set<uint32_t>& only_enable_;
-};
-
-Stmt IRTransform(Stmt ir_node, const ffi::Function& f_preorder, const 
ffi::Function& f_postorder,
-                 ffi::Optional<ffi::Array<ffi::String>> only_enable) {
-  std::unordered_set<uint32_t> only_type_index;
-  if (only_enable.has_value()) {
-    for (auto s : only_enable.value()) {
-      only_type_index.insert(ffi::TypeKeyToIndex(s.c_str()));
-    }
-  }
-  IRTransformer transform(f_preorder, f_postorder, only_type_index);
-  return transform(std::move(ir_node));
-}
-
 class IRSubstitute : public StmtExprMutator {
  public:
   explicit IRSubstitute(std::function<ffi::Optional<Expr>(const Var&)> vmap) : 
vmap_(vmap) {}
@@ -1006,7 +952,6 @@ PrimExpr SubstituteWithDataTypeLegalization(
 TVM_FFI_STATIC_INIT_BLOCK() {
   namespace refl = tvm::ffi::reflection;
   refl::GlobalDef()
-      .def("tirx.IRTransform", IRTransform)
       .def("tirx.PostOrderVisit",
            [](ffi::ObjectRef node, ffi::Function f) {
              tirx::PostOrderVisit(node, [f](const ffi::ObjectRef& n) { f(n); 
});
diff --git a/tests/python/tirx-base/test_tir_stmt_functor_ir_transform.py 
b/tests/python/tirx-base/test_tir_stmt_functor_ir_transform.py
deleted file mode 100644
index 0c9b667aea..0000000000
--- a/tests/python/tirx-base/test_tir_stmt_functor_ir_transform.py
+++ /dev/null
@@ -1,63 +0,0 @@
-# Licensed to the Apache Software Foundation (ASF) under one
-# or more contributor license agreements.  See the NOTICE file
-# distributed with this work for additional information
-# regarding copyright ownership.  The ASF licenses this file
-# to you under the Apache License, Version 2.0 (the
-# "License"); you may not use this file except in compliance
-# with the License.  You may obtain a copy of the License at
-#
-#   http://www.apache.org/licenses/LICENSE-2.0
-#
-# Unless required by applicable law or agreed to in writing,
-# software distributed under the License is distributed on an
-# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
-# KIND, either express or implied.  See the License for the
-# specific language governing permissions and limitations
-# under the License.
-import tvm
-from tvm.script import ir as I
-from tvm.script import tirx as T
-
-
-def test_ir_transform():
-    @I.ir_module
-    class Module:
-        @T.prim_func(s_tir=True)
-        def main(n: T.int32):
-            for i in T.serial(n):
-                for j in T.serial(10):
-                    # Inline call_extern to avoid Let binding (x must be the 
Call node itself)
-                    T.evaluate(
-                        T.call_extern(
-                            "int32", "TestB", T.call_extern("int32", "TestA", 
i * 3 + j * 1)
-                        )
-                    )
-                    T.evaluate(
-                        T.call_extern(
-                            "int32", "TestC", T.call_extern("int32", "TestA", 
i * 3 + j * 1)
-                        )
-                    )
-
-    body = Module["main"].body
-    builtin_call_extern = tvm.ir.Op.get("tirx.call_extern")
-
-    def preorder(op):
-        if op.op.same_as(builtin_call_extern) and op.args[0].value == "TestC":
-            return tvm.tirx.const(42, "int32")
-        return None
-
-    def postorder(op):
-        assert isinstance(op, tvm.ir.Call)
-        assert tvm.ir.is_prim_expr(op)
-        if op.op.same_as(builtin_call_extern) and op.args[0].value == "TestA":
-            return tvm.tirx.call_extern("int32", "TestB", op.args[1] + 1)
-        return op
-
-    body = tvm.tirx.stmt_functor.ir_transform(body, preorder, postorder, 
["ir.Call"])
-    stmt_list = tvm.tirx.stmt_list(body.body.body)
-    assert stmt_list[0].value.args[1].args[0].value == "TestB"
-    assert stmt_list[1].value.value == 42
-
-
-if __name__ == "__main__":
-    test_ir_transform()
diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py
index 268ef0eae6..e46c995a38 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -22,7 +23,6 @@ from tvm.ir import assert_structural_equal as 
_assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -33,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git 
a/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py
index b5e8a6554a..3993663b58 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -22,7 +23,6 @@ from tvm.ir import assert_structural_equal as 
_assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -33,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py
index 308048a081..206cc3be51 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py
@@ -15,13 +15,14 @@
 # specific language governing permissions and limitations
 # under the License.
 
+import tvm_ffi
+
 import tvm
 import tvm.testing
 from tvm.ir import assert_structural_equal as _assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -32,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py
index 8397627a88..8515eaa2e0 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -22,7 +23,6 @@ from tvm.ir import assert_structural_equal as 
_assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -33,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git 
a/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py
index 36da370d10..adcf465cf3 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -22,7 +23,6 @@ from tvm.ir import assert_structural_equal as 
_assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -33,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py
index 477620eb7a..9e4580e612 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py
@@ -15,13 +15,14 @@
 # specific language governing permissions and limitations
 # under the License.
 
+import tvm_ffi
+
 import tvm
 import tvm.testing
 from tvm.ir import assert_structural_equal as _assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -32,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):
diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py 
b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py
index a6557c346a..5b476b99b3 100644
--- a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py
+++ b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -22,7 +23,6 @@ from tvm.ir import assert_structural_equal as 
_assert_structural_equal
 from tvm.script import tirx as T
 from tvm.script.tirx import tile as Tx
 from tvm.tirx.layout import F, P, S, TileLayout
-from tvm.tirx.stmt_functor import ir_transform
 
 target = tvm.target.Target("aws/trn1/trn1.2xlarge")
 
@@ -33,12 +33,7 @@ def _strip_exec_scope_stmt(stmt):
             return node.body
         return node
 
-    return ir_transform(
-        stmt,
-        preorder=lambda _node: None,
-        postorder=_postorder,
-        only_enable=["tirx.AttrStmt"],
-    )
+    return tvm_ffi.structural_map(stmt, (tvm.tirx.AttrStmt, _postorder))
 
 
 def assert_structural_equal(lhs, rhs, *args, **kwargs):

Reply via email to