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):