From ea0e6b7bf6ed962f52feed280eaff1d3bbe7c7ca Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 00:31:16 +0000 Subject: [PATCH 1/2] Remove legacy TIRx IRTransform --- 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 ------------------- .../tile_primitive/trn/test_binary_trn.py | 9 +-- .../tile_primitive/trn/test_compose_op_trn.py | 9 +-- .../tile_primitive/trn/test_copy_trn.py | 9 +-- .../tile_primitive/trn/test_gemm_trn.py | 9 +-- .../tile_primitive/trn/test_reduction_trn.py | 9 +-- .../tile_primitive/trn/test_select_trn.py | 9 +-- .../tile_primitive/trn/test_unary_trn.py | 9 +-- 11 files changed, 15 insertions(+), 214 deletions(-) delete mode 100644 tests/python/tirx-base/test_tir_stmt_functor_ir_transform.py diff --git a/include/tvm/tirx/stmt_functor.h b/include/tvm/tirx/stmt_functor.h index 78ebec843408..54cf329a0a43 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> 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 83405d7eea7f..eec1c184f9fc 100644 --- a/python/tvm/tirx/stmt_functor.py +++ b/python/tvm/tirx/stmt_functor.py @@ -980,34 +980,6 @@ def visit_expr(self, expr): 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 ee0161d42d67..251f9d217ea7 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 f) : f_(f) {} @@ -780,60 +780,6 @@ void PostOrderVisit(const ffi::ObjectRef& node, std::function& only_enable) - : f_preorder_(f_preorder), f_postorder_(f_postorder), only_enable_(only_enable) {} - - Stmt VisitStmt(const Stmt& stmt) final { - return MutateInternal(stmt, [this](const Stmt& s) { return this->BaseVisitStmt(s); }); - } - Expr VisitExpr(const Expr& expr) final { - return MutateInternal(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 - 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(); - if (pre.defined()) return pre; - } - T new_node = fmutate(node); - if (f_postorder_ != nullptr) { - T post = f_postorder_(new_node).template cast(); - 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& only_enable_; -}; - -Stmt IRTransform(Stmt ir_node, const ffi::Function& f_preorder, const ffi::Function& f_postorder, - ffi::Optional> only_enable) { - std::unordered_set 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(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 0c9b667aeabd..000000000000 --- 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 268ef0eae6f3..f758e59d0289 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 @@ -18,11 +18,11 @@ import tvm import tvm.testing +import tvm_ffi 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 _postorder(node): 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 b5e8a6554a64..55ea6658bbb4 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 @@ -18,11 +18,11 @@ import tvm import tvm.testing +import tvm_ffi 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 _postorder(node): 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 308048a081a0..accb3200307d 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 @@ -17,11 +17,11 @@ import tvm import tvm.testing +import tvm_ffi 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 +32,7 @@ def _postorder(node): 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 8397627a88c4..e68c7b68e8af 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 @@ -18,11 +18,11 @@ import tvm import tvm.testing +import tvm_ffi 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 _postorder(node): 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 36da370d10e8..acf53cbdf7ac 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 @@ -18,11 +18,11 @@ import tvm import tvm.testing +import tvm_ffi 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 _postorder(node): 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 477620eb7a9d..7f14f615849e 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 @@ -17,11 +17,11 @@ import tvm import tvm.testing +import tvm_ffi 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 +32,7 @@ def _postorder(node): 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 a6557c346a1a..2d0ca2e1dc47 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 @@ -18,11 +18,11 @@ import tvm import tvm.testing +import tvm_ffi 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 _postorder(node): 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): From 7a2c6ea2f3cc234bf8890c334aafb70b9de954c8 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 01:53:44 +0000 Subject: [PATCH 2/2] Order tvm_ffi imports for lint --- .../python/tirx/operator/tile_primitive/trn/test_binary_trn.py | 2 +- .../tirx/operator/tile_primitive/trn/test_compose_op_trn.py | 2 +- tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py | 3 ++- tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py | 2 +- .../tirx/operator/tile_primitive/trn/test_reduction_trn.py | 2 +- .../python/tirx/operator/tile_primitive/trn/test_select_trn.py | 3 ++- .../python/tirx/operator/tile_primitive/trn/test_unary_trn.py | 2 +- 7 files changed, 9 insertions(+), 7 deletions(-) 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 f758e59d0289..e46c995a3801 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,10 +15,10 @@ # specific language governing permissions and limitations # under the License. import pytest +import tvm_ffi import tvm import tvm.testing -import tvm_ffi 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 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 55ea6658bbb4..3993663b5886 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,10 +15,10 @@ # specific language governing permissions and limitations # under the License. import pytest +import tvm_ffi import tvm import tvm.testing -import tvm_ffi 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 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 accb3200307d..206cc3be5195 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,9 +15,10 @@ # specific language governing permissions and limitations # under the License. +import tvm_ffi + import tvm import tvm.testing -import tvm_ffi 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 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 e68c7b68e8af..8515eaa2e046 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,10 +15,10 @@ # specific language governing permissions and limitations # under the License. import pytest +import tvm_ffi import tvm import tvm.testing -import tvm_ffi 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 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 acf53cbdf7ac..adcf465cf39f 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,10 +15,10 @@ # specific language governing permissions and limitations # under the License. import pytest +import tvm_ffi import tvm import tvm.testing -import tvm_ffi 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 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 7f14f615849e..9e4580e612b1 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,9 +15,10 @@ # specific language governing permissions and limitations # under the License. +import tvm_ffi + import tvm import tvm.testing -import tvm_ffi 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 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 2d0ca2e1dc47..5b476b99b3df 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,10 +15,10 @@ # specific language governing permissions and limitations # under the License. import pytest +import tvm_ffi import tvm import tvm.testing -import tvm_ffi 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