Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 0 additions & 18 deletions include/tvm/tirx/stmt_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
28 changes: 0 additions & 28 deletions python/tvm/tirx/stmt_functor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
57 changes: 1 addition & 56 deletions src/tirx/ir/stmt_functor.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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) {}
Expand Down Expand Up @@ -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) {}
Expand Down Expand Up @@ -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); });
Expand Down
63 changes: 0 additions & 63 deletions tests/python/tirx-base/test_tir_stmt_functor_ir_transform.py

This file was deleted.

Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,14 @@
# specific language governing permissions and limitations
# under the License.
import pytest
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")

Expand All @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,14 @@
# specific language governing permissions and limitations
# under the License.
import pytest
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")

Expand All @@ -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):
Expand Down
10 changes: 3 additions & 7 deletions tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand All @@ -32,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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,14 @@
# specific language governing permissions and limitations
# under the License.
import pytest
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")

Expand All @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,14 @@
# specific language governing permissions and limitations
# under the License.
import pytest
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")

Expand All @@ -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):
Expand Down
10 changes: 3 additions & 7 deletions tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand All @@ -32,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):
Expand Down
Loading
Loading