Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
39 changes: 0 additions & 39 deletions include/tvm/tirx/script/ir_builder/frame.h
Original file line number Diff line number Diff line change
Expand Up @@ -531,45 +531,6 @@ class DeclBufferFrame : public TIRFrame {
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(DeclBufferFrame, TIRFrame, DeclBufferFrameNode);
};

/*!
* \brief A frame that represents a hint directive for the sketch language.
*
* \sa HintFrame
*/
class HintFrameNode : public TIRFrameNode {
public:
/*! \brief The free-form hint message string. */
ffi::String message;
/*! \brief Optional structured key-value attributes. */
ffi::Map<ffi::String, ffi::Any> attrs;

static void RegisterReflection() {
namespace refl = tvm::ffi::reflection;
refl::ObjectDef<HintFrameNode>()
.def_ro("message", &HintFrameNode::message)
.def_ro("attrs", &HintFrameNode::attrs);
}
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.HintFrame", HintFrameNode,
TIRFrameNode);

public:
void ExitWithScope() final;
};

/*!
* \brief Managed reference to HintFrameNode.
*
* \sa HintFrameNode
*/
class HintFrame : public TIRFrame {
public:
explicit HintFrame(ffi::ObjectPtr<HintFrameNode> data) : TIRFrame(ffi::UnsafeInit{}) {
TVM_FFI_ICHECK(data != nullptr);
data_ = std::move(data);
}
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(HintFrame, TIRFrame, HintFrameNode);
};

} // namespace tirx
} // namespace ir_builder
} // namespace script
Expand Down
6 changes: 0 additions & 6 deletions include/tvm/tirx/stmt.h
Original file line number Diff line number Diff line change
Expand Up @@ -809,12 +809,6 @@ constexpr const char* device_id = "device_id";
constexpr const char* device_scope = "device_scope";
/*! \brief The device type. */
constexpr const char* device_type = "device_type";
/*!
* \brief Mark the scope as generated by extern primitive.
* Such scope can contain arbitrary ir program and we need to be careful
* when making certain assumptions about the structure of the program.
*/
constexpr const char* extern_scope = "extern_scope";
/*! \brief Pragma: auto-unroll, max_step */
constexpr const char* pragma_auto_unroll_max_step = "pragma_auto_unroll_max_step";
/*! \brief Import C source or file into the final code gen module */
Expand Down
33 changes: 20 additions & 13 deletions python/tvm/backend/cuda/lang/alloc_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,17 +400,18 @@ class SMEMPool:
Parameters
----------
ptr : Var or None, optional
If omitted, an ``alloc_buffer([0], "uint8", scope="shared.dyn")`` is
created automatically and ``commit()`` must be called after all
allocations to emit the size annotation.
If omitted, a ``uint8`` backing allocation in ``shared.dyn`` is created
automatically. Call ``commit()`` after all allocations to finalize its
byte extent before the enclosing builder scope closes.
If a ``Var`` is provided, the caller manages the backing buffer and
``commit()`` is a no-op.
"""

def __init__(self, ptr=_POOL_UNSET):
ir = _get_ir()
self._committed_size = None
if ptr is _POOL_UNSET:
self.buf = ir.alloc_buffer([0], "uint8", scope="shared.dyn")
self.buf = ir._alloc_buffer_deferred(self._allocation_shape, "uint8", "shared.dyn")
self.ptr = self.buf.data
self._owns_buffer = True
else:
Expand All @@ -420,6 +421,13 @@ def __init__(self, ptr=_POOL_UNSET):
self.offset = 0
self.max_offset = 0

def _allocation_shape(self):
if self._committed_size is None:
raise ValueError("SMEMPool.commit() must be called before leaving its scope")
from tvm.tirx import IntImm

return [IntImm("int64", self._committed_size)]

def alloc(
self,
shape,
Expand All @@ -429,6 +437,8 @@ def alloc(
align=0,
layout="default",
):
if self._owns_buffer and self._committed_size is not None:
raise ValueError("Cannot allocate from SMEMPool after commit()")
ir = _get_ir()
if align > 0:
self.offset = (self.offset + align - 1) // align * align
Expand Down Expand Up @@ -472,12 +482,14 @@ def alloc_tcgen05_mma_AB(self, shape, dtype="float16", swizzle_mode="auto", alig
return self.alloc(shape, dtype, align=align, layout=layout)

def move_base_to(self, offset):
if self._owns_buffer and self._committed_size is not None:
raise ValueError("Cannot move SMEMPool base after commit()")
self.offset = offset
if self._owns_buffer:
self.max_offset = max(self.max_offset, self.offset)

def commit(self, size=None):
"""Emit pool size annotation into the IR.
"""Finalize the backing allocation's byte extent.

Must be called after all ``alloc()`` / ``move_base_to()`` calls.

Expand All @@ -489,16 +501,11 @@ def commit(self, size=None):
"""
if not self._owns_buffer:
return
if self._committed_size is not None:
raise ValueError("SMEMPool.commit() can only be called once")
resolved = size if size is not None else self.max_offset
assert resolved >= self.max_offset, (
f"Specified smem size ({resolved}) is smaller than "
f"the pool high-water mark ({self.max_offset})"
)
import tvm.tirx
from tvm.tirx.script.ir_builder.parser_protocol import add_to_parent

