Skip to content

Commit 7d36570

Browse files
authored
[REFACTOR][TIRx] Represent buffer definitions as operation bindings (#20482)
Represent buffer allocation and declaration as ordinary bindings of buffer-returning operation calls. Their shape, data type, and storage scope are explicit operands; declaration also carries its existing pointer operand. Allocation annotations live in call attributes. Update traversal, lowering, code generation, and script construction and printing to consume the shared representation.
1 parent 2a11297 commit 7d36570

146 files changed

Lines changed: 3139 additions & 2829 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎include/tvm/relax/distributed/axis_group_graph.h‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,7 @@ class BufferAxisGraphExtractor : public s_tir::StmtExprVisitor {
7474
ffi::Map<BufferVar, Var> inverse_buffer_map;
7575
for (const Var& param : prim_func->params) {
7676
if (param->ty.as<BufferTypeNode>()) {
77-
inverse_buffer_map.Set(BufferVar(param), param);
77+
inverse_buffer_map.Set(param.as_or_throw<BufferVar>(), param);
7878
}
7979
}
8080
std::vector<std::vector<TIRVarAxis>> tir_var_axis_group_list;
@@ -83,7 +83,7 @@ class BufferAxisGraphExtractor : public s_tir::StmtExprVisitor {
8383
if (!param->ty.as<BufferTypeNode>()) {
8484
continue;
8585
}
86-
BufferVar buffer(param);
86+
BufferVar buffer = param.as_or_throw<BufferVar>();
8787
for (int i = 0; i < static_cast<int>(buffer->shape.size()); i++) {
8888
if (extractor->buffer_axis_graph_.count({buffer, i})) {
8989
std::vector<BufferAxis> buffer_axis_group;

‎include/tvm/tirx/builtin.h‎

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,53 @@ namespace tirx {
4141

4242
/*! \brief Collection of builtin intrinsics as ops */
4343
namespace builtin {
44+
/*!
45+
* \brief Allocate a buffer: alloc_buffer(shape, dtype, scope) -> BufferType.
46+
*
47+
* Arguments, in order:
48+
* - args[0]: shape, Tuple of integer extents (IntImm or symbolic integer expressions).
49+
* - args[1]: dtype, DataTypeImm with a DLDataType payload for the element type.
50+
* - args[2]: scope, StringImm naming the storage scope.
51+
*
52+
* DictAttrs directly holds the allocation annotations, defaulting to an empty dictionary.
53+
* The BufferType result agrees with the operands and retains buffer access/storage metadata.
54+
*
55+
* \code
56+
* // Example pattern match code for a given Binding:
57+
* if (const auto* call = binding->value.as<CallNode>();
58+
* call && call->op.same_as(builtin::alloc_buffer())) {
59+
* tvm::Tuple shape = call->args[0].as_or_throw<tvm::Tuple>();
60+
* DLDataType dtype = call->args[1].as_or_throw<DataTypeImm>()->value;
61+
* ffi::String scope = call->args[2].as_or_throw<StringImm>()->value;
62+
* DictAttrs annotations = call->attrs.as_or_throw<DictAttrs>();
63+
* }
64+
* \endcode
65+
*/
66+
TVM_DLL const Op& alloc_buffer();
67+
/*!
68+
* \brief Declare a buffer view: decl_buffer(data, shape, dtype, scope) -> BufferType.
69+
*
70+
* Arguments, in order:
71+
* - args[0]: data, Expr for the existing physical pointer backing the buffer view.
72+
* - args[1]: shape, Tuple of integer extents (IntImm or symbolic integer expressions).
73+
* - args[2]: dtype, DataTypeImm with a DLDataType payload for the element type.
74+
* - args[3]: scope, StringImm naming the storage scope.
75+
*
76+
* There are no attributes. The BufferType result agrees with the operands and retains
77+
* buffer access/storage metadata. The operation binds a view without allocating memory.
78+
*
79+
* \code
80+
* // Example pattern match code for a given Binding:
81+
* if (const auto* call = binding->value.as<CallNode>();
82+
* call && call->op.same_as(builtin::decl_buffer())) {
83+
* Expr data = call->args[0];
84+
* tvm::Tuple shape = call->args[1].as_or_throw<tvm::Tuple>();
85+
* DLDataType dtype = call->args[2].as_or_throw<DataTypeImm>()->value;
86+
* ffi::String scope = call->args[3].as_or_throw<StringImm>()->value;
87+
* }
88+
* \endcode
89+
*/
90+
TVM_DLL const Op& decl_buffer();
4491
/*!
4592
* \brief Return from a GPU thread without returning a function value.
4693
*/

‎include/tvm/tirx/expr.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,7 +117,7 @@ class BufferVar : public Var {
117117
*
118118
* If flattening changes the type, the result is a fresh BufferVar. Callers
119119
* that use it as a view over this buffer must bind the returned variable with
120-
* `DeclBuffer(flattened, this->data())`.
120+
* a `Bind` of `flattened` to a `decl_buffer` Call over `this->data()`.
121121
*/
122122
BufferVar GetFlattenedBuffer() const;
123123

‎include/tvm/tirx/script/ir_builder/ir.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -325,7 +325,7 @@ DeclBufferFrame DeclBuffer(ffi::Array<PrimExpr> shape, PrimType dtype, ffi::Stri
325325
ffi::Optional<PrimExpr> allocated_addr = std::nullopt);
326326

327327
/*!
328-
* \brief Statement-level buffer allocation (creates an AllocBuffer IR node).
328+
* \brief Statement-level buffer allocation (binds a buffer-returning allocation Call).
329329
* \param shape The shape of the buffer to allocate.
330330
* \param dtype The data type of buffer elements.
331331
* \param storage_scope The storage scope (e.g., "global", "shared").

‎include/tvm/tirx/stmt.h‎

Lines changed: 0 additions & 76 deletions
Original file line numberDiff line numberDiff line change
@@ -231,82 +231,6 @@ class BufferStore : public Stmt {
231231
TVM_DEFINE_OBJECT_REF_COW_METHOD(BufferStoreNode);
232232
};
233233

234-
/*! \brief Declare a buffer that can be used in the body */
235-
class DeclBufferNode : public StmtNode {
236-
public:
237-
/*! \brief The buffer being declared */
238-
BufferVar buffer;
239-
/*! \brief Physical pointer expression backing the declaration. */
240-
Expr data;
241-
242-
static void RegisterReflection() {
243-
namespace refl = tvm::ffi::reflection;
244-
refl::ObjectDef<DeclBufferNode>()
245-
.def_ro("buffer", &DeclBufferNode::buffer, refl::AttachFieldFlag::SEqHashDefSimple())
246-
.def_ro("data", &DeclBufferNode::data);
247-
}
248-
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.DeclBuffer", DeclBufferNode, StmtNode);
249-
};
250-
251-
/*! \brief Managed reference to DeclBufferNode */
252-
class DeclBuffer : public Stmt {
253-
public:
254-
TVM_DLL DeclBuffer(BufferVar buffer, Expr data, Span span = Span());
255-
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(DeclBuffer, Stmt, DeclBufferNode);
256-
TVM_DEFINE_OBJECT_REF_COW_METHOD(DeclBufferNode);
257-
};
258-
259-
/*! \brief Allocate a buffer and declare it in scope */
260-
class AllocBufferNode : public StmtNode {
261-
public:
262-
/*! \brief The buffer being allocated and declared */
263-
BufferVar buffer;
264-
/*!
265-
* \brief Additional annotations about the allocation.
266-
*
267-
* These annotations can be used as auxiliary hint
268-
* to future transformations.
269-
*/
270-
ffi::Map<ffi::String, ffi::Any> annotations;
271-
272-
static void RegisterReflection() {
273-
namespace refl = tvm::ffi::reflection;
274-
refl::ObjectDef<AllocBufferNode>()
275-
.def_ro("buffer", &AllocBufferNode::buffer, refl::AttachFieldFlag::SEqHashDefSimple())
276-
.def_ro("annotations", &AllocBufferNode::annotations);
277-
}
278-
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.AllocBuffer", AllocBufferNode, StmtNode);
279-
};
280-
281-
/*! \brief Managed reference to AllocBufferNode */
282-
class AllocBuffer : public Stmt {
283-
public:
284-
TVM_DLL AllocBuffer(
285-
BufferVar buffer,
286-
ffi::Map<ffi::String, ffi::Any> annotations = ffi::Map<ffi::String, ffi::Any>(),
287-
Span span = Span());
288-
/*!
289-
* \brief If the buffer's shape is constant, return the total number of elements.
290-
* \return The product of all shape extents if all are constant, std::nullopt otherwise.
291-
*/
292-
std::optional<int64_t> ConstantAllocationSize() const {
293-
int64_t result = 1;
294-
for (const PrimExpr& extent : (*this)->buffer->shape) {
295-
if (const auto* int_size = extent.as<IntImmNode>()) {
296-
auto product = (result * int_size->value).as<int64_t>();
297-
if (!product.has_value()) return std::nullopt;
298-
result = *product;
299-
} else {
300-
return std::nullopt;
301-
}
302-
}
303-
return result;
304-
}
305-
306-
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(AllocBuffer, Stmt, AllocBufferNode);
307-
TVM_DEFINE_OBJECT_REF_COW_METHOD(AllocBufferNode);
308-
};
309-
310234
/*!
311235
* \brief The container of seq statement.
312236
* Represent a sequence of statements.

‎include/tvm/tirx/stmt_functor.h‎

Lines changed: 0 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -98,12 +98,6 @@ class StmtFunctor<R(const Stmt&, Args...)> {
9898
virtual R Dispatch_(const ContinueNode* node, Args... args) {
9999
return DispatchDefault_(node, std::forward<Args>(args)...);
100100
}
101-
virtual R Dispatch_(const AllocBufferNode* node, Args... args) {
102-
return DispatchDefault_(node, std::forward<Args>(args)...);
103-
}
104-
virtual R Dispatch_(const DeclBufferNode* node, Args... args) {
105-
return DispatchDefault_(node, std::forward<Args>(args)...);
106-
}
107101
virtual R Dispatch_(const BufferStoreNode* node, Args... args) {
108102
return DispatchDefault_(node, std::forward<Args>(args)...);
109103
}
@@ -147,8 +141,6 @@ class StmtFunctor<R(const Stmt&, Args...)> {
147141
SetDispatch<TSelf, ReturnNode>(vtable);
148142
SetDispatch<TSelf, BreakNode>(vtable);
149143
SetDispatch<TSelf, ContinueNode>(vtable);
150-
SetDispatch<TSelf, AllocBufferNode>(vtable);
151-
SetDispatch<TSelf, DeclBufferNode>(vtable);
152144
SetDispatch<TSelf, BufferStoreNode>(vtable);
153145
SetDispatch<TSelf, AssertStmtNode>(vtable);
154146
SetDispatch<TSelf, SeqStmtNode>(vtable);
@@ -206,8 +198,6 @@ class TVM_DLL StmtExprVisitor : public tvm::ExprVisitor {
206198
virtual ffi::Optional<VisitInterrupt> Visit_(const ReturnNode* op);
207199
virtual ffi::Optional<VisitInterrupt> Visit_(const BreakNode* op);
208200
virtual ffi::Optional<VisitInterrupt> Visit_(const ContinueNode* op);
209-
virtual ffi::Optional<VisitInterrupt> Visit_(const AllocBufferNode* op);
210-
virtual ffi::Optional<VisitInterrupt> Visit_(const DeclBufferNode* op);
211201
virtual ffi::Optional<VisitInterrupt> Visit_(const BufferStoreNode* op);
212202
virtual ffi::Optional<VisitInterrupt> Visit_(const AssertStmtNode* op);
213203
virtual ffi::Optional<VisitInterrupt> Visit_(const SeqStmtNode* op);
@@ -227,9 +217,6 @@ class TVM_DLL StmtExprVisitor : public tvm::ExprVisitor {
227217
ffi::Optional<VisitInterrupt> Visit_(const prim::BroadcastNode* op) override;
228218
ffi::Optional<VisitInterrupt> Visit_(const prim::ShuffleNode* op) override;
229219

230-
/*! \brief Visit definition metadata as uses, separately from the buffer Var definition. */
231-
ffi::Optional<VisitInterrupt> VisitBufferMetadata(const BufferVar& buffer);
232-
233220
protected:
234221
explicit StmtExprVisitor(const VTable* vtable) : tvm::ExprVisitor(vtable) {}
235222
static void InitVTable(VTable* vtable);
@@ -270,8 +257,6 @@ class TVM_DLL StmtExprMutator : public tvm::ExprMutator {
270257
virtual UnchangedOr<Stmt> Mutate_(const ReturnNode* op, InplaceMode inplace_mode);
271258
virtual UnchangedOr<Stmt> Mutate_(const BreakNode* op, InplaceMode inplace_mode);
272259
virtual UnchangedOr<Stmt> Mutate_(const ContinueNode* op, InplaceMode inplace_mode);
273-
virtual UnchangedOr<Stmt> Mutate_(const AllocBufferNode* op, InplaceMode inplace_mode);
274-
virtual UnchangedOr<Stmt> Mutate_(const DeclBufferNode* op, InplaceMode inplace_mode);
275260
virtual UnchangedOr<Stmt> Mutate_(const BufferStoreNode* op, InplaceMode inplace_mode);
276261
virtual UnchangedOr<Stmt> Mutate_(const AssertStmtNode* op, InplaceMode inplace_mode);
277262
virtual UnchangedOr<Stmt> Mutate_(const SeqStmtNode* op, InplaceMode inplace_mode);

‎include/tvm/tirx/type.h‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,8 @@
2828
#include <tvm/ir/type.h>
2929
#include <tvm/tirx/layout.h>
3030

31+
#include <optional>
32+
3133
namespace tvm::tirx {
3234

3335
#ifndef TVM_INDEX_DEFAULT_I64
@@ -127,6 +129,12 @@ class BufferTypeNode : public TypeNode {
127129
/*! \return type of the physical pointer projected by buffer_data. */
128130
PointerType DataPointerType() const { return PointerType(dtype, storage_scope); }
129131

132+
/*! \brief Whether this type supports scalar buffer syntax. */
133+
TVM_DLL bool IsScalar(bool alloc_or_decl = true) const;
134+
135+
/*! \brief Return the constant element count, or nullopt for symbolic extents or overflow. */
136+
TVM_DLL std::optional<int64_t> ConstantAllocationSize() const;
137+
130138
/*! \brief Determine the offset in the buffer of the given index.
131139
*
132140
* Returns the buffer offset, in number of elements of type dtype,

‎python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py‎

Lines changed: 20 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -109,14 +109,15 @@
109109
import operator
110110

111111
import tvm
112+
from tvm.ir import Call, DataTypeImm, DictAttrs, StringImm, Tuple
112113
from tvm.runtime import DataType
113114
from tvm.script import tirx as T
114115
from tvm.sym import Analyzer
115116
from tvm.tirx import Buffer, PrimFunc
116117
from tvm.tirx.layout import ComposeLayout, TCol, TileLayout, TLane
117118
from tvm.tirx.layout import m as m_axis
118119
from tvm.tirx.operator.tile_primitive import DispatchContext, predicate, register_dispatch
119-
from tvm.tirx.stmt import AllocBuffer, Evaluate, SeqStmt
120+
from tvm.tirx.stmt import Bind, Evaluate, SeqStmt
120121
from tvm.tirx.tile_primitive import TilePrimitiveCall
121122

122123
from ..copy import _single_thread_exec
@@ -679,7 +680,24 @@ def _get_or_create_desc(sctx, s_buf, ldo, sdo, swizzle):
679680
encode_call = T.cuda.tcgen05.encode_matrix_descriptor(
680681
desc_buf.data, T.reinterpret("handle", T.uint64(0)), ldo, sdo, swizzle
681682
)
682-
wrap = SeqStmt([AllocBuffer(desc_buf), Evaluate(encode_call)])
683+
wrap = SeqStmt(
684+
[
685+
Bind(
686+
desc_buf,
687+
Call(
688+
"tirx.alloc_buffer",
689+
[
690+
Tuple(desc_buf.ty.shape),
691+
DataTypeImm(desc_buf.ty.dtype.dtype),
692+
StringImm(desc_buf.scope()),
693+
],
694+
attrs=DictAttrs({}),
695+
ret_ty=desc_buf.ty,
696+
),
697+
),
698+
Evaluate(encode_call),
699+
]
700+
)
683701
sctx.add_post_buffer_def_stmt(s_buf, wrap)
684702
sctx.cache_set(cache_key, desc_buf)
685703
return desc_buf

‎python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py‎

Lines changed: 44 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@
2626
import operator
2727

2828
import tvm
29-
from tvm.ir import TensorRegion
29+
from tvm.ir import Call, DataTypeImm, DictAttrs, StringImm, TensorRegion, Tuple
3030
from tvm.runtime import DataType
3131
from tvm.script import tirx as T
3232
from tvm.sym.analyzer import Analyzer
@@ -44,7 +44,7 @@
4444
tmem_mma_operand_layout,
4545
)
4646
from tvm.tirx.operator.tile_primitive import DispatchContext, predicate, register_dispatch
47-
from tvm.tirx.stmt import AllocBuffer, Evaluate, SeqStmt
47+
from tvm.tirx.stmt import Bind, Evaluate, SeqStmt
4848
from tvm.tirx.tile_primitive import TilePrimitiveCall
4949

5050
from ...cpp.descriptors import (
@@ -1089,8 +1089,32 @@ def _make_lo_uniform(desc_buf):
10891089
pack = T.ptx.mov.b64(desc_buf[0], desc_lo[0], desc_hi[0])
10901090
return SeqStmt(
10911091
[
1092-
AllocBuffer(desc_lo),
1093-
AllocBuffer(desc_hi),
1092+
Bind(
1093+
desc_lo,
1094+
Call(
1095+
"tirx.alloc_buffer",
1096+
[
1097+
Tuple(desc_lo.ty.shape),
1098+
DataTypeImm(desc_lo.ty.dtype.dtype),
1099+
StringImm(desc_lo.scope()),
1100+
],
1101+
attrs=DictAttrs({}),
1102+
ret_ty=desc_lo.ty,
1103+
),
1104+
),
1105+
Bind(
1106+
desc_hi,
1107+
Call(
1108+
"tirx.alloc_buffer",
1109+
[
1110+
Tuple(desc_hi.ty.shape),
1111+
DataTypeImm(desc_hi.ty.dtype.dtype),
1112+
StringImm(desc_hi.scope()),
1113+
],
1114+
attrs=DictAttrs({}),
1115+
ret_ty=desc_hi.ty,
1116+
),
1117+
),
10941118
Evaluate(unpack),
10951119
Evaluate(shuffle),
10961120
Evaluate(pack),
@@ -1112,7 +1136,22 @@ def _make_desc(smem_buf, ldo, sdo, swizzle_val, name):
11121136
sdo,
11131137
swizzle_val,
11141138
)
1115-
wrap_stmts = [AllocBuffer(desc_buf), Evaluate(encode_call)]
1139+
wrap_stmts = [
1140+
Bind(
1141+
desc_buf,
1142+
Call(
1143+
"tirx.alloc_buffer",
1144+
[
1145+
Tuple(desc_buf.ty.shape),
1146+
DataTypeImm(desc_buf.ty.dtype.dtype),
1147+
StringImm(desc_buf.scope()),
1148+
],
1149+
attrs=DictAttrs({}),
1150+
ret_ty=desc_buf.ty,
1151+
),
1152+
),
1153+
Evaluate(encode_call),
1154+
]
11161155
if warp_scope:
11171156
wrap_stmts.append(_make_lo_uniform(desc_buf))
11181157
wrap_stmts.append(_krp)

0 commit comments

Comments
 (0)