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
15 changes: 8 additions & 7 deletions docs/tirx/native_basics/cuda/functions.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
~~~~~~~~~
Expand Down
3 changes: 3 additions & 0 deletions include/tvm/tirx/builtin.h
Original file line number Diff line number Diff line change
Expand Up @@ -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.
*
Expand Down
11 changes: 0 additions & 11 deletions include/tvm/tirx/script/ir_builder/ir.h
Original file line number Diff line number Diff line change
Expand Up @@ -247,17 +247,6 @@ Var Bind(Expr value, ffi::Optional<Type> 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 while loop.
* \param condition The termination condition of the loop.
Expand Down
12 changes: 1 addition & 11 deletions include/tvm/tirx/stmt.h
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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 */
Expand Down Expand Up @@ -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
Expand Down
15 changes: 8 additions & 7 deletions python/tvm/backend/trn/transform/private_buffer_alloc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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):
Expand All @@ -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",
)

Expand Down
25 changes: 6 additions & 19 deletions python/tvm/tirx/script/ir_builder/parser_protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 region("tirx.device_entry", [])


def check_well_formed_(function: _tir.PrimFunc) -> None:
Expand Down
4 changes: 2 additions & 2 deletions src/s_tir/transform/decorate_device_scope.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
7 changes: 4 additions & 3 deletions src/tirx/analysis/verify_tirx_well_formed.cc
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
#include <tvm/runtime/logging.h>
#include <tvm/sym/analyzer.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/builtin.h>
#include <tvm/tirx/exec_scope.h>
#include <tvm/tirx/op_attr_types.h>
#include <tvm/tirx/stmt.h>
Expand Down Expand Up @@ -65,10 +66,10 @@ class ScopeIdVerifier : public Verifier<ScopeIdVerifier> {
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();
Expand Down
12 changes: 12 additions & 0 deletions src/tirx/ir/stmt.cc
Original file line number Diff line number Diff line change
Expand Up @@ -809,13 +809,25 @@ RegionStmt::RegionStmt(Op op, ffi::Array<Expr> args, ffi::Array<Var> 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<const VarNode*> definitions;
for (const auto& vars : {body_params, result_vars}) {
for (const Var& var : vars) {
TVM_FFI_CHECK(definitions.insert(var.get()).second, ValueError)
<< "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(),
Expand Down
5 changes: 5 additions & 0 deletions src/tirx/op/builtin.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -267,6 +268,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
.set_attr<TScriptDtypePrintLocation>("TScriptDtypePrintLocation",
static_cast<int64_t>(ScriptDtypePrintLocation::kFirst));

OpDef("tirx.device_entry", "Mark a device entry containing scope definitions.")
.signature()
.set_attr<TIRxOpCategory>("TIRxOpCategory", ffi::String("builtin"));

OpDef("tirx.launch_thread", "Bind a thread index within a body with a launch extent.")
.set_attr<TIRxOpCategory>("TIRxOpCategory", ffi::String("builtin"));
OpDef("tirx.thread_return")
Expand Down
28 changes: 0 additions & 28 deletions src/tirx/script/ir_builder/ir.cc
Original file line number Diff line number Diff line change
Expand Up @@ -439,33 +439,6 @@ 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<PrimFuncFrame> pf_frame = builder->FindFrame<PrimFuncFrame>();
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<IRBuilderFrameNode*>(static_cast<const IRBuilderFrameNode*>(frame_ref.get()))
->ExitWithScope();
});
return frame;
}

WhileFrame While(PrimExpr condition) {
ffi::ObjectPtr<WhileFrameNode> n = ffi::make_object<WhileFrameNode>(condition);
return WhileFrame(n);
Expand Down Expand Up @@ -733,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)
Expand Down
4 changes: 3 additions & 1 deletion src/tirx/script/printer/stmt.cc
Original file line number Diff line number Diff line change
Expand Up @@ -360,7 +360,9 @@ ffi::Optional<ExprDoc> 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<StringImm>()->value, std::nullopt),
Expand Down
39 changes: 8 additions & 31 deletions src/tirx/transform/bind_target.cc
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,9 @@ class FunctionClassifierVisitor : public StmtExprVisitor {
}

ffi::Optional<VisitInterrupt> 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;
Expand All @@ -125,19 +127,6 @@ class FunctionClassifierVisitor : public StmtExprVisitor {
return result;
}

ffi::Optional<VisitInterrupt> 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;
Expand Down Expand Up @@ -214,30 +203,18 @@ class CallSubstitutor : public StmtExprMutator {
}

UnchangedOr<Stmt> 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>();
PrimExpr extent = Mutate(old_extent, inplace_mode).ValueOrUnchanged(old_extent);
auto args =
Mutate(op->args, inplace_mode).ValueOrUnchanged(op->args).as_or_throw<ffi::Array<Expr>>();
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<Stmt> 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> 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 */
Expand Down
Loading
Loading