add_to_parent(
tvm.tirx.AttrStmt(
0, "tirx.dyn_smem_bytes", tvm.tirx.IntImm("int64", resolved), tvm.tirx.Evaluate(0)
)
)
self._committed_size = resolved
4 changes: 0 additions & 4 deletions python/tvm/tirx/script/ir_builder/frame.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,3 @@ class LaunchThreadFrame(TIRFrame):
def __enter__(self) -> Var:
super().__enter__()
return self.iter_var.var


@_register_object("script.ir_builder.tirx.HintFrame")
class HintFrame(TIRFrame): ...
7 changes: 7 additions & 0 deletions python/tvm/tirx/script/ir_builder/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -459,6 +459,13 @@ def thread_id_in_wg(
return tuple(ret)


def _alloc_buffer_deferred(shape, dtype, scope):
"""Allocate a buffer whose shape is resolved when its builder scope closes."""
buf = _ffi_api.AllocBufferDeferred(shape, _normalize_prim_type(dtype), scope)
_record_meta_resource(buf, skip_frames=2)
return buf


@_register_mutable_decl("tirx.alloc_buffer")
def alloc_buffer(
shape: list[Expr] | tuple[Expr] | Expr | Integral,
Expand Down
19 changes: 0 additions & 19 deletions python/tvm/tirx/script/ir_builder/parser_protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -503,24 +503,6 @@ def attr(
return _ffi_api.Attr(node_or_dict, attr_key, value) # type: ignore[attr-defined] # pylint: disable=no-member


def hint(message: str = "", **attrs) -> frame.HintFrame:
"""Universal directive primitive for the sketch language.

Parameters
----------
message : str
Free-form directive string that the agent interprets.
**attrs
Optional structured key-value attributes for known patterns.

Returns
-------
res : frame.HintFrame
Usable as context manager (with T.hint("msg"):) or bare statement (T.hint("msg")).
"""
return _ffi_api.Hint(message, attrs or {}) # type: ignore[attr-defined] # pylint: disable=no-member


def buffer_store(
buffer: Buffer, # pylint: disable=redefined-outer-name
value: Expr,
Expand Down Expand Up @@ -1079,7 +1061,6 @@ def env_thread(thread_tag: str, dtype: str = "int32") -> Var:
"ge_",
"grid",
"gt_",
"hint",
"if_",
"if_then_else_",
"launch_thread",
Expand Down
14 changes: 0 additions & 14 deletions src/backend/cuda/codegen/codegen_cuda.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1619,20 +1619,6 @@ void CodeGenCUDA::Dispatch_(const AttrStmtNode* op) {
TVM_FFI_ICHECK(inner);
this->Dispatch(inner->body);
return;
} else if (op->attr_key == "disable_unroll") {
PrintIndent();
stream << "#pragma unroll 1\n";
this->Dispatch(op->body);
return;
} else if (op->attr_key == "pragma_unroll") {
PrintIndent();
stream << "#pragma unroll";
if (const auto* count = op->value.as<IntImmNode>(); count && count->value != 1) {
stream << " " << count->value;
}
stream << "\n";
this->Dispatch(op->body);
return;
} else if (op->attr_key == tirx::attr::thread_extent) {
}
CodeGenC::Dispatch_(op);
Expand Down
3 changes: 0 additions & 3 deletions src/s_tir/transform/merge_shared_memory_allocations.cc
Original file line number Diff line number Diff line change
Expand Up @@ -295,8 +295,6 @@ class SharedMemLinearAccessPatternFinder final : public StmtExprVisitor {
in_thread_env_ = true;
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op));
in_thread_env_ = false;
} else if (op->attr_key == tirx::attr::extern_scope) {
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op));
} else if (op->attr_key == s_tir::attr::virtual_thread) {
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op));
} else {
Expand Down Expand Up @@ -461,7 +459,6 @@ class SharedMemoryRewriter : public StmtExprMutator {
}

in_thread_env_ = false;

// 6. If this scope has no shmem allocs, skip the wrapper.
if (scope.shmem_allocs.empty()) {
scope_stack_.pop_back();
Expand Down
16 changes: 0 additions & 16 deletions src/tirx/script/ir_builder/frame.cc
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
ThenFrameNode::RegisterReflection();
ElseFrameNode::RegisterReflection();
DeclBufferFrameNode::RegisterReflection();
HintFrameNode::RegisterReflection();
}

namespace {
Expand Down Expand Up @@ -291,21 +290,6 @@ void DeclBufferFrameNode::ExitWithScope() {
}
}

void HintFrameNode::ExitWithScope() {
TIRFrameNode::ExitWithScope();
// Always store attrs as a structured Map in the node field
ffi::Map<ffi::String, Any> full_attrs;
if (!message.empty()) {
full_attrs.Set("message", ffi::String(message));
}
for (const auto& [k, v] : attrs) {
full_attrs.Set(k, v);
}
AddToParent(
tvm::tirx::AttrStmt(full_attrs, "tirx_hint", IntImm::Int32(1), AsStmt(stmts), source_span),
source_span);
}

} // namespace tirx
} // namespace ir_builder
} // namespace script
Expand Down
33 changes: 25 additions & 8 deletions src/tirx/script/ir_builder/ir.cc
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
#include <tvm/ffi/cast.h>
#include <tvm/ffi/container/array.h>
#include <tvm/ffi/container/variant.h>
#include <tvm/ffi/extra/structural_mutate.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/op.h>
#include <tvm/ir/prim/builtin.h>
Expand Down Expand Up @@ -513,13 +514,6 @@ ElseFrame Else() {
return ElseFrame(n);
}

