From 290830af5cc2bb946bb9eb8edbc9a2fde3da6aff Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Tue, 6 Oct 2026 04:01:37 +0000 Subject: [PATCH 1/4] [TIRx] Represent device boundaries with region ops Use distinct device entry and internal split regions to preserve lexical scope and target ownership. Route script construction, scope resolution, target binding and host/device extraction through these operators. --- docs/tirx/native_basics/cuda/functions.rst | 15 ++--- include/tvm/tirx/builtin.h | 3 + include/tvm/tirx/script/ir_builder/ir.h | 12 +--- include/tvm/tirx/stmt.h | 12 +--- .../trn/transform/private_buffer_alloc.py | 15 ++--- .../tirx/script/ir_builder/parser_protocol.py | 25 ++------- src/s_tir/transform/decorate_device_scope.cc | 4 +- src/tirx/analysis/verify_tirx_well_formed.cc | 7 ++- src/tirx/ir/stmt.cc | 12 ++++ src/tirx/op/builtin.cc | 5 ++ src/tirx/script/ir_builder/ir.cc | 27 +-------- src/tirx/script/printer/stmt.cc | 4 +- src/tirx/transform/bind_target.cc | 38 +++---------- src/tirx/transform/split_host_device.cc | 55 ++++++++++++------- src/tirx/transform/tile_primitive_dispatch.cc | 31 +++++++---- ...t_s_tir_transform_decorate_device_scope.py | 2 +- .../test_tir_transform_split_host_device.py | 34 ++++++------ .../tile_primitive/trn/test_binary_trn.py | 6 +- .../tile_primitive/trn/test_compose_op_trn.py | 6 +- .../tile_primitive/trn/test_copy_trn.py | 6 +- .../tile_primitive/trn/test_gemm_trn.py | 6 +- .../tile_primitive/trn/test_reduction_trn.py | 6 +- .../tile_primitive/trn/test_select_trn.py | 6 +- .../tile_primitive/trn/test_unary_trn.py | 6 +- 24 files changed, 154 insertions(+), 189 deletions(-) diff --git a/docs/tirx/native_basics/cuda/functions.rst b/docs/tirx/native_basics/cuda/functions.rst index 661d5877c4f9..c09ec847a37b 100644 --- a/docs/tirx/native_basics/cuda/functions.rst +++ b/docs/tirx/native_basics/cuda/functions.rst @@ -220,13 +220,14 @@ Launch parameters ``Tx.device_entry()`` ~~~~~~~~~~~~~~~~~~~~~ -``Tx.device_entry()`` is a flat marker (no ``with``) that starts the authored -device region: parameter binding and shape reads precede it, while the kernel -body follows it. The parser represents the marker as -``AttrStmt("tirx.device_entry", ...)``. ``LowerTIRx`` then removes the marker, -resolves scope ids, and wraps the device body in thread-extent attributes; -target binding and ``SplitHostDevice`` use those resulting device regions to -extract the kernel shown above. +``Tx.device_entry()`` starts the authored device region: parameter binding and +shape reads precede it, while the kernel body follows it. A flat call scopes the +remaining statements in the enclosing body; ``with Tx.device_entry():`` gives +an explicit boundary. Both forms create a ``RegionStmt`` with the +``tirx.device_entry`` op. ``LowerTIRx`` removes this region, resolves scope ids, +and wraps the device body in single-axis ``launch_thread`` regions. Target +binding and ``SplitHostDevice`` use those resulting device regions to extract +the kernel shown above. Scope ids ~~~~~~~~~ diff --git a/include/tvm/tirx/builtin.h b/include/tvm/tirx/builtin.h index 27033de7165c..e95700a5b402 100644 --- a/include/tvm/tirx/builtin.h +++ b/include/tvm/tirx/builtin.h @@ -47,6 +47,9 @@ namespace builtin { * Tags starting with vthread denote virtual threads. There are no attrs or results. */ TVM_DLL const Op& launch_thread(); + +/*! \brief Mark a user-facing device entry region. */ +TVM_DLL const Op& device_entry(); /*! * \brief Allocate a buffer: alloc_tensor(shape, dtype, scope) -> TensorType. * diff --git a/include/tvm/tirx/script/ir_builder/ir.h b/include/tvm/tirx/script/ir_builder/ir.h index 8bf566d37b32..47015626b078 100644 --- a/include/tvm/tirx/script/ir_builder/ir.h +++ b/include/tvm/tirx/script/ir_builder/ir.h @@ -247,16 +247,8 @@ Var Bind(Expr value, ffi::Optional type_annotation = std::nullopt, */ AttrFrame Attr(ffi::Any node, ffi::String attr_key, Expr value); -/*! - * \brief Mark the device-region entry within the enclosing PrimFunc body. - * Returns an AttrFrame keyed ``tirx.device_entry`` (value ``Bool(true)``). - * Subsequent stmts accumulate into the frame's body; the frame is closed - * by ``PrimFuncFrameNode::ExitWithScope`` which drains leftover frames. - * - * Python sugar: ``Tx.device_entry()`` is a flat-call (no ``with``), which - * auto-enters the frame. - */ -AttrFrame DeviceEntry(); +/*! \brief Create a device-entry region frame. */ +RegionFrame DeviceEntry(); /*! * \brief Create a while loop. diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h index 1ef75452cd15..fb97230872d2 100644 --- a/include/tvm/tirx/stmt.h +++ b/include/tvm/tirx/stmt.h @@ -825,7 +825,7 @@ class Continue : public Stmt { * * Each declaration is a flat stmt within the device-region body. The declared * ``Var``\ s are visible in subsequent stmts in the same enclosing scope - * (the AttrStmt ``kDeviceEntry`` body), analogous to ``BindNode``. + * (the ``tirx.device_entry`` region body), analogous to ``BindNode``. */ class ScopeIdDefStmtNode : public StmtNode { public: @@ -859,8 +859,6 @@ namespace attr { constexpr const char* compute_scope = "compute_scope"; /*! \brief The allocation device for global malloc in host. */ constexpr const char* device_id = "device_id"; -/*! \brief Mark that it is in the device scope. */ -constexpr const char* device_scope = "device_scope"; /*! \brief The device type. */ constexpr const char* device_type = "device_type"; /*! \brief Pragma: auto-unroll, max_step */ @@ -889,14 +887,6 @@ constexpr const char* tensorized_nki_instruction = "tensorized_nki_instruction"; */ constexpr const char* kPersistentKernel = "tirx.persistent_kernel"; -/*! - * \brief Mark the device-region entry within a PrimFunc body. The - * ``AttrStmt`` so-keyed has a body that is the device-side region; anything - * before the marker (within the PrimFunc body) is host code. Value is - * ``IntImm("bool", 1)`` -- a boolean marker, similar to ``kPersistentKernel``. - */ -constexpr const char* kDeviceEntry = "tirx.device_entry"; - /*! * \brief Check if attr_key is a pragma key extension * \param attr_key The attr key to be compared diff --git a/python/tvm/backend/trn/transform/private_buffer_alloc.py b/python/tvm/backend/trn/transform/private_buffer_alloc.py index 11547148fcbb..d08802841481 100644 --- a/python/tvm/backend/trn/transform/private_buffer_alloc.py +++ b/python/tvm/backend/trn/transform/private_buffer_alloc.py @@ -16,10 +16,9 @@ # under the License. import tvm_ffi -from tvm.ir import Call, DataTypeImm, DictAttrs, Range, StringImm, Tuple, Var +from tvm.ir import Call, DataTypeImm, DictAttrs, Op, Range, StringImm, Tuple, Var from tvm.target import Target from tvm.tirx.stmt import ( - AttrStmt, Bind, For, RegionStmt, @@ -87,11 +86,11 @@ def _inject_private_allocations( ) -> Stmt: is_outer_block = True - def visit_attr(op: AttrStmt): + def visit_region(op: RegionStmt): nonlocal is_outer_block - # AttrStmt(kDeviceEntry) marks the device-region root: inject the + # The device-entry region marks the root: inject the # collected init stmts + alloc_buffers into its body. - if op.attr_key == "tirx.device_entry": + if op.op.same_as(Op.get("tirx.device_entry")): is_outer = is_outer_block is_outer_block = False if is_outer: @@ -113,7 +112,9 @@ def visit_attr(op: AttrStmt): ), ) body = SeqStmt([allocation, body]) - return AttrStmt(op.node, op.attr_key, op.value, body) + return RegionStmt( + op.op, op.args, op.body_params, op.attrs, body, op.result_vars, op.span + ) return op def visit_op_call(op: TilePrimitiveCall): @@ -125,7 +126,7 @@ def visit_op_call(op: TilePrimitiveCall): return tvm_ffi.structural_map( stmt, - [(AttrStmt, visit_attr), (TilePrimitiveCall, visit_op_call)], + [(RegionStmt, visit_region), (TilePrimitiveCall, visit_op_call)], order="pre", ) diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index b9477286f85f..711b7a1b76cd 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -149,27 +149,14 @@ def func_attr(attrs: dict[str, Any]) -> None: _ffi_api.FuncAttrs(attrs) # type: ignore[attr-defined] # pylint: disable=no-member -def device_entry() -> None: - """Mark the device-region entry within the enclosing PrimFunc body. +def device_entry() -> frame.RegionFrame: + """Mark a device-entry region containing scope definitions. - Flat marker (no ``with``). Subsequent statements in the function body - accumulate into an ``AttrStmt("tirx.device_entry", True, body=...)``; - the wrapping is closed by the PrimFunc frame at function end. - - Anything written before this marker is host code (e.g. buffer layout setup); - anything after is device code. - - Example:: - - @T.prim_func - def kernel(...): - A = T.Tensor(...) - T.device_entry() # device region starts here - bx = T.cta_id([SM_COUNT]) # standalone scope-id def - ... + Use a flat ``T.device_entry()`` to scope the remaining statements in the + enclosing body, or ``with T.device_entry():`` for an explicit boundary. + Statements before the region remain host code. """ - attr_frame = _ffi_api.DeviceEntry() # type: ignore[attr-defined] # pylint: disable=no-member - attr_frame.__enter__() + return _ffi_api.DeviceEntry() def check_well_formed_(function: _tir.PrimFunc) -> None: diff --git a/src/s_tir/transform/decorate_device_scope.cc b/src/s_tir/transform/decorate_device_scope.cc index 80c5a51c8931..39ce7b7c86b9 100644 --- a/src/s_tir/transform/decorate_device_scope.cc +++ b/src/s_tir/transform/decorate_device_scope.cc @@ -31,8 +31,8 @@ namespace s_tir { using namespace tvm::tirx; Stmt DecorateDeviceScopeImpl(Stmt&& stmt) { - Stmt body = AttrStmt(0, tirx::attr::device_scope, IntImm::Int32(0), stmt); - return body; + static const Op device_scope = Op::Get("tirx.device_scope"); + return RegionStmt(device_scope, {}, {}, DictAttrs(), std::move(stmt)); } namespace transform { diff --git a/src/tirx/analysis/verify_tirx_well_formed.cc b/src/tirx/analysis/verify_tirx_well_formed.cc index 94d827bab8a2..1e3acb4bdef9 100644 --- a/src/tirx/analysis/verify_tirx_well_formed.cc +++ b/src/tirx/analysis/verify_tirx_well_formed.cc @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -65,10 +66,10 @@ class ScopeIdVerifier : public Verifier { private: using Verifier::Visit; - void Dispatch_(const AttrStmtNode* op, ffi::reflection::AccessPath path) override { - if (op->attr_key == tvm::tirx::attr::kDeviceEntry) { + void Dispatch_(const RegionStmtNode* op, ffi::reflection::AccessPath path) override { + if (op->op.same_as(tirx::builtin::device_entry())) { // Device-region marker: defs gathered from the body are verified when - // the AttrStmt exits, with launch-param sanity enforced as ``is_root``. + // the region exits, with launch-param sanity enforced as ``is_root``. size_t baseline = scope_id_def_.size(); Verifier::Dispatch_(op, path); size_t total = scope_id_def_.size(); diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc index 100f78a4dc99..d11518e97393 100644 --- a/src/tirx/ir/stmt.cc +++ b/src/tirx/ir/stmt.cc @@ -809,6 +809,10 @@ RegionStmt::RegionStmt(Op op, ffi::Array args, ffi::Array body_params : Stmt(ffi::UnsafeInit{}) { TVM_FFI_CHECK(op.defined() && body.defined(), ValueError) << "RegionStmt requires an operator and a body"; + CallNode signature(op); + signature.args = args; + signature.attrs = attrs; + op.Validate(&signature); std::unordered_set definitions; for (const auto& vars : {body_params, result_vars}) { for (const Var& var : vars) { @@ -816,6 +820,14 @@ RegionStmt::RegionStmt(Op op, ffi::Array args, ffi::Array body_params << "RegionStmt parameters and results must be distinct definitions"; } } + static const Op device_scope = Op::Get("tirx.device_scope"); + if (op.same_as(tirx::builtin::device_entry()) || op.same_as(device_scope)) { + TVM_FFI_CHECK(body_params.empty() && result_vars.empty(), ValueError) + << op->name << " expects no body parameters or results"; + if (op.same_as(tirx::builtin::device_entry())) { + TVM_FFI_CHECK(attrs->dict.empty(), ValueError) << "device_entry expects no attrs"; + } + } if (op.same_as(tirx::builtin::launch_thread())) { TVM_FFI_CHECK( args.size() == 2 && body_params.size() == 1 && result_vars.empty() && attrs->dict.empty(), diff --git a/src/tirx/op/builtin.cc b/src/tirx/op/builtin.cc index 3309f9bded2c..f3c7968ab2fd 100644 --- a/src/tirx/op/builtin.cc +++ b/src/tirx/op/builtin.cc @@ -189,6 +189,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { TVM_DEFINE_CACHED_OP_GETTER(reinterpret, "tirx.reinterpret") TVM_DEFINE_CACHED_OP_GETTER(launch_thread, "tirx.launch_thread") +TVM_DEFINE_CACHED_OP_GETTER(device_entry, "tirx.device_entry") TVM_DEFINE_CACHED_OP_GETTER(thread_return, "tirx.thread_return") TVM_DEFINE_CACHED_OP_GETTER(filter, "tirx.filter") TVM_DEFINE_CACHED_OP_GETTER(selector, "tirx.selector") @@ -267,6 +268,10 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TScriptDtypePrintLocation", static_cast(ScriptDtypePrintLocation::kFirst)); + OpDef("tirx.device_entry", "Mark a device entry containing scope definitions.") + .signature() + .set_attr("TIRxOpCategory", ffi::String("builtin")); + OpDef("tirx.launch_thread", "Bind a thread index within a body with a launch extent.") .set_attr("TIRxOpCategory", ffi::String("builtin")); OpDef("tirx.thread_return") diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index 5e7f6c0f1294..dd06936f73f9 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -439,32 +439,7 @@ AttrFrame Attr(ffi::Any node, ffi::String attr_key, Expr value) { return AttrFrame(n); } -AttrFrame DeviceEntry() { - // Flat marker: open an AttrFrame keyed ``tirx.device_entry`` with - // ``Bool(true)`` value. Subsequent stmts within the enclosing PrimFunc - // body accumulate into this frame's body. The Python wrapper auto-calls - // ``__enter__`` so users write a flat ``Tx.device_entry()`` (no ``with``). - // To close the AttrFrame at function end, register a callback on the - // enclosing PrimFuncFrame: ``IRBuilderFrameNode::ExitWithScope`` runs - // callbacks before popping itself, so the AttrFrame is closed and its - // emitted ``AttrStmt`` lands in the PrimFunc's body sequence. - AttrFrame frame = Attr(0, ffi::String(tvm::tirx::attr::kDeviceEntry), IntImm::Bool(true)); - IRBuilder builder = IRBuilder::Current(); - ffi::Optional pf_frame = builder->FindFrame(); - TVM_FFI_ICHECK(pf_frame.has_value()) - << "T.device_entry() must be called inside a @T.prim_func body"; - // Capture the AttrFrame by ObjectRef value so the lambda holds a strong - // reference while the callback runs. Without this, the only reference is - // the IRBuilder frame stack; ``ExitWithScope`` pops itself first and the - // AttrFrameNode would be destroyed mid-method (before the body-wrapping - // AddToParent runs). - AttrFrame frame_ref = frame; - pf_frame.value()->callbacks.push_back([frame_ref]() { - const_cast(static_cast(frame_ref.get())) - ->ExitWithScope(); - }); - return frame; -} +RegionFrame DeviceEntry() { return Region(tvm::tirx::builtin::device_entry(), {}, {}); } WhileFrame While(PrimExpr condition) { ffi::ObjectPtr n = ffi::make_object(condition); diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc index 81ba314a7247..5088fe5289b7 100644 --- a/src/tirx/script/printer/stmt.cc +++ b/src/tirx/script/printer/stmt.cc @@ -360,7 +360,9 @@ ffi::Optional RegionStmtDocTranslate(DocTranslatorObj* d, ffi::AnyView // Inputs, attributes, and parameter types are evaluated before the body // parameters enter scope. Explicit Var constructors preserve their exact types. ExprDoc rhs(ffi::UnsafeInit{}); - if (stmt->op.same_as(tirx::builtin::launch_thread())) { + if (stmt->op.same_as(tirx::builtin::device_entry())) { + rhs = NamespaceDoc("tirx")->Attr("device_entry")->Call({}); + } else if (stmt->op.same_as(tirx::builtin::launch_thread())) { rhs = NamespaceDoc("tirx") ->Attr("launch_thread") ->Call({LiteralDoc::Str(stmt->args[0].as_or_throw()->value, std::nullopt), diff --git a/src/tirx/transform/bind_target.cc b/src/tirx/transform/bind_target.cc index e84be7151e92..2855b0575cd0 100644 --- a/src/tirx/transform/bind_target.cc +++ b/src/tirx/transform/bind_target.cc @@ -116,7 +116,9 @@ class FunctionClassifierVisitor : public StmtExprVisitor { } ffi::Optional Visit_(const RegionStmtNode* op) final { - if (!op->op.same_as(tirx::builtin::launch_thread())) return StmtExprVisitor::Visit_(op); + if (!op->op.same_as(tirx::builtin::launch_thread()) && + !op->op.same_as(tirx::builtin::device_entry())) + return StmtExprVisitor::Visit_(op); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(Visit(op->args)); bool previous_scope = is_under_gpu_scope_; is_under_gpu_scope_ = true; @@ -125,19 +127,6 @@ class FunctionClassifierVisitor : public StmtExprVisitor { return result; } - ffi::Optional Visit_(const AttrStmtNode* op) final { - if (op->attr_key == attr::kDeviceEntry) { - // Enter the explicit device scope - bool last_is_under_gpu_scope = is_under_gpu_scope_; - is_under_gpu_scope_ = true; - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op)); - is_under_gpu_scope_ = last_is_under_gpu_scope; - } else { - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op)); - } - return std::nullopt; - } - private: /*! \brief Whether the current statement is under a GPU scope */ bool is_under_gpu_scope_ = false; @@ -214,30 +203,17 @@ class CallSubstitutor : public StmtExprMutator { } UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) final { - if (!op->op.same_as(tirx::builtin::launch_thread())) + if (!op->op.same_as(tirx::builtin::launch_thread()) && + !op->op.same_as(tirx::builtin::device_entry())) return StmtExprMutator::Mutate_(op, inplace_mode); - PrimExpr old_extent = op->args[1].as_or_throw(); - PrimExpr extent = Mutate(old_extent, inplace_mode).ValueOrUnchanged(old_extent); + auto args = Mutate(op->args, inplace_mode).ValueOrUnchanged(op->args); bool previous_scope = is_under_gpu_scope_; is_under_gpu_scope_ = true; Stmt body = Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); is_under_gpu_scope_ = previous_scope; - return RegionStmt(op->op, {op->args[0], extent}, op->body_params, op->attrs, body, - op->result_vars, op->span); + return RegionStmt(op->op, args, op->body_params, op->attrs, body, op->result_vars, op->span); } - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == attr::kDeviceEntry) { - // Enter the explicit device scope - bool last_is_under_gpu_scope = is_under_gpu_scope_; - is_under_gpu_scope_ = true; - UnchangedOr stmt = StmtExprMutator::Mutate_(op, inplace_mode); - is_under_gpu_scope_ = last_is_under_gpu_scope; - return stmt; - } else { - return StmtExprMutator::Mutate_(op, inplace_mode); - } - } /*! \brief Whether the current statement is under a GPU scope */ bool is_under_gpu_scope_ = false; /*! \brief Mapping from original functions to host-specific duplicates */ diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index 5858c380f7d1..ddd8f0f4bfb0 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -47,6 +47,11 @@ namespace tvm { namespace tirx { +TVM_FFI_STATIC_INIT_BLOCK() { + OpDef("tirx.device_scope", "Internal host/device splitting boundary.") + .signature(sig::call_attrs()); +} + // Device-region annotation class DeviceRegionAnnotater : public StmtExprMutator { @@ -60,27 +65,20 @@ class DeviceRegionAnnotater : public StmtExprMutator { explicit DeviceRegionAnnotater(Target device_target) : device_target_(device_target) {} UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) final { + static const Op device_scope = Op::Get("tirx.device_scope"); + if (op->op.same_as(device_scope)) { + if (op->attrs->dict.count(tvm::attr::kTarget)) return ffi::Unchanged(); + return RegionStmt(op->op, op->args, op->body_params, + DictAttrs({{tvm::attr::kTarget, device_target_}}), op->body, + op->result_vars, op->span); + } if (op->op.same_as(tirx::builtin::launch_thread())) { - return AttrStmt(device_target_, tvm::attr::kTarget, IntImm::Int32(0), ffi::GetRef(op)); + return RegionStmt(device_scope, {}, {}, DictAttrs({{tvm::attr::kTarget, device_target_}}), + ffi::GetRef(op)); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == tvm::attr::kTarget) { - // If a target attribute already exists, use it as-is. - return ffi::Unchanged(); - } else if (op->attr_key == attr::device_scope) { - // These attributes are only allowed in device-side code, so - // they should be annotated with the function's default target. - Stmt body = ffi::GetRef(op); - return AttrStmt(device_target_, tvm::attr::kTarget, IntImm::Int32(0), body); - } else { - // All other annotations are ignored. - return StmtExprMutator::Mutate_(op, inplace_mode); - } - } - private: Target device_target_; }; @@ -206,15 +204,30 @@ class HostDeviceSplitter : public StmtExprMutator { PrimFunc cur_func) : device_mod_(device_mod), var_supply_(var_supply), cur_func_(cur_func) {} - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == tvm::attr::kTarget) { - auto device_target = op->node.as().value().WithoutHost(); - return SplitDeviceFunc(op->body, device_target); + UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) final { + static const Op device_scope = Op::Get("tirx.device_scope"); + if (op->op.same_as(device_scope)) { + auto target = op->attrs->dict.Get(tvm::attr::kTarget); + if (!target) return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); + Target device_target = target.value().as_or_throw(); + return SplitDeviceFunc(op->body, device_target.WithoutHost()); } return StmtExprMutator::Mutate_(op, inplace_mode); } private: + class KernelBodyRewriter : public StmtExprMutator { + public: + using StmtExprMutator::Mutate_; + UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) final { + static const Op device_scope = Op::Get("tirx.device_scope"); + if (op->op.same_as(device_scope)) { + return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); + } + return StmtExprMutator::Mutate_(op, inplace_mode); + } + }; + Stmt SplitDeviceFunc(Stmt body, Target device_target) { auto [params, buffers_to_declare] = [&]() -> std::tuple, ffi::Array> { @@ -254,7 +267,7 @@ class HostDeviceSplitter : public StmtExprMutator { ffi::Array kernel_params; ffi::Array call_args; ffi::Map buffer_data_params; - auto kernel_rewriter = ffi::make_object(); + auto kernel_rewriter = ffi::make_object(); for (const Var& param : params) { if (param->ty.as()) { TensorVar buffer = param.as_or_throw(); diff --git a/src/tirx/transform/tile_primitive_dispatch.cc b/src/tirx/transform/tile_primitive_dispatch.cc index 0f7bceeb82fa..fa9b160d1596 100644 --- a/src/tirx/transform/tile_primitive_dispatch.cc +++ b/src/tirx/transform/tile_primitive_dispatch.cc @@ -67,8 +67,8 @@ class ScopeIdDefGather : public StmtExprVisitor { return std::move(gather->out_); } - ffi::Optional Visit_(const AttrStmtNode* op) override { - if (op->attr_key == tvm::tirx::attr::kDeviceEntry) { + ffi::Optional Visit_(const RegionStmtNode* op) override { + if (op->op.same_as(tirx::builtin::device_entry())) { return EnterSourceAndPartition(op, [&]() { return StmtExprVisitor::Visit_(op); }); } return StmtExprVisitor::Visit_(op); @@ -164,8 +164,8 @@ class ScopeIdVarFinder : public StmtExprVisitor { bool found_{false}; }; -// Remove any standalone ``ScopeIdDefStmt`` nodes; the resolved values are -// bound at kernel scope via Bind statements emitted separately. +// Remove resolved scope definitions and device-entry boundaries after gathering. +// Their values are bound at kernel scope via Bind statements emitted separately. class ScopeIdDefRemover : public StmtExprMutator { public: using StmtExprMutator::Mutate; @@ -176,6 +176,13 @@ class ScopeIdDefRemover : public StmtExprMutator { .ValueOrUnchanged(stmt); } + UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) override { + if (op->op.same_as(tirx::builtin::device_entry())) { + return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); + } + return StmtExprMutator::Mutate_(op, inplace_mode); + } + UnchangedOr Mutate_(const ScopeIdDefStmtNode* op, InplaceMode inplace_mode) override { // Drop the def stmt by replacing with a no-op Evaluate(0). It will be // flattened away by SeqStmt::Flatten elsewhere or stay as a benign @@ -251,14 +258,14 @@ class TilePrimitiveDispatcher : public StmtExprMutator { Stmt body_; }; - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == tirx::attr::kDeviceEntry) { + UnchangedOr Mutate_(const RegionStmtNode* op, InplaceMode inplace_mode) final { + if (op->op.same_as(tirx::builtin::device_entry())) { return ProcessDeviceEntry(op); } return StmtExprMutator::Mutate_(op, inplace_mode); } - Stmt ProcessDeviceEntry(const AttrStmtNode* entry_node) { + Stmt ProcessDeviceEntry(const RegionStmtNode* entry_node) { Stmt body_to_visit = entry_node->body; bool is_first_block = false; @@ -296,8 +303,8 @@ class TilePrimitiveDispatcher : public StmtExprMutator { if (body_unchanged) { return ffi::GetRef(entry_node); } - return AttrStmt(entry_node->node, entry_node->attr_key, entry_node->value, body, - entry_node->span); + return RegionStmt(entry_node->op, entry_node->args, entry_node->body_params, + entry_node->attrs, body, entry_node->result_vars, entry_node->span); } // Insert device init stmts into kernel body. @@ -617,9 +624,9 @@ class TilePrimitiveDispatcher : public StmtExprMutator { // resolution that used to live here is now in ``ResolveAllScopeBinds``, // which runs AFTER dispatch so it sees ScopeIdDefs introduced by // dispatched impls too. - void PrepareLaunchParams(const AttrStmtNode* entry_node, Stmt body, + void PrepareLaunchParams(const RegionStmtNode* entry_node, Stmt body, std::vector>* scope_binds) { - Stmt gather_target = AttrStmt(0, tvm::tirx::attr::kDeviceEntry, IntImm::Bool(true), body); + Stmt gather_target = RegionStmt(tirx::builtin::device_entry(), {}, {}, DictAttrs(), body); std::vector gathered = ScopeIdDefGather::Gather(gather_target); Array defs; defs.reserve(gathered.size()); @@ -647,7 +654,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { void ResolveAllScopeBinds(Stmt body, std::vector>* scope_binds) { // Gather from a temporary stmt synthesized as the device-entry marker // to retain nested-before-direct declaration order. - Stmt gather_target = AttrStmt(0, tvm::tirx::attr::kDeviceEntry, IntImm::Bool(true), body); + Stmt gather_target = RegionStmt(tirx::builtin::device_entry(), {}, {}, DictAttrs(), body); std::vector gathered = ScopeIdDefGather::Gather(gather_target); Array defs; defs.reserve(gathered.size()); diff --git a/tests/python/s_tir/transform/test_s_tir_transform_decorate_device_scope.py b/tests/python/s_tir/transform/test_s_tir_transform_decorate_device_scope.py index 6c5e39415ffd..9ac231a6faa1 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_decorate_device_scope.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_decorate_device_scope.py @@ -22,7 +22,7 @@ def test_decorate_device(): mod = tvm.IRModule.from_expr(tvm.tirx.PrimFunc([x], tvm.tirx.Evaluate(x))) stmt = tvm.s_tir.transform.DecorateDeviceScope()(mod)["main"].body - assert stmt.attr_key == "device_scope" + assert stmt.op.same_as(tvm.ir.Op.get("tirx.device_scope")) if __name__ == "__main__": diff --git a/tests/python/tirx-transform/test_tir_transform_split_host_device.py b/tests/python/tirx-transform/test_tir_transform_split_host_device.py index 5513764ed73e..bd95398420fe 100644 --- a/tests/python/tirx-transform/test_tir_transform_split_host_device.py +++ b/tests/python/tirx-transform/test_tir_transform_split_host_device.py @@ -54,7 +54,7 @@ class before: def main(): T.func_attr({"global_symbol": "main", "target": T.target("cuda", host="llvm")}) for i in range(16): - T.attr(0, "device_scope", 0) + T.region("tirx.device_scope", []) for j in range(16): T.evaluate(i) @@ -73,7 +73,7 @@ class Before: @T.prim_func def main(n: T.int32): T.func_attr({"target": T.target("cuda", host={"kind": "llvm", "opt-level": 0})}) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) T.evaluate(n) @I.ir_module @@ -109,7 +109,7 @@ class Before: @T.prim_func def main(n: T.int32): T.func_attr({"target": T.target("cuda", host={"kind": "llvm", "opt-level": 0})}) - T.attr(T.target("llvm"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("llvm")}) T.evaluate(n) @I.ir_module @@ -141,10 +141,11 @@ def test_device_kernel_nonzero_return_is_rejected(): device_target = tvm.target.Target({"kind": "cuda", "arch": "sm_100a"}) target = tvm.target.Target(device_target, host="llvm") - body = tvm.tirx.AttrStmt( - device_target, - "target", - 0, + body = tvm.tirx.RegionStmt( + tvm.ir.Op.get("tirx.device_scope"), + [], + [], + tvm.ir.DictAttrs({"target": device_target}), tvm.tirx.Return(tvm.tirx.IntImm("int32", 1)), ) func = tvm.tirx.PrimFunc([], body) @@ -167,7 +168,7 @@ class Before: @T.prim_func def main(n: T.int32): T.func_attr({"target": T.target("llvm")}) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) T.evaluate(n) @I.ir_module @@ -227,7 +228,7 @@ class Before: @T.prim_func def main(n: T.int32): T.func_attr({"target": T.target("cuda", host={"kind": "llvm", "opt-level": 0})}) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) T.evaluate(n) @T.prim_func @@ -295,7 +296,7 @@ def default_function( T.func_attr({"target": T.target("cuda")}) num_blocks: T.let[T.int32] = (seq_len + 127) // 128 - with T.attr(T.target("cuda"), "target", 0): + with T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}): blockIdx_x = T.launch_thread("blockIdx.x", num_blocks) threadIdx_x = T.launch_thread("threadIdx.x", 128) if blockIdx_x * 128 + threadIdx_x < seq_len: @@ -350,7 +351,7 @@ class Module: def main(A: T.Tensor((m,)), B: T.Tensor((m,))): T.func_attr({"target": T.target("cuda")}) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) blockIdx_x = T.launch_thread("blockIdx.x", m) B_1 = T.decl_tensor((m,), data=B.data) A_1 = T.decl_tensor((m,), data=A.data) @@ -367,7 +368,7 @@ class Before: @T.prim_func def main(A: T.Tensor((16,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - with T.attr(T.target("cuda"), "target", 0): + with T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}): T.evaluate(T.call_extern("consume", A.data, dtype="int32")) after = tvm.tirx.transform.SplitHostDevice()(Before) @@ -435,7 +436,7 @@ def main(A: T.Tensor(16, "float32")): ], } ) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) tx = T.launch_thread("threadIdx.x", 16) A[tx] = 0.0 @@ -469,7 +470,7 @@ class Before: @T.prim_func def main(A: T.Tensor(4, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - T.attr(T.target("cuda"), "target", 0) + T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) T.attr(0, "tirx.required_block_size", 1) with T.attr(0, "tirx.launch_bounds_min_blocks_per_sm", 1): bx = T.launch_thread("blockIdx.x", 4) @@ -501,7 +502,7 @@ class Before: @T.prim_func def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - with T.attr(T.target("cuda"), "target", 0): + with T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}): T.launch_thread("blockIdx.x", 4) T.launch_thread("clusterCtaIdx.x", 1) T.launch_thread("clusterCtaIdx.y", 1) @@ -532,7 +533,7 @@ class Before: @T.prim_func def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - T.attr(0, "device_scope", 0) + T.region("tirx.device_scope", []) A[0] = 0.0 @I.ir_module @@ -555,7 +556,6 @@ def main_kernel(A_data: T.handle("float32")): } ) A = T.decl_tensor(1, dtype="float32", data=A_data) - T.attr(0, "device_scope", 0) A[0] = 0.0 After = tvm.tirx.transform.SplitHostDevice()(Before) 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 9759510130cf..5a15976772f9 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 b01f6a4931ad..881fd8735b93 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 3d0953337c6a..c34fb6d5de8c 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 e39e37198de5..eae318588bb7 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 13cf4e74302d..dab2091f70ec 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 19d8bc6d70bb..bf5cf5527a88 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) 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 c61c2e07caec..8226d9d2b0c0 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 @@ -28,14 +28,14 @@ def _strip_exec_scope_stmt(stmt): - def _strip_attr(node: tvm.tirx.AttrStmt): - if node.attr_key == "tirx.device_entry": + def _strip_region(node: tvm.tirx.RegionStmt): + if node.op.same_as(tvm.ir.Op.get("tirx.device_entry")): return node.body return node return tvm_ffi.structural_map( stmt, - (tvm.tirx.AttrStmt, _strip_attr), + (tvm.tirx.RegionStmt, _strip_region), order="post", ) From ccf7eb374f7374d9eb434898575284de2a3e9711 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Tue, 6 Oct 2026 04:12:19 +0000 Subject: [PATCH 2/4] [TIRx] Preserve typed operands when rewriting device regions --- src/tirx/transform/bind_target.cc | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/tirx/transform/bind_target.cc b/src/tirx/transform/bind_target.cc index 2855b0575cd0..c80f277b6aa9 100644 --- a/src/tirx/transform/bind_target.cc +++ b/src/tirx/transform/bind_target.cc @@ -206,7 +206,8 @@ class CallSubstitutor : public StmtExprMutator { if (!op->op.same_as(tirx::builtin::launch_thread()) && !op->op.same_as(tirx::builtin::device_entry())) return StmtExprMutator::Mutate_(op, inplace_mode); - auto args = Mutate(op->args, inplace_mode).ValueOrUnchanged(op->args); + auto args = + Mutate(op->args, inplace_mode).ValueOrUnchanged(op->args).as_or_throw>(); bool previous_scope = is_under_gpu_scope_; is_under_gpu_scope_ = true; Stmt body = Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); From 2897ff8b40d725dc1f025ad2fa9c9f373b137d61 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Tue, 6 Oct 2026 05:40:01 +0000 Subject: [PATCH 3/4] [TIRx] Classify internal device scope as a builtin region --- src/tirx/transform/split_host_device.cc | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index ddd8f0f4bfb0..409afa600a77 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -49,7 +49,8 @@ namespace tirx { TVM_FFI_STATIC_INIT_BLOCK() { OpDef("tirx.device_scope", "Internal host/device splitting boundary.") - .signature(sig::call_attrs()); + .signature(sig::call_attrs()) + .set_attr("TIRxOpCategory", ffi::String("builtin")); } // Device-region annotation From a21157da4a2de540e54db93df6db4cd5572a823c Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Tue, 6 Oct 2026 12:41:42 +0000 Subject: [PATCH 4/4] [TIRx] Reuse generic region builder for device entry --- include/tvm/tirx/script/ir_builder/ir.h | 3 --- python/tvm/tirx/script/ir_builder/parser_protocol.py | 2 +- src/tirx/script/ir_builder/ir.cc | 3 --- 3 files changed, 1 insertion(+), 7 deletions(-) diff --git a/include/tvm/tirx/script/ir_builder/ir.h b/include/tvm/tirx/script/ir_builder/ir.h index 47015626b078..07122fb5dd12 100644 --- a/include/tvm/tirx/script/ir_builder/ir.h +++ b/include/tvm/tirx/script/ir_builder/ir.h @@ -247,9 +247,6 @@ Var Bind(Expr value, ffi::Optional type_annotation = std::nullopt, */ AttrFrame Attr(ffi::Any node, ffi::String attr_key, Expr value); -/*! \brief Create a device-entry region frame. */ -RegionFrame DeviceEntry(); - /*! * \brief Create a while loop. * \param condition The termination condition of the loop. diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 711b7a1b76cd..4f88013707c1 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -156,7 +156,7 @@ def device_entry() -> frame.RegionFrame: enclosing body, or ``with T.device_entry():`` for an explicit boundary. Statements before the region remain host code. """ - return _ffi_api.DeviceEntry() + return region("tirx.device_entry", []) def check_well_formed_(function: _tir.PrimFunc) -> None: diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index dd06936f73f9..636e079ff3ac 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -439,8 +439,6 @@ AttrFrame Attr(ffi::Any node, ffi::String attr_key, Expr value) { return AttrFrame(n); } -RegionFrame DeviceEntry() { return Region(tvm::tirx::builtin::device_entry(), {}, {}); } - WhileFrame While(PrimExpr condition) { ffi::ObjectPtr n = ffi::make_object(condition); return WhileFrame(n); @@ -708,7 +706,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { .def("script.ir_builder.tirx.Assert", Assert) .def("script.ir_builder.tirx.Bind", Bind) .def("script.ir_builder.tirx.Attr", Attr) - .def("script.ir_builder.tirx.DeviceEntry", DeviceEntry) .def("script.ir_builder.tirx.While", While) .def("script.ir_builder.tirx.Return", Return) .def("script.ir_builder.tirx.Break", Break)