From c06fcb895144e28981be3fd01f57e71f4e66cb97 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Tue, 6 Oct 2026 03:58:07 +0000 Subject: [PATCH] [TIRx][CUDA] Represent kernel constraints with CUDA leaf declarations --- docs/tirx/native_basics/cuda/functions.rst | 28 +-- include/tvm/tirx/function.h | 37 ---- src/backend/cuda/codegen/codegen_cuda.cc | 184 ++++++++++++------ src/backend/cuda/codegen/codegen_cuda.h | 13 ++ src/backend/cuda/op/target_builtin.cc | 15 ++ src/target/source/codegen_c.cc | 19 +- src/target/source/codegen_c.h | 6 + src/tirx/transform/split_host_device.cc | 116 +---------- .../test_tir_transform_split_host_device.py | 16 +- .../python/tirx/codegen/test_codegen_cuda.py | 46 ++--- 10 files changed, 218 insertions(+), 262 deletions(-) diff --git a/docs/tirx/native_basics/cuda/functions.rst b/docs/tirx/native_basics/cuda/functions.rst index c09ec847a37b..4debc34cf3aa 100644 --- a/docs/tirx/native_basics/cuda/functions.rst +++ b/docs/tirx/native_basics/cuda/functions.rst @@ -287,13 +287,12 @@ By default, the block size also drives the kernel's ``__launch_bounds__``. The first argument (max threads per block) is set automatically from the thread extent. To also set the second argument — the minimum blocks per SM, an occupancy hint — add -``Tx.attr({"tirx.launch_bounds_min_blocks_per_sm": N})`` in the device region (note: -``Tx.attr``, not ``func_attr``): +``Tx.cuda.launch_bounds_min_blocks_per_sm(N)`` in the device region: .. code-block:: python Tx.device_entry() - Tx.attr({"tirx.launch_bounds_min_blocks_per_sm": 2}) # second launch-bounds arg + Tx.cuda.launch_bounds_min_blocks_per_sm(2) # second launch-bounds arg bx = Tx.cta_id([1]); tx = Tx.thread_id([256]) ... @@ -301,16 +300,23 @@ occupancy hint — add extern "C" __global__ void __launch_bounds__(256, 2) scale_kernel(...) { ... } -Without the attr the second argument is omitted (just ``__launch_bounds__(256)``). +Without this declaration the second argument is omitted (just ``__launch_bounds__(256)``). + +``Tx.cuda.launch_bounds_max_blocks_per_cluster(N)`` supplies the third operand +and requires a minimum-blocks declaration. ``Tx.cuda.max_registers_per_thread(N)`` +emits ``__maxnreg__(N)`` and cannot accompany launch bounds. These declarations +accept positive integer constants. Matching repetitions are allowed; conflicting +values are rejected. They configure the containing kernel and produce no runtime +instructions at their textual position. Some kernels require an exact block and cluster shape instead of an advisory -maximum. Set ``tirx.required_block_size`` to ``1`` to make the thread and -cluster extents a compile-time launch contract: +maximum. Use ``Tx.cuda.required_block_size`` with the three thread dimensions followed +by the three cluster dimensions to declare a compile-time launch contract: .. code-block:: python Tx.device_entry() - Tx.attr({"tirx.required_block_size": 1}) + Tx.cuda.required_block_size(128, 1, 1, 1, 2, 1) bx, by = Tx.cta_id([4, 2]) _, cy = Tx.cta_id_in_cluster([1, 2]) tx = Tx.thread_id([128]) @@ -321,15 +327,15 @@ cluster extents a compile-time launch contract: extern "C" __global__ void __block_size__((128, 1, 1), (1, 2, 1)) kernel(...) { ... } This requires CUDA Toolkit 13 or newer. All thread and cluster dimensions must -be static; CUDA lowers ``__block_size__`` to PTX ``.reqntid`` and checks the +be positive constants matching the declared launch extents; CUDA lowers ``__block_size__`` to PTX ``.reqntid`` and checks the same dimensions at launch. A preferred cluster dimension must be absent or equal to the required cluster dimension, and each logical block-grid dimension must be divisible by its cluster dimension. -``tirx.required_block_size`` can be combined with the launch-bounds attributes +``Tx.cuda.required_block_size`` can be combined with the launch-bounds declarations when an occupancy hint is also needed; code generation then emits both ``__block_size__`` and ``__launch_bounds__``. It cannot be combined with -``tirx.max_registers``. +``Tx.cuda.max_registers_per_thread``. At run time the kernel is launched through the **CUDA Driver API**. TVM's CUDA runtime loads the module (``cuModuleLoadData``), fetches the function @@ -338,7 +344,7 @@ runtime loads the module (``cuModuleLoadData``), fetches the function the config carries a list of launch *attributes* — the thread-block **cluster dimension** and **preferred cluster dimension** (Hopper/Blackwell), plus optional programmatic-dependent-launch and cooperative-launch flags. Kernels with -``tirx.required_block_size`` instead use CUDA's required-block sentinel; their +``Tx.cuda.required_block_size`` instead use CUDA's required-block sentinel; their compile-time cluster shape replaces the ordinary runtime cluster attribute. In outline, ``src/backend/cuda/runtime/cuda_module.cc`` follows this path: diff --git a/include/tvm/tirx/function.h b/include/tvm/tirx/function.h index 29204fe7a09c..e6e68a0954ee 100644 --- a/include/tvm/tirx/function.h +++ b/include/tvm/tirx/function.h @@ -225,43 +225,6 @@ namespace attr { */ constexpr const char* kKernelLaunchParams = "tirx.kernel_launch_params"; -/*! - * \brief CUDA launch bound minimum CTAs per SM. - * - * Type: IntImm - */ -constexpr const char* kLaunchBoundsMinBlocksPerSM = "tirx.launch_bounds_min_blocks_per_sm"; - -/*! - * \brief CUDA launch bound maximum CTAs per cluster. - * - * Type: IntImm - */ -constexpr const char* kLaunchBoundsMaxBlocksPerCluster = - "tirx.launch_bounds_max_blocks_per_cluster"; - -/*! - * \brief CUDA maximum registers per thread. - * - * Emits the CUDA 13 ``__maxnreg__`` kernel qualifier. This attribute is - * mutually exclusive with the launch-bounds attributes. - * - * Type: IntImm - */ -constexpr const char* kMaxRegisters = "tirx.max_registers"; - -/*! - * \brief Require CUDA to use the statically-declared block and cluster dimensions. - * - * Emits the CUDA 13 ``__block_size__`` kernel qualifier. Unlike - * ``__launch_bounds__``, this is an exact launch contract: CUDA derives the - * PTX ``.reqntid`` directive from the thread extents, and interprets the - * launch grid in clusters using the cluster-CTA extents. - * - * Type: IntImm (must be 1) - */ -constexpr const char* kRequiredBlockSize = "tirx.required_block_size"; - /*! * \brief Whether to set noalias rule on the function arguments. * diff --git a/src/backend/cuda/codegen/codegen_cuda.cc b/src/backend/cuda/codegen/codegen_cuda.cc index 9affe0ec570a..c03c00382f9f 100644 --- a/src/backend/cuda/codegen/codegen_cuda.cc +++ b/src/backend/cuda/codegen/codegen_cuda.cc @@ -175,6 +175,11 @@ CodeGenCUDA::CodeGenCUDA(Target target) : target(target) { restrict_keyword_ = " void CodeGenCUDA::PrintFunctionSignature(const ffi::String& function_name, const PrimFunc& func, std::ostream& os) { + PrintFunctionPrefix(func, os); + CodeGenC::PrintFunctionSignature(function_name, func, os); +} + +void CodeGenCUDA::PrintFunctionPrefix(const PrimFunc& func, std::ostream& os) { CallingConv calling_conv = func->GetAttr(tvm::attr::kCallingConv, CallingConv::kDefault).value(); in_kernel_launch_ = (calling_conv == CallingConv::kDeviceKernelLaunch); @@ -186,7 +191,6 @@ void CodeGenCUDA::PrintFunctionSignature(const ffi::String& function_name, const TVM_FFI_THROW(InternalError) << "Unsupported calling convention for cuda codegen: " << static_cast(calling_conv); } - CodeGenC::PrintFunctionSignature(function_name, func, os); } class ThreadIdxExtractor : public tirx::StmtExprVisitor { @@ -239,71 +243,98 @@ class ThreadIdxExtractor : public tirx::StmtExprVisitor { PrimExpr clusterCtaIdx_z_ext = IntImm::Int32(1); }; -void CodeGenCUDA::PrintExtraAttrs(const PrimFunc& f, std::ostream& os) { +void CodeGenCUDA::InitFuncState(const PrimFunc& func) { + CodeGenC::InitFuncState(func); + min_blocks_per_sm_.reset(); + max_blocks_per_cluster_.reset(); + max_registers_per_thread_.reset(); + required_block_size_.reset(); + in_kernel_launch_ = + func->GetAttr(tvm::attr::kCallingConv, CallingConv::kDefault).value() == + CallingConv::kDeviceKernelLaunch; + // Thread metadata is also needed while emitting cluster index accesses. auto extractor = ffi::make_object(); - extractor->Visit(f->body); + extractor->Visit(func->body); + launch_dimensions_ = {extractor->threadIdx_x_ext, extractor->threadIdx_y_ext, + extractor->threadIdx_z_ext, extractor->clusterCtaIdx_x_ext, + extractor->clusterCtaIdx_y_ext, extractor->clusterCtaIdx_z_ext}; sym::Analyzer analyzer; - PrimExpr threadIdx_ext = analyzer->Simplify( - extractor->threadIdx_x_ext * extractor->threadIdx_y_ext * extractor->threadIdx_z_ext); - PrimExpr cluster_cta_yz_ext = - analyzer->Simplify(extractor->clusterCtaIdx_y_ext * extractor->clusterCtaIdx_z_ext); - if (const IntImmNode* const cluster_cta_yz_ext_int = cluster_cta_yz_ext.as()) { - cluster_cta_x_is_linear_rank_ = cluster_cta_yz_ext_int->value == 1; - } else { - cluster_cta_x_is_linear_rank_ = false; - } - auto max_registers = f->GetAttr(tirx::attr::kMaxRegisters); - auto required_block_size = f->GetAttr(tirx::attr::kRequiredBlockSize); - if (required_block_size.has_value()) { - TVM_FFI_ICHECK_EQ(required_block_size.value(), 1); - TVM_FFI_ICHECK(!max_registers.has_value()) - << tirx::attr::kRequiredBlockSize << " cannot be combined with maximum registers"; - const auto* tx = extractor->threadIdx_x_ext.as(); - const auto* ty = extractor->threadIdx_y_ext.as(); - const auto* tz = extractor->threadIdx_z_ext.as(); - const auto* cx = extractor->clusterCtaIdx_x_ext.as(); - const auto* cy = extractor->clusterCtaIdx_y_ext.as(); - const auto* cz = extractor->clusterCtaIdx_z_ext.as(); - TVM_FFI_ICHECK(tx && ty && tz && cx && cy && cz) - << tirx::attr::kRequiredBlockSize << " requires static thread and cluster dimensions"; - os << " __block_size__((" << tx->value << ", " << ty->value << ", " << tz->value << "), (" - << cx->value << ", " << cy->value << ", " << cz->value << "))"; - if (!f->GetAttr(tirx::attr::kLaunchBoundsMinBlocksPerSM).has_value()) { - TVM_FFI_ICHECK(!f->GetAttr(tirx::attr::kLaunchBoundsMaxBlocksPerCluster).has_value()) - << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " requires " - << tirx::attr::kLaunchBoundsMinBlocksPerSM; - return; - } + cluster_cta_x_is_linear_rank_ = + is_one(analyzer->Simplify(launch_dimensions_[4] * launch_dimensions_[5])); +} + +void CodeGenCUDA::DeclareFunction(const GlobalVar& gvar, const PrimFunc& func) { + if (!RegisterFunctionName(gvar, func)) return; + // Definitions supply their body-derived qualifiers in AddFunction. + if (!func->body.has_value()) { + InitFuncState(func); + PrintFunctionSignature(GetFunctionName(gvar), func, fwd_decl_stream); + fwd_decl_stream << ";\n"; } - if (max_registers.has_value()) { - TVM_FFI_ICHECK_GT(max_registers.value(), 0); - TVM_FFI_ICHECK(!f->GetAttr(tirx::attr::kLaunchBoundsMinBlocksPerSM).has_value() && - !f->GetAttr(tirx::attr::kLaunchBoundsMaxBlocksPerCluster).has_value()) - << tirx::attr::kMaxRegisters << " cannot be combined with CUDA launch bounds"; - os << " __maxnreg__(" << max_registers.value() << ")"; +} + +void CodeGenCUDA::AddFunction(const GlobalVar& gvar, const PrimFunc& func) { + DeclareFunction(gvar, func); + if (!func->body.has_value()) return; + InitFuncState(func); + std::ostringstream parameters; + PrintFunctionParameters(func, parameters); + + // Generate the body once, retaining the settings declared by its leaf ops. + std::ostringstream previous_functions; + stream.swap(previous_functions); + stream << " {\n"; + PreFunctionBody(func); + int func_scope = BeginScope(); + PrintStmt(func->body.value()); + EndScope(func_scope); + PrintIndent(); + stream << "}\n\n"; + std::string body = stream.str(); + stream.swap(previous_functions); + + std::ostringstream signature; + PrintFunctionPrefix(func, signature); + PrintType(func->ret_type, signature); + PrintExtraAttrs(func, signature); + signature << " " << GetFunctionName(gvar) << parameters.str(); + fwd_decl_stream << signature.str() << ";\n"; + stream << signature.str() << body; +} + +void CodeGenCUDA::PrintExtraAttrs(const PrimFunc& f, std::ostream& os) { + TVM_FFI_ICHECK(!max_blocks_per_cluster_ || min_blocks_per_sm_) + << "CUDA maximum blocks per cluster requires minimum blocks per SM"; + TVM_FFI_ICHECK(!max_registers_per_thread_ || (!min_blocks_per_sm_ && !max_blocks_per_cluster_)) + << "Maximum registers per thread cannot be combined with CUDA launch bounds"; + TVM_FFI_ICHECK(!max_registers_per_thread_ || !required_block_size_) + << "Required block size cannot be combined with maximum registers"; + if (required_block_size_) { + for (size_t i = 0; i < launch_dimensions_.size(); ++i) { + const auto* dim = launch_dimensions_[i].as(); + TVM_FFI_ICHECK(dim) << "Required block size requires static thread and cluster dimensions"; + TVM_FFI_ICHECK_EQ(dim->value, (*required_block_size_)[i]) + << "Required block size must agree with declared launch extents"; + } + const auto& dims = *required_block_size_; + os << " __block_size__((" << dims[0] << ", " << dims[1] << ", " << dims[2] << "), (" << dims[3] + << ", " << dims[4] << ", " << dims[5] << "))"; + if (!min_blocks_per_sm_) return; + } + if (max_registers_per_thread_) { + os << " __maxnreg__(" << *max_registers_per_thread_ << ")"; return; } - if (const IntImmNode* const threadIdx_ext_int = threadIdx_ext.as()) { - if (threadIdx_ext_int->value == 1) { - // unable to extract the number of threads per block, hence directly return - return; - } - auto min_blocks_per_sm = f->GetAttr(tirx::attr::kLaunchBoundsMinBlocksPerSM); - auto max_blocks_per_cluster = f->GetAttr(tirx::attr::kLaunchBoundsMaxBlocksPerCluster); - if (min_blocks_per_sm.has_value()) { - TVM_FFI_ICHECK_GT(min_blocks_per_sm.value(), 0); - os << " __launch_bounds__(" << threadIdx_ext_int->value << ", " << min_blocks_per_sm.value(); - if (max_blocks_per_cluster.has_value()) { - TVM_FFI_ICHECK_GT(max_blocks_per_cluster.value(), 0); - os << ", " << max_blocks_per_cluster.value(); - } - os << ")"; - } else { - TVM_FFI_ICHECK(!max_blocks_per_cluster.has_value()) - << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " requires " - << tirx::attr::kLaunchBoundsMinBlocksPerSM; - os << " __launch_bounds__(" << threadIdx_ext_int->value << ")"; + sym::Analyzer analyzer; + PrimExpr threads = + analyzer->Simplify(launch_dimensions_[0] * launch_dimensions_[1] * launch_dimensions_[2]); + if (const auto* count = threads.as(); count && count->value != 1) { + os << " __launch_bounds__(" << count->value; + if (min_blocks_per_sm_) { + os << ", " << *min_blocks_per_sm_; + if (max_blocks_per_cluster_) os << ", " << *max_blocks_per_cluster_; } + os << ")"; } } @@ -1684,6 +1715,39 @@ void CodeGenCUDA::DispatchAllocTensor(const BindNode* op, const CallNode* buffer } void CodeGenCUDA::Dispatch_(const EvaluateNode* op) { + if (const auto* call = op->value.as()) { + static const Op min_blocks = Op::Get("tirx.cuda.launch_bounds_min_blocks_per_sm"); + static const Op max_blocks = Op::Get("tirx.cuda.launch_bounds_max_blocks_per_cluster"); + static const Op max_registers = Op::Get("tirx.cuda.max_registers_per_thread"); + static const Op required_block = Op::Get("tirx.cuda.required_block_size"); + std::optional* setting = nullptr; + if (call->op.same_as(min_blocks)) { + setting = &min_blocks_per_sm_; + } else if (call->op.same_as(max_blocks)) { + setting = &max_blocks_per_cluster_; + } else if (call->op.same_as(max_registers)) { + setting = &max_registers_per_thread_; + } + if (setting) { + int64_t value = static_cast(call->args[0].as_or_throw()->value); + TVM_FFI_ICHECK_GT(value, 0) << call->op << " must be positive"; + TVM_FFI_ICHECK(!setting->has_value() || setting->value() == value) + << "Conflicting " << call->op << " values"; + *setting = value; + return; + } + if (call->op.same_as(required_block)) { + std::array dimensions; + for (size_t i = 0; i < dimensions.size(); ++i) { + dimensions[i] = static_cast(call->args[i].as_or_throw()->value); + TVM_FFI_ICHECK_GT(dimensions[i], 0) << "Required block dimensions must be positive"; + } + TVM_FFI_ICHECK(!required_block_size_ || *required_block_size_ == dimensions) + << "Conflicting required block size values"; + required_block_size_ = dimensions; + return; + } + } if (auto value = op->value.as(); value && is_const_int(value.value())) return; CodeGenC::Dispatch_(op); } diff --git a/src/backend/cuda/codegen/codegen_cuda.h b/src/backend/cuda/codegen/codegen_cuda.h index bfcca6c2a949..919079483617 100644 --- a/src/backend/cuda/codegen/codegen_cuda.h +++ b/src/backend/cuda/codegen/codegen_cuda.h @@ -28,6 +28,8 @@ #include #include +#include +#include #include #include @@ -48,6 +50,9 @@ class CodeGenCUDA final : public CodeGenC { }); } // override behavior + void DeclareFunction(const GlobalVar& gvar, const PrimFunc& func) final; + void AddFunction(const GlobalVar& gvar, const PrimFunc& func) final; + void InitFuncState(const PrimFunc& func) final; void PrintFunctionSignature(const ffi::String& function_name, const PrimFunc& func, std::ostream& os) final; void PrintExtraAttrs(const PrimFunc& f, std::ostream& os) final; // NOLINT(*) @@ -90,6 +95,14 @@ class CodeGenCUDA final : public CodeGenC { bool skip_first_arg, std::ostream& os) final; // NOLINT(*) private: + void PrintFunctionPrefix(const PrimFunc& func, std::ostream& os); + std::array launch_dimensions_{IntImm::Int32(1), IntImm::Int32(1), IntImm::Int32(1), + IntImm::Int32(1), IntImm::Int32(1), IntImm::Int32(1)}; + std::optional min_blocks_per_sm_; + std::optional max_blocks_per_cluster_; + std::optional max_registers_per_thread_; + std::optional> required_block_size_; + // Handle volatile loads void HandleVolatileLoads(const std::string& value, const TensorLoadNode* op, std::ostream& os) final; diff --git a/src/backend/cuda/op/target_builtin.cc b/src/backend/cuda/op/target_builtin.cc index 0c33c27bb526..0306dfb38f24 100644 --- a/src/backend/cuda/op/target_builtin.cc +++ b/src/backend/cuda/op/target_builtin.cc @@ -260,6 +260,21 @@ OpDef& RegisterDeviceIntrinsic(OpDef&& def, const char* op_namespace, CallEffect } void RegisterDeviceIntrinsics() { + // Kernel configuration declarations survive lowering until CUDA body generation. + RegisterDeviceIntrinsic(OpDef("tirx.cuda.launch_bounds_min_blocks_per_sm"), "cuda", + CallEffectKind::kEmbedInfo, sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.launch_bounds_max_blocks_per_cluster"), "cuda", + CallEffectKind::kEmbedInfo, sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.max_registers_per_thread"), "cuda", + CallEffectKind::kEmbedInfo, sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic( + OpDef("tirx.cuda.required_block_size"), "cuda", CallEffectKind::kEmbedInfo, + sig::arg("thread_x"), sig::arg("thread_y"), sig::arg("thread_z"), + sig::arg("cluster_x"), sig::arg("cluster_y"), sig::arg("cluster_z")) + .set_attr("TFixedReturnType", PrimType::Void()); RegisterDeviceIntrinsic(OpDef("tirx.cuda.any_sync"), "cuda", CallEffectKind::kPure, sig::arg("mask"), sig::arg("pred")) .set_attr("TFixedReturnType", PrimType::Int(32)); diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index aa6313b76b5b..fa3a57d30e69 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc @@ -89,7 +89,12 @@ void CodeGenC::PrintFunctionSignature(const ffi::String& function_name, const Pr PrintFuncPrefix(os); PrintType(func->ret_type, os); PrintExtraAttrs(func, os); - os << " " << function_name << "("; + os << " " << function_name; + PrintFunctionParameters(func, os); +} + +void CodeGenC::PrintFunctionParameters(const PrimFunc& func, std::ostream& os) { + os << "("; for (size_t i = 0; i < func->params.size(); ++i) { tirx::Var v = func->params[i]; @@ -139,10 +144,8 @@ void CodeGenC::PrintFunctionSignature(const ffi::String& function_name, const Pr } } -void CodeGenC::DeclareFunction(const GlobalVar& gvar, const PrimFunc& func) { - if (internal_functions_.count(gvar)) { - return; - } +bool CodeGenC::RegisterFunctionName(const GlobalVar& gvar, const PrimFunc& func) { + if (internal_functions_.count(gvar)) return false; auto function_name = [&]() -> ffi::String { if (auto global_symbol = func->GetAttr(tvm::attr::kGlobalSymbol)) { @@ -161,9 +164,13 @@ void CodeGenC::DeclareFunction(const GlobalVar& gvar, const PrimFunc& func) { has_tvm_ffi_main_func_ = true; } internal_functions_.insert({gvar, function_name}); + return true; +} +void CodeGenC::DeclareFunction(const GlobalVar& gvar, const PrimFunc& func) { + if (!RegisterFunctionName(gvar, func)) return; InitFuncState(func); - PrintFunctionSignature(function_name, func, fwd_decl_stream); + PrintFunctionSignature(GetFunctionName(gvar), func, fwd_decl_stream); fwd_decl_stream << ";\n"; } diff --git a/src/target/source/codegen_c.h b/src/target/source/codegen_c.h index c33782706b3c..bf96b0a2c57d 100644 --- a/src/target/source/codegen_c.h +++ b/src/target/source/codegen_c.h @@ -167,6 +167,12 @@ class CodeGenC : public tirx::ExprFunctor, * \param f The function to be compiled. */ virtual void InitFuncState(const PrimFunc& f); + + // Register names independently of emitting declarations for body-derived qualifiers. + bool RegisterFunctionName(const GlobalVar& gvar, const PrimFunc& func); + // Prints parameters and initializes their variable IDs and handle types once. + void PrintFunctionParameters(const PrimFunc& func, std::ostream& os); + // expression void Dispatch_(const VarNode* op, std::ostream& os) override; // NOLINT(*) void Dispatch_(const TensorLoadNode* op, std::ostream& os) override; // NOLINT(*) diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index dd009492d238..cb8ebb7bb1b1 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -104,95 +104,6 @@ PrimFunc AnnotateDeviceRegionsForSplit(PrimFunc func) { // Host/device function extraction -class LaunchBoundsAttrExtractor : 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); - } - Stmt Extract(Stmt stmt) { - min_blocks_per_sm_.reset(); - max_blocks_per_cluster_.reset(); - max_registers_.reset(); - required_block_size_.reset(); - Stmt result = Mutate(stmt, InplaceMode::kAllow).ValueOrUnchanged(stmt); - TVM_FFI_ICHECK(!max_blocks_per_cluster_.has_value() || min_blocks_per_sm_.has_value()) - << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " requires " - << tirx::attr::kLaunchBoundsMinBlocksPerSM; - TVM_FFI_ICHECK(!max_registers_.has_value() || - (!min_blocks_per_sm_.has_value() && !max_blocks_per_cluster_.has_value())) - << tirx::attr::kMaxRegisters << " cannot be combined with CUDA launch bounds"; - TVM_FFI_ICHECK(!required_block_size_.has_value() || !max_registers_.has_value()) - << tirx::attr::kRequiredBlockSize << " cannot be combined with maximum registers"; - return result; - } - - std::optional min_blocks_per_sm() const { return min_blocks_per_sm_; } - std::optional max_blocks_per_cluster() const { return max_blocks_per_cluster_; } - std::optional max_registers() const { return max_registers_; } - std::optional required_block_size() const { return required_block_size_; } - - private: - UnchangedOr Mutate_(const AttrStmtNode* op, InplaceMode inplace_mode) final { - if (op->attr_key == tirx::attr::kLaunchBoundsMinBlocksPerSM) { - const auto* min_blocks_per_sm = op->value.as(); - TVM_FFI_ICHECK(min_blocks_per_sm) - << tirx::attr::kLaunchBoundsMinBlocksPerSM << " expects an integer value"; - TVM_FFI_ICHECK_GT(min_blocks_per_sm->value, 0) - << tirx::attr::kLaunchBoundsMinBlocksPerSM << " must be positive"; - if (min_blocks_per_sm_.has_value()) { - TVM_FFI_ICHECK_EQ(min_blocks_per_sm_.value(), min_blocks_per_sm->value) - << "Conflicting " << tirx::attr::kLaunchBoundsMinBlocksPerSM << " values"; - } - min_blocks_per_sm_ = static_cast(min_blocks_per_sm->value); - return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); - } else if (op->attr_key == tirx::attr::kLaunchBoundsMaxBlocksPerCluster) { - const auto* max_blocks_per_cluster = op->value.as(); - TVM_FFI_ICHECK(max_blocks_per_cluster) - << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " expects an integer value"; - TVM_FFI_ICHECK_GT(max_blocks_per_cluster->value, 0) - << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " must be positive"; - if (max_blocks_per_cluster_.has_value()) { - TVM_FFI_ICHECK_EQ(max_blocks_per_cluster_.value(), max_blocks_per_cluster->value) - << "Conflicting " << tirx::attr::kLaunchBoundsMaxBlocksPerCluster << " values"; - } - max_blocks_per_cluster_ = static_cast(max_blocks_per_cluster->value); - return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); - } else if (op->attr_key == tirx::attr::kMaxRegisters) { - const auto* max_registers = op->value.as(); - TVM_FFI_ICHECK(max_registers) << tirx::attr::kMaxRegisters << " expects an integer value"; - TVM_FFI_ICHECK_GT(max_registers->value, 0) - << tirx::attr::kMaxRegisters << " must be positive"; - if (max_registers_.has_value()) { - TVM_FFI_ICHECK_EQ(max_registers_.value(), max_registers->value) - << "Conflicting " << tirx::attr::kMaxRegisters << " values"; - } - max_registers_ = static_cast(max_registers->value); - return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); - } else if (op->attr_key == tirx::attr::kRequiredBlockSize) { - const auto* required_block_size = op->value.as(); - TVM_FFI_ICHECK(required_block_size) - << tirx::attr::kRequiredBlockSize << " expects an integer value"; - TVM_FFI_ICHECK_EQ(required_block_size->value, 1) - << tirx::attr::kRequiredBlockSize << " must be 1"; - if (required_block_size_.has_value()) { - TVM_FFI_ICHECK_EQ(required_block_size_.value(), required_block_size->value) - << "Conflicting " << tirx::attr::kRequiredBlockSize << " values"; - } - required_block_size_ = static_cast(required_block_size->value); - return Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body); - } - return StmtExprMutator::Mutate_(op, inplace_mode); - } - - std::optional min_blocks_per_sm_; - std::optional max_blocks_per_cluster_; - std::optional max_registers_; - std::optional required_block_size_; -}; - class HostDeviceSplitter : public StmtExprMutator { public: using StmtExprMutator::Mutate; @@ -317,8 +228,6 @@ class HostDeviceSplitter : public StmtExprMutator { {})), std::move(body)); } - auto launch_bounds_attr = ffi::make_object(); - body = launch_bounds_attr->Extract(std::move(body)); PrimFunc device_func(kernel_params, body, kernel_ret_type); device_func = WithAttrs(std::move(device_func), {{tvm::attr::kTarget, device_target}, {tirx::attr::kNoAlias, true}, @@ -332,24 +241,6 @@ class HostDeviceSplitter : public StmtExprMutator { device_func = WithAttr(std::move(device_func), tirx::attr::kKernelLaunchParams, launch_params.value()); } - if (device_target->kind->name == "cuda") { - if (launch_bounds_attr->min_blocks_per_sm().has_value()) { - device_func = WithAttr(std::move(device_func), tirx::attr::kLaunchBoundsMinBlocksPerSM, - launch_bounds_attr->min_blocks_per_sm().value()); - } - if (launch_bounds_attr->max_blocks_per_cluster().has_value()) { - device_func = WithAttr(std::move(device_func), tirx::attr::kLaunchBoundsMaxBlocksPerCluster, - launch_bounds_attr->max_blocks_per_cluster().value()); - } - if (launch_bounds_attr->max_registers().has_value()) { - device_func = WithAttr(std::move(device_func), tirx::attr::kMaxRegisters, - launch_bounds_attr->max_registers().value()); - } - if (launch_bounds_attr->required_block_size().has_value()) { - device_func = WithAttr(std::move(device_func), tirx::attr::kRequiredBlockSize, - launch_bounds_attr->required_block_size().value()); - } - } auto num_inputs = cur_func_->GetAttr(tvm::attr::kNumInputs); if (num_inputs.has_value()) { device_func = WithAttr(std::move(device_func), tvm::attr::kNumInputs, num_inputs); @@ -441,8 +332,6 @@ class DeviceInfoCollector : public StmtExprVisitor { } } } - collector->use_required_block_dimension_ = - func->GetAttr(tirx::attr::kRequiredBlockSize).value_or(0) == 1; collector->Visit(func->body); @@ -533,7 +422,12 @@ class DeviceInfoCollector : public StmtExprVisitor { } ffi::Optional Visit_(const EvaluateNode* op) final { + static const Op required_block_size = Op::Get("tirx.cuda.required_block_size"); static const Op dyn_smem_bytes = Op::Get("tirx.cuda.dyn_smem_bytes"); + if (const auto* call = op->value.as(); + call && call->op.same_as(required_block_size)) { + use_required_block_dimension_ = true; + } if (const auto* call = op->value.as(); call && call->op.same_as(dyn_smem_bytes)) { // The declaration supplies the launch size even when the backing // shared.dyn allocation is an extern placeholder. diff --git a/tests/python/tirx-transform/test_tir_transform_split_host_device.py b/tests/python/tirx-transform/test_tir_transform_split_host_device.py index bd95398420fe..2c90bf735ed5 100644 --- a/tests/python/tirx-transform/test_tir_transform_split_host_device.py +++ b/tests/python/tirx-transform/test_tir_transform_split_host_device.py @@ -471,17 +471,17 @@ class Before: def main(A: T.Tensor(4, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) T.region("tirx.device_scope", [], attrs={"target": T.target("cuda")}) - T.attr(0, "tirx.required_block_size", 1) - with T.attr(0, "tirx.launch_bounds_min_blocks_per_sm", 1): - bx = T.launch_thread("blockIdx.x", 4) - tx = T.launch_thread("threadIdx.x", 128) - if tx == 0: - A[bx] = 0.0 + T.cuda.required_block_size(128, 1, 1, 1, 1, 1) + T.cuda.launch_bounds_min_blocks_per_sm(1) + bx = T.launch_thread("blockIdx.x", 4) + tx = T.launch_thread("threadIdx.x", 128) + if tx == 0: + A[bx] = 0.0 after = tvm.tirx.transform.SplitHostDevice()(Before) kernel = after["main_kernel"] - assert int(kernel.attrs["tirx.required_block_size"]) == 1 - assert int(kernel.attrs["tirx.launch_bounds_min_blocks_per_sm"]) == 1 + assert "T.cuda.required_block_size(128, 1, 1, 1, 1, 1)" in kernel.script() + assert "T.cuda.launch_bounds_min_blocks_per_sm(1)" in kernel.script() assert list(kernel.attrs["tirx.kernel_launch_params"]) == [ "blockIdx.x", "threadIdx.x", diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py b/tests/python/tirx/codegen/test_codegen_cuda.py index 6d6405377760..c84790f48edc 100644 --- a/tests/python/tirx/codegen/test_codegen_cuda.py +++ b/tests/python/tirx/codegen/test_codegen_cuda.py @@ -253,11 +253,11 @@ def main(A: T.Tensor((4,), "int32")): assert "__launch_bounds__(128, 1)" not in src -def test_tirx_launch_bounds_min_blocks_attr_sets_one_block_per_sm(): +def test_tirx_launch_bounds_min_blocks_sets_one_block_per_sm(): @T.prim_func def main(A: T.Tensor((4,), "int32")): T.device_entry() - T.attr({"tirx.launch_bounds_min_blocks_per_sm": 1}) + T.cuda.launch_bounds_min_blocks_per_sm(1) bx = T.cta_id([4]) tx = T.thread_id([128]) if tx == 0: @@ -265,19 +265,15 @@ def main(A: T.Tensor((4,), "int32")): src, _ = _get_source(main) assert 'extern "C" __global__ void __launch_bounds__(128, 1) main_kernel' in src - assert "tirx.launch_bounds_min_blocks_per_sm" not in src + assert "tirx.cuda.launch_bounds_min_blocks_per_sm" not in src def test_tirx_launch_bounds_max_blocks_per_cluster_emits_third_operand(): @T.prim_func def main(A: T.Tensor((4,), "int32")): T.device_entry() - T.attr( - { - "tirx.launch_bounds_min_blocks_per_sm": 1, - "tirx.launch_bounds_max_blocks_per_cluster": 1, - } - ) + T.cuda.launch_bounds_min_blocks_per_sm(1) + T.cuda.launch_bounds_max_blocks_per_cluster(1) bx = T.cta_id([4]) tx = T.thread_id([384]) if tx == 0: @@ -285,14 +281,14 @@ def main(A: T.Tensor((4,), "int32")): src, _ = _get_source(main) assert 'extern "C" __global__ void __launch_bounds__(384, 1, 1) main_kernel' in src - assert "tirx.launch_bounds_max_blocks_per_cluster" not in src + assert "tirx.cuda.launch_bounds_max_blocks_per_cluster" not in src -def test_tirx_max_registers_attr_emits_cuda_maxnreg(): +def test_tirx_max_registers_emits_cuda_maxnreg(): @T.prim_func def main(A: T.Tensor((4,), "int32")): T.device_entry() - T.attr({"tirx.max_registers": 92}) + T.cuda.max_registers_per_thread(92) bx = T.cta_id([4]) tx = T.thread_id([128]) if tx == 0: @@ -301,19 +297,15 @@ def main(A: T.Tensor((4,), "int32")): src, _ = _get_source(main) assert 'extern "C" __global__ void __maxnreg__(92) main_kernel' in src assert "__launch_bounds__" not in src - assert "tirx.max_registers" not in src + assert "tirx.cuda.max_registers_per_thread" not in src def test_tirx_max_registers_rejects_launch_bounds(): @T.prim_func def main(A: T.Tensor((4,), "int32")): T.device_entry() - T.attr( - { - "tirx.max_registers": 92, - "tirx.launch_bounds_min_blocks_per_sm": 1, - } - ) + T.cuda.max_registers_per_thread(92) + T.cuda.launch_bounds_min_blocks_per_sm(1) bx = T.cta_id([4]) tx = T.thread_id([128]) if tx == 0: @@ -327,7 +319,7 @@ def test_tirx_required_block_size_emits_cuda_block_size(): @T.prim_func def main(A: T.Tensor((8,), "int32")): T.device_entry() - T.attr({"tirx.required_block_size": 1}) + T.cuda.required_block_size(128, 1, 1, 1, 2, 1) bx, by = T.cta_id([4, 2]) _, cy = T.cta_id_in_cluster([1, 2]) tx = T.thread_id([128]) @@ -337,19 +329,15 @@ def main(A: T.Tensor((8,), "int32")): src, _ = _get_source(main) assert 'extern "C" __global__ void __block_size__((128, 1, 1), (1, 2, 1)) main_kernel' in src assert "__launch_bounds__" not in src - assert "tirx.required_block_size" not in src + assert "tirx.cuda.required_block_size" not in src def test_tirx_required_block_size_emits_launch_bounds_when_requested(): @T.prim_func def main(A: T.Tensor((4,), "int32")): T.device_entry() - T.attr( - { - "tirx.required_block_size": 1, - "tirx.launch_bounds_min_blocks_per_sm": 1, - } - ) + T.cuda.required_block_size(128, 1, 1, 1, 1, 1) + T.cuda.launch_bounds_min_blocks_per_sm(1) bx = T.cta_id([4]) tx = T.thread_id([128]) if tx == 0: @@ -360,8 +348,8 @@ def main(A: T.Tensor((4,), "int32")): 'extern "C" __global__ void __block_size__((128, 1, 1), (1, 1, 1)) ' "__launch_bounds__(128, 1) main_kernel" in src ) - assert "tirx.required_block_size" not in src - assert "tirx.launch_bounds_min_blocks_per_sm" not in src + assert "tirx.cuda.required_block_size" not in src + assert "tirx.cuda.launch_bounds_min_blocks_per_sm" not in src def test_tirx_cuda_kernel_return_zero_codegen_is_void_early_return():