From ad9b6c51bb8a243c775ce6386c447b328a3d7b19 Mon Sep 17 00:00:00 2001 From: tqchen Date: Tue, 29 Sep 2026 05:35:59 +0000 Subject: [PATCH 1/2] [REFACTOR][TIR] Remove obsolete AttrStmt protocols Derive dynamic shared-memory launch bytes from ordinary allocation extents and finalize pool shapes during IR construction. Remove unused hint and external-scope interfaces and obsolete unroll and debug AttrStmt consumers while retaining loop annotations. --- include/tvm/tirx/script/ir_builder/frame.h | 39 ---- include/tvm/tirx/stmt.h | 6 - python/tvm/backend/cuda/lang/alloc_pool.py | 33 +-- python/tvm/tirx/script/ir_builder/frame.py | 4 - python/tvm/tirx/script/ir_builder/ir.py | 7 + .../tirx/script/ir_builder/parser_protocol.py | 19 -- src/backend/cuda/codegen/codegen_cuda.cc | 14 -- .../merge_shared_memory_allocations.cc | 3 - src/tirx/script/ir_builder/frame.cc | 16 -- src/tirx/script/ir_builder/ir.cc | 33 ++- src/tirx/script/printer/stmt.cc | 17 -- src/tirx/transform/remove_no_op.cc | 4 +- src/tirx/transform/split_host_device.cc | 57 +---- src/tirx/transform/storage_rewrite.cc | 13 -- .../cuda/copy_async/test_tma.py | 4 - tests/python/tirx/test_alloc_pool.py | 71 +++++++ tests/python/tirx/test_hint.py | 199 +----------------- 17 files changed, 131 insertions(+), 408 deletions(-) diff --git a/include/tvm/tirx/script/ir_builder/frame.h b/include/tvm/tirx/script/ir_builder/frame.h index 295a59870301..97f47f2aa53c 100644 --- a/include/tvm/tirx/script/ir_builder/frame.h +++ b/include/tvm/tirx/script/ir_builder/frame.h @@ -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 attrs; - - static void RegisterReflection() { - namespace refl = tvm::ffi::reflection; - refl::ObjectDef() - .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 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 diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h index 212ad12b980a..7abfb6621040 100644 --- a/include/tvm/tirx/stmt.h +++ b/include/tvm/tirx/stmt.h @@ -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 */ diff --git a/python/tvm/backend/cuda/lang/alloc_pool.py b/python/tvm/backend/cuda/lang/alloc_pool.py index f3d7eba2659f..ee70456f46d4 100644 --- a/python/tvm/backend/cuda/lang/alloc_pool.py +++ b/python/tvm/backend/cuda/lang/alloc_pool.py @@ -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: @@ -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, @@ -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 @@ -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. @@ -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 diff --git a/python/tvm/tirx/script/ir_builder/frame.py b/python/tvm/tirx/script/ir_builder/frame.py index 2aaedb4e71ff..c46a97af0bb7 100644 --- a/python/tvm/tirx/script/ir_builder/frame.py +++ b/python/tvm/tirx/script/ir_builder/frame.py @@ -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): ... diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index 58134a768611..bc833bdd76ca 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -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, diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 90fa9184c729..268dabd861ff 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -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, @@ -1079,7 +1061,6 @@ def env_thread(thread_tag: str, dtype: str = "int32") -> Var: "ge_", "grid", "gt_", - "hint", "if_", "if_then_else_", "launch_thread", diff --git a/src/backend/cuda/codegen/codegen_cuda.cc b/src/backend/cuda/codegen/codegen_cuda.cc index 7cba4d4c1dfa..e71b272d71fa 100644 --- a/src/backend/cuda/codegen/codegen_cuda.cc +++ b/src/backend/cuda/codegen/codegen_cuda.cc @@ -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(); 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); diff --git a/src/s_tir/transform/merge_shared_memory_allocations.cc b/src/s_tir/transform/merge_shared_memory_allocations.cc index e9b2db0c0b3d..fd843692393d 100644 --- a/src/s_tir/transform/merge_shared_memory_allocations.cc +++ b/src/s_tir/transform/merge_shared_memory_allocations.cc @@ -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 { @@ -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(); diff --git a/src/tirx/script/ir_builder/frame.cc b/src/tirx/script/ir_builder/frame.cc index cc10fb0fea93..760fcfd539d1 100644 --- a/src/tirx/script/ir_builder/frame.cc +++ b/src/tirx/script/ir_builder/frame.cc @@ -49,7 +49,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { ThenFrameNode::RegisterReflection(); ElseFrameNode::RegisterReflection(); DeclBufferFrameNode::RegisterReflection(); - HintFrameNode::RegisterReflection(); } namespace { @@ -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 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 diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index 675bf816d1d2..5937b01889d6 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -19,6 +19,7 @@ #include #include #include +#include #include #include #include @@ -513,13 +514,6 @@ ElseFrame Else() { return ElseFrame(n); } -HintFrame Hint(ffi::String message, ffi::Map attrs) { - ffi::ObjectPtr n = ffi::make_object(); - 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); @@ -660,6 +654,29 @@ BufferVar AllocBuffer(ffi::Array 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()> 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().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> { + if (var.same_as(buffer)) return ffi::Any(resolved); + return ffi::Unchanged(); + }; + frame->stmts = ffi::StructuralMap(frame->stmts, replace) + .as_or_throw>(); + }); + return buffer; +} + tvm::tirx::Stmt Evaluate(Expr value) { tvm::tirx::Stmt stmt = tvm::tirx::Evaluate(value); AddToParent(stmt); @@ -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) @@ -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); diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc index 1c2a14b81e92..070cdca0c495 100644 --- a/src/tirx/script/printer/stmt.cc +++ b/src/tirx/script/printer/stmt.cc @@ -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::Array args; - ffi::Array kwargs_keys; - ffi::Array kwargs_values; - for (const auto& [k, v] : map_node.value()) { - if (k == "message") { - auto s = v.as().value(); - args.push_back(LiteralDoc::Str(s, stmt_p->Attr("node"))); - } else { - kwargs_keys.push_back(k); - kwargs_values.push_back(d->AsDoc(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)) { diff --git a/src/tirx/transform/remove_no_op.cc b/src/tirx/transform/remove_no_op.cc index 99e5c1d928e1..b01d9894916d 100644 --- a/src/tirx/transform/remove_no_op.cc +++ b/src/tirx/transform/remove_no_op.cc @@ -92,9 +92,7 @@ class NoOpRemover : public IRMutatorWithAnalyzer { private: UnchangedOr 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; diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index a2ab6f41a1ee..3bb27f15f112 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -429,20 +429,8 @@ class DeviceInfoCollector : public StmtExprVisitor { collector->info_.launch_params.push_back( tvm::runtime::launch_param::kUseRequiredBlockDimension); } - // The dynamic shared memory is required to be the last of the kernel - // launch parameters. An explicit tirx.dyn_smem_bytes declaration wins; - // otherwise fall back to the size inferred from the allocation extent. - // A zero-extent allocation is a pool-style extern placeholder, so having - // neither a declaration nor a usable extent is an authoring error. - if (!collector->dyn_shmem_size.has_value() && collector->inferred_shmem_size_.has_value()) { - const auto* inferred = collector->inferred_shmem_size_.value().as(); - TVM_FFI_ICHECK(!(inferred && inferred->value == 0)) - << "PrimFunc " << gvar->name_hint - << " allocates dynamic shared memory with a placeholder extent but does not declare " - "its size; annotate the kernel with tirx.dyn_smem_bytes (SMEMPool.commit() emits " - "it)."; - collector->dyn_shmem_size = collector->inferred_shmem_size_; - } + // Dynamic shared memory is the last kernel launch parameter. Its size is + // derived from the backing allocation, including pool allocations. if (collector->dyn_shmem_size) { collector->info_.launch_params.push_back( tvm::runtime::launch_param::kUseDynamicSharedMemoryTag); @@ -468,7 +456,7 @@ class DeviceInfoCollector : public StmtExprVisitor { if (launch_param == tvm::runtime::launch_param::kUseDynamicSharedMemoryTag) { TVM_FFI_ICHECK(dyn_shmem_size.has_value()) << "Compute kernel requires launch parameter \"" << launch_param - << "\", but PrimFunc did not declare tirx.dyn_smem_bytes."; + << "\", but PrimFunc has no dynamic shared memory allocation."; return dyn_shmem_size.value(); } @@ -502,15 +490,6 @@ class DeviceInfoCollector : public StmtExprVisitor { } ffi::Optional Visit_(const AttrStmtNode* op) final { - if (op->attr_key == "tirx.dyn_smem_bytes") { - // Kernel-level declaration of the dynamic shared memory launch size. - // The backing shared.dyn allocation is an extern placeholder; this - // attribute is the single source of truth for the launch parameter. - TVM_FFI_ICHECK(!dyn_shmem_size.has_value()) - << "Only one tirx.dyn_smem_bytes declaration is allowed per kernel."; - TVM_FFI_ICHECK(op->value.as()) << "tirx.dyn_smem_bytes must be an IntImm"; - dyn_shmem_size = op->value.as_or_throw(); - } if (op->attr_key == attr::thread_extent) { ffi::String thread_tag; if (auto iv = op->node.as()) { @@ -552,10 +531,6 @@ class DeviceInfoCollector : public StmtExprVisitor { << "Only one dynamic shared memory allocation is allowed."; saw_dyn_shared_alloc_ = true; - // Fallback launch size inferred from the allocation extent, used when - // no tirx.dyn_smem_bytes declaration is present (e.g. s_tir schedules - // allocate shared.dyn with a concrete extent). A zero extent is a - // pool-style extern placeholder and carries no size information. TVM_FFI_ICHECK_GT(op->buffer->shape.size(), 0); PrimExpr dyn_size = IntImm::Int32(1); for (const auto& extent : op->buffer->shape) { @@ -570,7 +545,7 @@ class DeviceInfoCollector : public StmtExprVisitor { dyn_size = ffi::StructuralMap(dyn_size, f_substitute) .as_or_throw(); } - inferred_shmem_size_ = dyn_size; + dyn_shmem_size = dyn_size; } return StmtExprVisitor::Visit_(op); } @@ -585,9 +560,6 @@ class DeviceInfoCollector : public StmtExprVisitor { ffi::Optional dyn_shmem_size{std::nullopt}; // Whether a shared.dyn allocation was seen. bool saw_dyn_shared_alloc_{false}; - // Launch size inferred from the allocation extent (fallback when no - // tirx.dyn_smem_bytes declaration is present). - ffi::Optional inferred_shmem_size_{std::nullopt}; // Flag-only launch attributes requested by the original PrimFunc. bool use_programmatic_dependent_launch_{false}; bool use_cooperative_launch_{false}; @@ -710,27 +682,6 @@ class DeviceKernelMutator : public StmtExprMutator { Target target = func->GetAttr(tvm::attr::kTarget).value(); bool preserve_early_returns = target->kind->name == "cuda"; write_ptr->body = ReturnRemover::Apply(write_ptr->body, !preserve_early_returns); - // The dyn-smem size declaration was consumed by DeviceInfoCollector; - // it has no meaning inside the kernel body. - class StripDynSmemAttr : public StmtExprMutator { - public: - using StmtExprMutator::Mutate; - using StmtExprMutator::Mutate_; - UnchangedOr Mutate(ffi::AnyView input, InplaceMode inplace_mode) override { - if (input.as()) return ffi::Unchanged(); - return StmtExprMutator::Mutate(input, inplace_mode); - } - - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == "tirx.dyn_smem_bytes") { - return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); - } - return StmtExprMutator::Mutate_(op, inplace_mode); - } - }; - write_ptr->body = ffi::make_object() - ->Mutate(write_ptr->body, InplaceMode::kAllow) - .ValueOrUnchanged(write_ptr->body); } func = WithAttrs(std::move(func), diff --git a/src/tirx/transform/storage_rewrite.cc b/src/tirx/transform/storage_rewrite.cc index d6a6d7bdd275..453261fa40c6 100644 --- a/src/tirx/transform/storage_rewrite.cc +++ b/src/tirx/transform/storage_rewrite.cc @@ -240,8 +240,6 @@ class LinearAccessPatternFinder 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 == attr::extern_scope) { - TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else if (op->attr_key == tvm::tirx::attr::virtual_thread) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(VisitNewScope(op)); } else { @@ -391,15 +389,6 @@ class InplaceOpVerifier : public StmtExprVisitor { return std::nullopt; } - ffi::Optional Visit_(const AttrStmtNode* op) final { - // always reject extern code - if (op->attr_key == attr::extern_scope) { - result_ = false; - return std::nullopt; - } - return StmtExprVisitor::Visit_(op); - } - ffi::Optional Visit_(const AllocBufferNode* op) final { // reject inplace for volatile buffers if (op->annotations.count(attr::kVolatile)) { @@ -1091,8 +1080,6 @@ class StoragePlanRewriter : public StmtExprMutator { if (op->attr_key == attr::thread_extent || op->attr_key == tvm::tirx::attr::virtual_thread || attr::IsPragmaKey(op->attr_key)) { PlanNewScope(op); - } else { - TVM_FFI_ICHECK(op->attr_key == attr::extern_scope); } } else if (s.stmt->IsInstance()) { const auto* op = static_cast(s.stmt); diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py index ab2e7b49f857..20e24d9e5019 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py @@ -1532,7 +1532,6 @@ def rank_change(A: T.Buffer((8, 8), 'float16')): T.cta_id([1]) tid = T.thread_id([1]) dyn = T.alloc_buffer((65,), "uint64", scope="shared.dyn") - T.attr({"tirx.dyn_smem_bytes": 65 * 8}) A_smem = T.decl_buffer((64,), "float16", dyn.data, layout=T.TileLayout(T.S[64])) mbar = T.decl_buffer((1,), "uint64", dyn.data, elem_offset=16) if tid == 0: @@ -1559,7 +1558,6 @@ def selector_gather( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((520,), "uint64", scope="shared.dyn") - T.attr({"tirx.dyn_smem_bytes": 520 * 8}) A_smem = T.decl_buffer( (4, 64), "bfloat16", dyn.data, layout=T.TileLayout(T.S[4, 64]) ) @@ -1730,7 +1728,6 @@ def kernel( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((shared_bytes + 8,), "uint8", scope="shared.dyn") - T.attr({"tirx.dyn_smem_bytes": shared_bytes + 8}) q_smem = T.decl_buffer( (64, 512), "bfloat16", dyn.data, scope="shared.dyn", layout=q_layout ) @@ -2004,7 +2001,6 @@ def kernel( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((shared_bytes + 64,), "uint8", scope="shared.dyn") - T.attr({"tirx.dyn_smem_bytes": shared_bytes + 64}) A_smem = T.decl_buffer( (4, cols), dtype, dyn.data, layout=T.TileLayout(T.S[4, cols]) ) diff --git a/tests/python/tirx/test_alloc_pool.py b/tests/python/tirx/test_alloc_pool.py index 9614531ec08a..6902a64fda6d 100644 --- a/tests/python/tirx/test_alloc_pool.py +++ b/tests/python/tirx/test_alloc_pool.py @@ -17,7 +17,10 @@ """Tests for CUDA allocation pool validation.""" import pytest +import tvm_ffi +import tvm +from tvm.script import tirx as T from tvm.tirx.cuda.lang.alloc_pool import _validate_mma_alloc_shape from tvm.tirx.cuda.tile_primitive.tma_utils import SwizzleMode @@ -113,5 +116,73 @@ def test_swizzle_none_skips_validation(self): _validate_mma_alloc_shape((128,), "bfloat16", SwizzleMode.SWIZZLE_NONE) +@pytest.mark.parametrize("size", [None, 128]) +def test_smem_pool_commits_allocation_extent(size): + @T.prim_func + def kernel(): + T.func_attr({"target": T.target("cuda", host="c")}) + T.attr(T.target("cuda"), "target", 0) + pool = T.SMEMPool() + first = pool.alloc((3,), "float4_e2m1fn") + second = pool.alloc((4,), "float32", align=16) + pool.move_base_to(64) + pool.commit(size) + T.evaluate(first.data) + T.evaluate(second.data) + + allocations = [] + views = [] + + def collect(node): + if isinstance(node, tvm.tirx.AllocBuffer): + allocations.append(node) + if isinstance(node, tvm.tirx.DeclBuffer): + views.append(node) + + tvm_ffi.structural_walk(kernel.body, collect) + assert len(allocations) == 1 + backing = allocations[0].buffer + assert int(backing.shape[0]) == (64 if size is None else size) + assert len(views) == 2 + assert all(view.data.args[0].same_as(backing) for view in views) + assert int(views[1].buffer.elem_offset) == 4 + tvm.ir.assert_structural_equal( + kernel, tvm.script.from_source(kernel.script(), extra_vars={"T": T}) + ) + + split = tvm.tirx.transform.SplitHostDevice()(tvm.IRModule({"kernel": kernel})) + calls = [] + + def collect_launch(node): + if isinstance(node, tvm.ir.Call) and node.op.name == "tirx.call_ffi_kernel": + calls.append(node) + + tvm_ffi.structural_walk(split["kernel"].body, collect_launch) + assert len(calls) == 1 + analyzer = tvm.sym.Analyzer() + assert int(analyzer.simplify(calls[0].args[-1])) == (64 if size is None else size) + + +def test_smem_pool_requires_commit(): + with pytest.raises(ValueError, match=r"SMEMPool.commit\(\) must be called"): + + @T.prim_func + def kernel(): + pool = T.SMEMPool() + view = pool.alloc((16,), "uint8") + T.evaluate(view.data) + + +def test_smem_pool_commit_rejects_small_size(): + with pytest.raises(AssertionError, match="smaller than"): + + @T.prim_func + def kernel(): + pool = T.SMEMPool() + view = pool.alloc((16,), "uint8") + pool.commit(8) + T.evaluate(view.data) + + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/python/tirx/test_hint.py b/tests/python/tirx/test_hint.py index a02377f655df..25840ef0e8e1 100644 --- a/tests/python/tirx/test_hint.py +++ b/tests/python/tirx/test_hint.py @@ -14,147 +14,19 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -"""Tests for T.hint() — universal directive primitive for TIRx sketch language.""" - -import tvm_ffi +"""Tests for hint configuration on tile primitive calls.""" import tvm import tvm.script import tvm.testing -from tvm.ir import TensorRegion, assert_structural_equal +from tvm.ir import assert_structural_equal from tvm.script import tirx as T -from tvm.tirx import AttrStmt def from_source(code): return tvm.script.from_source(code, extra_vars={"I": tvm.script.ir, "T": tvm.script.tirx}) -def test_hint_statement(): - """T.hint("msg") as a bare statement produces an AttrStmt with attr_key=tirx_hint.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("persistent tile scheduler with L2 swizzle") - T.evaluate(0) - - # Walk the IR to find the AttrStmt with tirx_hint - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - # node is now a Map with "message" key - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "persistent tile scheduler with L2 swizzle" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_context_manager(): - """with T.hint("msg"): scopes its body inside the AttrStmt.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - with T.hint("software pipeline, depth 4"): - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "software pipeline, depth 4" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_with_attrs(): - """T.hint("msg", key="value") passes structured attrs in Map node.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("scheduler", mode="persistent", depth="4") - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "scheduler" - assert str(stmt.node["mode"]) == "persistent" - assert str(stmt.node["depth"]) == "4" - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' not found" - - -def test_hint_printer_roundtrip_statement(): - """Verify T.hint("msg") prints as T.hint("msg") and roundtrips through script/parse.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("persistent tile scheduler with L2 swizzle") - T.evaluate(0) - - code = func.script() - assert 'hint("persistent tile scheduler with L2 swizzle")' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - -def test_hint_printer_roundtrip_context_manager(): - """Verify with T.hint("msg"): prints correctly and roundtrips.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - with T.hint("software pipeline, depth 4"): - T.evaluate(0) - - code = func.script() - assert 'hint("software pipeline, depth 4")' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - -def test_hint_printer_roundtrip_with_attrs(): - """Verify T.hint("msg", key="val") prints with kwargs and roundtrips.""" - - @T.prim_func - def func(_A: T.Buffer((64,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("scheduler", mode="persistent") - T.evaluate(0) - - code = func.script() - assert 'hint("scheduler"' in code - assert 'mode="persistent"' in code - reparsed = from_source(code) - assert_structural_equal(func, reparsed) - - def test_hint_keyword_arg_on_tx_op(): """Tx.op(..., hint="msg") stores hint in TilePrimitiveCall.config.""" from tvm.tirx.buffer import decl_buffer @@ -191,70 +63,5 @@ def func( assert_structural_equal(func, reparsed) -def test_hint_no_message(): - """T.hint(access=...) with no message string.""" - - @T.prim_func - def func(A: T.Buffer((128,), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([1, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint(access=A[0:64]) - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - # Should have "access" key but no "message" key - assert "access" in stmt.node - assert "message" not in stmt.node - - assert isinstance(stmt.node["access"], TensorRegion) - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with attr_key='tirx_hint' containing access not found" - - -def test_hint_access_buffer_region(): - """T.hint(access=A[region]) stores the BufferRegion structurally in the IR.""" - - @T.prim_func - def func(A: T.Buffer((128, 64), "float32", scope="global")) -> None: - bx, by, bz = T.cta_id([2, 1, 1]) - warp_id = T.warp_id([1]) - lane_id = T.lane_id([32]) - T.hint("partition", access=A[bx * 64 : (bx + 1) * 64, 0:64]) - T.evaluate(0) - - found = [False] - - def visit(stmt): - if isinstance(stmt, AttrStmt) and stmt.attr_key == "tirx_hint": - assert isinstance(stmt.node, tvm.ir.Map) - assert str(stmt.node["message"]) == "partition" - assert "access" in stmt.node - - assert isinstance(stmt.node["access"], TensorRegion) - br = stmt.node["access"] - assert br.source.name == "A" - assert len(br.region) == 2 - found[0] = True - - tvm_ffi.structural_walk(func.body, visit) - assert found[0], "Expected AttrStmt with structured BufferRegion access not found" - - if __name__ == "__main__": - test_hint_statement() - test_hint_context_manager() - test_hint_with_attrs() - test_hint_printer_roundtrip_statement() - test_hint_printer_roundtrip_context_manager() - test_hint_printer_roundtrip_with_attrs() - test_hint_keyword_arg_on_tx_op() - test_hint_keyword_arg_on_tx_op_roundtrip() - test_hint_no_message() - test_hint_access_buffer_region() + tvm.testing.main() From 751df8bc3fd5217142042be9364f161727342009 Mon Sep 17 00:00:00 2001 From: tqchen Date: Tue, 29 Sep 2026 11:43:29 +0000 Subject: [PATCH 2/2] [FIX][TIR] Preserve explicit dynamic shared-memory sizing Keep launch size independent of placeholder allocation extents through the existing declaration protocol. Use the original pool allocation and commit path without deferred buffer construction. --- python/tvm/backend/cuda/lang/alloc_pool.py | 33 ++++----- python/tvm/tirx/script/ir_builder/ir.py | 7 -- src/tirx/script/ir_builder/ir.cc | 25 ------- src/tirx/transform/split_host_device.cc | 57 +++++++++++++-- .../cuda/copy_async/test_tma.py | 4 ++ tests/python/tirx/test_alloc_pool.py | 71 ------------------- 6 files changed, 70 insertions(+), 127 deletions(-) diff --git a/python/tvm/backend/cuda/lang/alloc_pool.py b/python/tvm/backend/cuda/lang/alloc_pool.py index ee70456f46d4..f3d7eba2659f 100644 --- a/python/tvm/backend/cuda/lang/alloc_pool.py +++ b/python/tvm/backend/cuda/lang/alloc_pool.py @@ -400,18 +400,17 @@ class SMEMPool: Parameters ---------- ptr : Var or None, optional - 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 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 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_deferred(self._allocation_shape, "uint8", "shared.dyn") + self.buf = ir.alloc_buffer([0], "uint8", scope="shared.dyn") self.ptr = self.buf.data self._owns_buffer = True else: @@ -421,13 +420,6 @@ 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, @@ -437,8 +429,6 @@ 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 @@ -482,14 +472,12 @@ 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): - """Finalize the backing allocation's byte extent. + """Emit pool size annotation into the IR. Must be called after all ``alloc()`` / ``move_base_to()`` calls. @@ -501,11 +489,16 @@ 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})" ) - self._committed_size = resolved + 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) + ) + ) diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index bc833bdd76ca..58134a768611 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -459,13 +459,6 @@ 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, diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index 5937b01889d6..b720a71c5f34 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -19,7 +19,6 @@ #include #include #include -#include #include #include #include @@ -654,29 +653,6 @@ BufferVar AllocBuffer(ffi::Array 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()> 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().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> { - if (var.same_as(buffer)) return ffi::Any(resolved); - return ffi::Unchanged(); - }; - frame->stmts = ffi::StructuralMap(frame->stmts, replace) - .as_or_throw>(); - }); - return buffer; -} - tvm::tirx::Stmt Evaluate(Expr value) { tvm::tirx::Stmt stmt = tvm::tirx::Evaluate(value); AddToParent(stmt); @@ -768,7 +744,6 @@ 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) diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index 3bb27f15f112..a2ab6f41a1ee 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -429,8 +429,20 @@ class DeviceInfoCollector : public StmtExprVisitor { collector->info_.launch_params.push_back( tvm::runtime::launch_param::kUseRequiredBlockDimension); } - // Dynamic shared memory is the last kernel launch parameter. Its size is - // derived from the backing allocation, including pool allocations. + // The dynamic shared memory is required to be the last of the kernel + // launch parameters. An explicit tirx.dyn_smem_bytes declaration wins; + // otherwise fall back to the size inferred from the allocation extent. + // A zero-extent allocation is a pool-style extern placeholder, so having + // neither a declaration nor a usable extent is an authoring error. + if (!collector->dyn_shmem_size.has_value() && collector->inferred_shmem_size_.has_value()) { + const auto* inferred = collector->inferred_shmem_size_.value().as(); + TVM_FFI_ICHECK(!(inferred && inferred->value == 0)) + << "PrimFunc " << gvar->name_hint + << " allocates dynamic shared memory with a placeholder extent but does not declare " + "its size; annotate the kernel with tirx.dyn_smem_bytes (SMEMPool.commit() emits " + "it)."; + collector->dyn_shmem_size = collector->inferred_shmem_size_; + } if (collector->dyn_shmem_size) { collector->info_.launch_params.push_back( tvm::runtime::launch_param::kUseDynamicSharedMemoryTag); @@ -456,7 +468,7 @@ class DeviceInfoCollector : public StmtExprVisitor { if (launch_param == tvm::runtime::launch_param::kUseDynamicSharedMemoryTag) { TVM_FFI_ICHECK(dyn_shmem_size.has_value()) << "Compute kernel requires launch parameter \"" << launch_param - << "\", but PrimFunc has no dynamic shared memory allocation."; + << "\", but PrimFunc did not declare tirx.dyn_smem_bytes."; return dyn_shmem_size.value(); } @@ -490,6 +502,15 @@ class DeviceInfoCollector : public StmtExprVisitor { } ffi::Optional Visit_(const AttrStmtNode* op) final { + if (op->attr_key == "tirx.dyn_smem_bytes") { + // Kernel-level declaration of the dynamic shared memory launch size. + // The backing shared.dyn allocation is an extern placeholder; this + // attribute is the single source of truth for the launch parameter. + TVM_FFI_ICHECK(!dyn_shmem_size.has_value()) + << "Only one tirx.dyn_smem_bytes declaration is allowed per kernel."; + TVM_FFI_ICHECK(op->value.as()) << "tirx.dyn_smem_bytes must be an IntImm"; + dyn_shmem_size = op->value.as_or_throw(); + } if (op->attr_key == attr::thread_extent) { ffi::String thread_tag; if (auto iv = op->node.as()) { @@ -531,6 +552,10 @@ class DeviceInfoCollector : public StmtExprVisitor { << "Only one dynamic shared memory allocation is allowed."; saw_dyn_shared_alloc_ = true; + // Fallback launch size inferred from the allocation extent, used when + // no tirx.dyn_smem_bytes declaration is present (e.g. s_tir schedules + // allocate shared.dyn with a concrete extent). A zero extent is a + // pool-style extern placeholder and carries no size information. TVM_FFI_ICHECK_GT(op->buffer->shape.size(), 0); PrimExpr dyn_size = IntImm::Int32(1); for (const auto& extent : op->buffer->shape) { @@ -545,7 +570,7 @@ class DeviceInfoCollector : public StmtExprVisitor { dyn_size = ffi::StructuralMap(dyn_size, f_substitute) .as_or_throw(); } - dyn_shmem_size = dyn_size; + inferred_shmem_size_ = dyn_size; } return StmtExprVisitor::Visit_(op); } @@ -560,6 +585,9 @@ class DeviceInfoCollector : public StmtExprVisitor { ffi::Optional dyn_shmem_size{std::nullopt}; // Whether a shared.dyn allocation was seen. bool saw_dyn_shared_alloc_{false}; + // Launch size inferred from the allocation extent (fallback when no + // tirx.dyn_smem_bytes declaration is present). + ffi::Optional inferred_shmem_size_{std::nullopt}; // Flag-only launch attributes requested by the original PrimFunc. bool use_programmatic_dependent_launch_{false}; bool use_cooperative_launch_{false}; @@ -682,6 +710,27 @@ class DeviceKernelMutator : public StmtExprMutator { Target target = func->GetAttr(tvm::attr::kTarget).value(); bool preserve_early_returns = target->kind->name == "cuda"; write_ptr->body = ReturnRemover::Apply(write_ptr->body, !preserve_early_returns); + // The dyn-smem size declaration was consumed by DeviceInfoCollector; + // it has no meaning inside the kernel body. + class StripDynSmemAttr : public StmtExprMutator { + public: + using StmtExprMutator::Mutate; + using StmtExprMutator::Mutate_; + UnchangedOr Mutate(ffi::AnyView input, InplaceMode inplace_mode) override { + if (input.as()) return ffi::Unchanged(); + return StmtExprMutator::Mutate(input, inplace_mode); + } + + UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { + if (op->attr_key == "tirx.dyn_smem_bytes") { + return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); + } + return StmtExprMutator::Mutate_(op, inplace_mode); + } + }; + write_ptr->body = ffi::make_object() + ->Mutate(write_ptr->body, InplaceMode::kAllow) + .ValueOrUnchanged(write_ptr->body); } func = WithAttrs(std::move(func), diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py index 20e24d9e5019..ab2e7b49f857 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py @@ -1532,6 +1532,7 @@ def rank_change(A: T.Buffer((8, 8), 'float16')): T.cta_id([1]) tid = T.thread_id([1]) dyn = T.alloc_buffer((65,), "uint64", scope="shared.dyn") + T.attr({"tirx.dyn_smem_bytes": 65 * 8}) A_smem = T.decl_buffer((64,), "float16", dyn.data, layout=T.TileLayout(T.S[64])) mbar = T.decl_buffer((1,), "uint64", dyn.data, elem_offset=16) if tid == 0: @@ -1558,6 +1559,7 @@ def selector_gather( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((520,), "uint64", scope="shared.dyn") + T.attr({"tirx.dyn_smem_bytes": 520 * 8}) A_smem = T.decl_buffer( (4, 64), "bfloat16", dyn.data, layout=T.TileLayout(T.S[4, 64]) ) @@ -1728,6 +1730,7 @@ def kernel( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((shared_bytes + 8,), "uint8", scope="shared.dyn") + T.attr({"tirx.dyn_smem_bytes": shared_bytes + 8}) q_smem = T.decl_buffer( (64, 512), "bfloat16", dyn.data, scope="shared.dyn", layout=q_layout ) @@ -2001,6 +2004,7 @@ def kernel( T.cta_id([1]) tid = T.thread_id([128]) dyn = T.alloc_buffer((shared_bytes + 64,), "uint8", scope="shared.dyn") + T.attr({"tirx.dyn_smem_bytes": shared_bytes + 64}) A_smem = T.decl_buffer( (4, cols), dtype, dyn.data, layout=T.TileLayout(T.S[4, cols]) ) diff --git a/tests/python/tirx/test_alloc_pool.py b/tests/python/tirx/test_alloc_pool.py index 6902a64fda6d..9614531ec08a 100644 --- a/tests/python/tirx/test_alloc_pool.py +++ b/tests/python/tirx/test_alloc_pool.py @@ -17,10 +17,7 @@ """Tests for CUDA allocation pool validation.""" import pytest -import tvm_ffi -import tvm -from tvm.script import tirx as T from tvm.tirx.cuda.lang.alloc_pool import _validate_mma_alloc_shape from tvm.tirx.cuda.tile_primitive.tma_utils import SwizzleMode @@ -116,73 +113,5 @@ def test_swizzle_none_skips_validation(self): _validate_mma_alloc_shape((128,), "bfloat16", SwizzleMode.SWIZZLE_NONE) -@pytest.mark.parametrize("size", [None, 128]) -def test_smem_pool_commits_allocation_extent(size): - @T.prim_func - def kernel(): - T.func_attr({"target": T.target("cuda", host="c")}) - T.attr(T.target("cuda"), "target", 0) - pool = T.SMEMPool() - first = pool.alloc((3,), "float4_e2m1fn") - second = pool.alloc((4,), "float32", align=16) - pool.move_base_to(64) - pool.commit(size) - T.evaluate(first.data) - T.evaluate(second.data) - - allocations = [] - views = [] - - def collect(node): - if isinstance(node, tvm.tirx.AllocBuffer): - allocations.append(node) - if isinstance(node, tvm.tirx.DeclBuffer): - views.append(node) - - tvm_ffi.structural_walk(kernel.body, collect) - assert len(allocations) == 1 - backing = allocations[0].buffer - assert int(backing.shape[0]) == (64 if size is None else size) - assert len(views) == 2 - assert all(view.data.args[0].same_as(backing) for view in views) - assert int(views[1].buffer.elem_offset) == 4 - tvm.ir.assert_structural_equal( - kernel, tvm.script.from_source(kernel.script(), extra_vars={"T": T}) - ) - - split = tvm.tirx.transform.SplitHostDevice()(tvm.IRModule({"kernel": kernel})) - calls = [] - - def collect_launch(node): - if isinstance(node, tvm.ir.Call) and node.op.name == "tirx.call_ffi_kernel": - calls.append(node) - - tvm_ffi.structural_walk(split["kernel"].body, collect_launch) - assert len(calls) == 1 - analyzer = tvm.sym.Analyzer() - assert int(analyzer.simplify(calls[0].args[-1])) == (64 if size is None else size) - - -def test_smem_pool_requires_commit(): - with pytest.raises(ValueError, match=r"SMEMPool.commit\(\) must be called"): - - @T.prim_func - def kernel(): - pool = T.SMEMPool() - view = pool.alloc((16,), "uint8") - T.evaluate(view.data) - - -def test_smem_pool_commit_rejects_small_size(): - with pytest.raises(AssertionError, match="smaller than"): - - @T.prim_func - def kernel(): - pool = T.SMEMPool() - view = pool.alloc((16,), "uint8") - pool.commit(8) - T.evaluate(view.data) - - if __name__ == "__main__": pytest.main([__file__, "-v"])