HintFrame Hint(ffi::String message, ffi::Map<ffi::String, ffi::Any> attrs) {
ffi::ObjectPtr<HintFrameNode> n = ffi::make_object<HintFrameNode>();
n->message = message;
n->attrs = attrs;
return HintFrame(n);
}

Var EnvThread(ffi::String thread_tag, PrimType dtype) {
IterVar iter_var(Range{nullptr}, tvm::PrimVar("", dtype), tvm::tirx::IterVarType::kThreadIndex,
thread_tag);
Expand Down Expand Up @@ -660,6 +654,29 @@ BufferVar AllocBuffer(ffi::Array<PrimExpr> shape, PrimType dtype, ffi::String st
return buffer;
}

// Resolve a builder-only allocation shape after its enclosing scope has been
// emitted. Replacing the buffer definition and every use together preserves the
// allocation's original position and the identity shared by its views.
BufferVar AllocBufferDeferred(ffi::TypedFunction<ffi::Array<PrimExpr>()> shape, PrimType dtype,
ffi::String storage_scope) {
BufferVar buffer = AllocBuffer({IntImm::Int64(0)}, dtype, storage_scope, std::nullopt);
auto* frame = IRBuilder::Current()->frames.back().as_or_throw<TIRFrame>().get();
frame->callbacks.push_back([frame, buffer, shape, dtype, storage_scope]() {
Var resolved =
buffer.CopyWithType(BufferDecl(shape(), dtype, buffer.name(), std::nullopt, std::nullopt,
std::nullopt, storage_scope, 0, 0, std::nullopt, {})
.type());
auto replace = [&buffer,
&resolved](const Var& var) -> ffi::Expected<ffi::UnchangedOr<ffi::Any>> {
if (var.same_as(buffer)) return ffi::Any(resolved);
return ffi::Unchanged();
};
frame->stmts = ffi::StructuralMap<ffi::WalkOrder::kPreOrder>(frame->stmts, replace)
.as_or_throw<ffi::Array<tvm::tirx::Stmt>>();
});
return buffer;
}

tvm::tirx::Stmt Evaluate(Expr value) {
tvm::tirx::Stmt stmt = tvm::tirx::Evaluate(value);
AddToParent(stmt);
Expand Down Expand Up @@ -751,6 +768,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
ffi::String cur,
PrimType dtype) { return ScopeId(extents, parent, name, cur, dtype); })
.def("script.ir_builder.tirx.AllocBuffer", AllocBuffer)
.def("script.ir_builder.tirx.AllocBufferDeferred", AllocBufferDeferred)
.def("script.ir_builder.tirx.Serial", Serial)
.def("script.ir_builder.tirx.Parallel", Parallel)
.def("script.ir_builder.tirx.Vectorized", Vectorized)
Expand Down Expand Up @@ -782,7 +800,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
}
})
.def("script.ir_builder.tirx.EnvThread", EnvThread)
.def("script.ir_builder.tirx.Hint", Hint)
.def("script.ir_builder.tirx.BufferStore", BufferStore)
.def("script.ir_builder.tirx.Evaluate", Evaluate)
.def("script.ir_builder.tirx.Ptr", Ptr);
Expand Down
17 changes: 0 additions & 17 deletions src/tirx/script/printer/stmt.cc
Original file line number Diff line number Diff line change
Expand Up @@ -786,23 +786,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
rhs = DocsifyLaunchThread(stmt, stmt_p, &define_var, d);
}
}
if (stmt->attr_key == "tirx_hint") {
if (auto map_node = stmt->node.as<ffi::Map<ffi::String, ffi::Any>>()) {
ffi::Array<ExprDoc> args;
ffi::Array<ffi::String> kwargs_keys;
ffi::Array<ExprDoc> kwargs_values;
for (const auto& [k, v] : map_node.value()) {
if (k == "message") {
auto s = v.as<ffi::String>().value();
args.push_back(LiteralDoc::Str(s, stmt_p->Attr("node")));
} else {
kwargs_keys.push_back(k);
kwargs_values.push_back(d->AsDoc<ExprDoc>(v, stmt_p->Attr("node")));
}
}
rhs = TIR(d, "hint")->Call(args, kwargs_keys, kwargs_values);
}
}
if (!rhs.has_value()) {
// Try to collapse consecutive dict-attr-pattern AttrStmts into T.attr({...})
if (IsDictAttrPattern(stmt)) {
Expand Down
4 changes: 1 addition & 3 deletions src/tirx/transform/remove_no_op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -92,9 +92,7 @@ class NoOpRemover : public IRMutatorWithAnalyzer {

private:
UnchangedOr<Stmt> Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final {
if (op->attr_key == "pragma_debug_skip_region") {
return MakeEvaluate(IntImm::Int32(0));
} else if (op->attr_key == tvm::tirx::attr::async_wait_queue_scope) {
if (op->attr_key == tvm::tirx::attr::async_wait_queue_scope) {
auto wait_attrs = GetAsyncWaitAttributes(op);
auto wait_cnt = wait_attrs.second;
sym::Analyzer ana;
Expand Down
Loading
Loading