From 527e71443cf72e1fba173e4bbaf60543c6802b7a Mon Sep 17 00:00:00 2001 From: Neroued Date: Sun, 13 Sep 2026 21:04:00 +0200 Subject: [PATCH] build: lower CUDA floor to 12.9, stage oversized GEMM smem dynamically Toolchains before CUDA 13 cap statically allocated shared memory at 48 KiB per block. W8 and NVFP4 GEMM kernels whose staged storage (activations, weight codes, scales) exceeds the cap now stage it in dynamic shared memory, sized through the matching dynamic-bytes helpers, and every launch site raises cudaFuncAttributeMaxDynamicSharedMemorySize when the storage exceeds the cap. Kernels whose storage fits the cap keep their static allocation. --- CMakeLists.txt | 4 +-- .../nvfp4/nvfp4_attn_input_w4a4.cu | 17 +++++++--- .../w8/w8_attn_input_gemm_mma.cu | 16 ++++++--- .../w8/w8_attn_input_gemm_splitk.cu | 33 +++++++++++++++++-- .../w8/w8_dflash2_attn_input.cu | 20 +++++++++-- ...8_dynamic_grouped_conv_add_materialized.cu | 9 ++--- .../nvfp4/nvfp4_gdn_input_w4a4.cu | 19 ++++++++--- .../w8/w8_gdn_input_gemm_mma.cu | 16 ++++++--- .../w8/w8_gdn_input_gemm_splitk.cu | 22 +++++++++++-- src/ops/linear/nvfp4/nvfp4_w4a4.cu | 16 ++++++--- src/ops/linear/nvfp4/nvfp4_w4a4_mma.cuh | 19 ++++++++++- src/ops/linear/w8/w8_config.h | 6 ++++ src/ops/linear/w8/w8_feature.cu | 28 +++++++++++++--- .../w8/w8_rowsplit_gemm_medium_t_splitk.cuh | 32 ++++++++++++++++-- src/ops/linear/w8/w8_rowsplit_gemm_mma.cu | 14 ++++++-- src/ops/linear/w8/w8_rowsplit_gemm_mma.cuh | 33 ++++++++++++++++--- src/ops/linear/w8/w8_rowsplit_gemm_splitk.cu | 22 +++++++++++-- src/ops/linear/w8/w8_small_t.cu | 22 +++++++++++-- src/ops/linear/w8/w8_small_t_mma.cuh | 18 ++++++++-- .../linear_add/nvfp4/nvfp4_linear_add_w4a4.cu | 19 ++++++++--- .../linear_add/w8/w8_linear_add_gemm_mma.cu | 9 ++++- .../w8/w8_linear_add_gemm_splitk.cu | 22 +++++++++++-- src/ops/linear_pair/w8/w8_pair_gemm_concat.cu | 9 ++++- src/ops/linear_pair/w8/w8_pair_gemm_splitk.cu | 27 ++++++++++++--- .../nvfp4/nvfp4_linear_swiglu_w4a4.cu | 17 +++++++--- .../w8/w8_dflash2_linear_swiglu.cu | 11 ++++++- .../w8/w8_linear_swiglu_gemm_mma.cu | 16 ++++++--- .../w8/w8_linear_swiglu_gemm_splitk.cu | 12 ++++++- 28 files changed, 429 insertions(+), 79 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 98ac0bbd65..9e40aa94bf 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -33,9 +33,9 @@ if(CMAKE_GENERATOR MATCHES "Ninja") set(CMAKE_JOB_POOL_LINK cuda_link) endif() -if(CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 13.1) +if(CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 12.9) message(FATAL_ERROR - "NInfer requires CUDA 13.1 or newer; found CUDA compiler " + "NInfer requires CUDA 12.9 or newer; found CUDA compiler " "${CMAKE_CUDA_COMPILER_VERSION}") endif() diff --git a/src/ops/attn_input_proj/nvfp4/nvfp4_attn_input_w4a4.cu b/src/ops/attn_input_proj/nvfp4/nvfp4_attn_input_w4a4.cu index 56835368cf..b1cf913e5c 100644 --- a/src/ops/attn_input_proj/nvfp4/nvfp4_attn_input_w4a4.cu +++ b/src/ops/attn_input_proj/nvfp4/nvfp4_attn_input_w4a4.cu @@ -71,10 +71,19 @@ void launch_gemm(const Weight& weight, Tensor& q, Tensor& gate, Tensor& k, Tenso static_cast<__nv_bfloat16*>(v.data), }; const float alpha = 1.0F / (weight.input_scale_divisor * weight.weight_scale_divisor); - nvfp4_w4a4_mma_kernel<<>>( - activation, static_cast(weight.qdata), - static_cast(weight.scales), tokens, alpha, Nvfp4IdentityEpilogue{}, - output); + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kNvfp4W4a4StaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + nvfp4_w4a4_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + nvfp4_w4a4_mma_kernel + <<>>( + activation, static_cast(weight.qdata), + static_cast(weight.scales), tokens, alpha, + Nvfp4IdentityEpilogue{}, output); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/attn_input_proj/w8/w8_attn_input_gemm_mma.cu b/src/ops/attn_input_proj/w8/w8_attn_input_gemm_mma.cu index 2a4be492df..51b6232011 100644 --- a/src/ops/attn_input_proj/w8/w8_attn_input_gemm_mma.cu +++ b/src/ops/attn_input_proj/w8/w8_attn_input_gemm_mma.cu @@ -16,11 +16,19 @@ using CompanionOutput = W8SplitOutput3<4096, 1024, 1024>; template void launch_variant(const Tensor& x, const Weight& weight, Output output, cudaStream_t stream) { const dim3 grid(Rows / Schedule::BM, static_cast(div_up(x.ne[1], Schedule::BN)), 1u); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>(static_cast(x.data), - static_cast(weight.qdata), - static_cast(weight.scales), - output, Rows, kHidden, x.ne[1], kHidden); + <<>>( + static_cast(x.data), + static_cast(weight.qdata), + static_cast(weight.scales), output, Rows, kHidden, x.ne[1], + kHidden); } template diff --git a/src/ops/attn_input_proj/w8/w8_attn_input_gemm_splitk.cu b/src/ops/attn_input_proj/w8/w8_attn_input_gemm_splitk.cu index 20776a1d3e..05ff6ca7e1 100644 --- a/src/ops/attn_input_proj/w8/w8_attn_input_gemm_splitk.cu +++ b/src/ops/attn_input_proj/w8/w8_attn_input_gemm_splitk.cu @@ -36,8 +36,17 @@ void launch_output(const Tensor& x, const Weight& weight, Output output, cudaStr : 48; using Geometry = W8LinearGeometry; using Schedule = W8SmallTMmaDefaultSchedule; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output); @@ -87,8 +96,17 @@ void launch_target_medium_cols(const Tensor& x, const Weight& weight, Tensor& q, const TargetOutput output{ static_cast<__nv_bfloat16*>(q.data), static_cast<__nv_bfloat16*>(k.data), static_cast<__nv_bfloat16*>(gate.data), static_cast<__nv_bfloat16*>(v.data)}; + constexpr std::size_t kDynamicBytes = + w8_rowsplit_medium_t_splitk_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_medium_t_splitk_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_medium_t_splitk_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, x.ne[1]); @@ -101,8 +119,17 @@ void launch_companion_medium_cols(const Tensor& x, const Weight& weight, Tensor& const CompanionOutput output{static_cast<__nv_bfloat16*>(q.data), static_cast<__nv_bfloat16*>(k.data), static_cast<__nv_bfloat16*>(v.data)}; + constexpr std::size_t kDynamicBytes = + w8_rowsplit_medium_t_splitk_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_medium_t_splitk_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_medium_t_splitk_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, x.ne[1]); diff --git a/src/ops/attn_input_proj/w8/w8_dflash2_attn_input.cu b/src/ops/attn_input_proj/w8/w8_dflash2_attn_input.cu index 4c41bf7361..9d53d34e03 100644 --- a/src/ops/attn_input_proj/w8/w8_dflash2_attn_input.cu +++ b/src/ops/attn_input_proj/w8/w8_dflash2_attn_input.cu @@ -42,9 +42,18 @@ void launch_small(const Tensor& x, const Weight& weight, Tensor& q, Tensor& k, T const Output output{static_cast<__nv_bfloat16*>(q.data), static_cast<__nv_bfloat16*>(k.data), static_cast<__nv_bfloat16*>(v.data)}; constexpr int kBlocks = Geometry::kOutputRows / Schedule::kRowsPerCta; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, W8SmallTMmaStoreEpilogue{}, @@ -69,8 +78,15 @@ void launch_mma_slice(const Tensor& x, const Weight& weight, Tensor& q, Tensor& static_cast<__nv_bfloat16*>(v.data)}; const dim3 grid(Geometry::kOutputRows / Schedule::BM, static_cast(div_up(x.ne[1], Schedule::BN)), 1u); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, Geometry::kOutputRows, diff --git a/src/ops/dynamic_grouped_conv/w8/w8_dynamic_grouped_conv_add_materialized.cu b/src/ops/dynamic_grouped_conv/w8/w8_dynamic_grouped_conv_add_materialized.cu index 992ea7f899..a0a3aa81c0 100644 --- a/src/ops/dynamic_grouped_conv/w8/w8_dynamic_grouped_conv_add_materialized.cu +++ b/src/ops/dynamic_grouped_conv/w8/w8_dynamic_grouped_conv_add_materialized.cu @@ -40,12 +40,13 @@ void tiled_projection(const Tensor& x, const Weight& weight, Tensor& out, cudaSt using Geometry = W8LinearGeometry; using Schedule = W8SmallTMmaSchedule; - constexpr int SharedBytes = TileColumns > 64 ? sizeof(W8SmallTMmaSharedStorage) : 0; - if constexpr (SharedBytes > 0) { + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { static const cudaError_t attribute = cudaFuncSetAttribute( w8_small_t_mma_kernel, - cudaFuncAttributeMaxDynamicSharedMemorySize, SharedBytes); + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); CUDA_CHECK(attribute); } const int columns = x.ne[1]; @@ -53,7 +54,7 @@ void tiled_projection(const Tensor& x, const Weight& weight, Tensor& out, cudaSt const dim3 grid(kRows / 16, (columns + TileColumns - 1) / TileColumns); w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, W8SmallTMmaStoreEpilogue{}, diff --git a/src/ops/gdn_input_proj/nvfp4/nvfp4_gdn_input_w4a4.cu b/src/ops/gdn_input_proj/nvfp4/nvfp4_gdn_input_w4a4.cu index f56a4cd0fe..ca9193327b 100644 --- a/src/ops/gdn_input_proj/nvfp4/nvfp4_gdn_input_w4a4.cu +++ b/src/ops/gdn_input_proj/nvfp4/nvfp4_gdn_input_w4a4.cu @@ -25,11 +25,20 @@ void launch_gemm(const Weight& weight, Tensor& qkv, Tensor& z, Nvfp4W4a4Workspac (tokens + Schedule::kBlockM - 1) / Schedule::kBlockM); const Nvfp4W4a4MaterializedActivation activation{workspace.codes, workspace.scales}; const float alpha = 1.0F / (weight.input_scale_divisor * weight.weight_scale_divisor); - nvfp4_w4a4_mma_kernel<<>>( - activation, static_cast(weight.qdata), - static_cast(weight.scales), tokens, alpha, Nvfp4IdentityEpilogue{}, - Nvfp4GdnInputOutput{static_cast<__nv_bfloat16*>(qkv.data), - static_cast<__nv_bfloat16*>(z.data)}); + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kNvfp4W4a4StaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + nvfp4_w4a4_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + nvfp4_w4a4_mma_kernel + <<>>( + activation, static_cast(weight.qdata), + static_cast(weight.scales), tokens, alpha, + Nvfp4IdentityEpilogue{}, + Nvfp4GdnInputOutput{static_cast<__nv_bfloat16*>(qkv.data), + static_cast<__nv_bfloat16*>(z.data)}); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_mma.cu b/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_mma.cu index 9737cb30cb..03739d8c99 100644 --- a/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_mma.cu +++ b/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_mma.cu @@ -18,11 +18,19 @@ void launch_variant(const Tensor& x, const Weight& weight, Tensor& qkv, Tensor& static_assert((8192 % Schedule::BM) == 0 && (4096 % Schedule::BM) == 0); const Output output{static_cast<__nv_bfloat16*>(qkv.data), static_cast<__nv_bfloat16*>(z.data)}; const dim3 grid(kRows / Schedule::BM, static_cast(div_up(x.ne[1], Schedule::BN)), 1u); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>(static_cast(x.data), - static_cast(weight.qdata), - static_cast(weight.scales), - output, kRows, kHidden, x.ne[1], kHidden); + <<>>( + static_cast(x.data), + static_cast(weight.qdata), + static_cast(weight.scales), output, kRows, kHidden, x.ne[1], + kHidden); } } // namespace diff --git a/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_splitk.cu b/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_splitk.cu index f75f68bef8..d2729dd62c 100644 --- a/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_splitk.cu +++ b/src/ops/gdn_input_proj/w8/w8_gdn_input_gemm_splitk.cu @@ -281,8 +281,17 @@ void launch_active_cols(const Tensor& x, const Weight& weight, Tensor& qkv, Tens using Schedule = W8SmallTMmaDefaultSchedule; static_assert((8192 % kRowsPerCta) == 0 && (4096 % kRowsPerCta) == 0); const Output output{static_cast<__nv_bfloat16*>(qkv.data), static_cast<__nv_bfloat16*>(z.data)}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output); @@ -320,8 +329,17 @@ void launch_active_cols_conv(const Tensor& x, const Weight& weight, const Tensor }, static_cast<__nv_bfloat16*>(z.data), }; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, W8SmallTMmaIdentityRows>, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel> - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), ignored_output, epilogue); diff --git a/src/ops/linear/nvfp4/nvfp4_w4a4.cu b/src/ops/linear/nvfp4/nvfp4_w4a4.cu index 156871fbbe..e575199c79 100644 --- a/src/ops/linear/nvfp4/nvfp4_w4a4.cu +++ b/src/ops/linear/nvfp4/nvfp4_w4a4.cu @@ -29,10 +29,18 @@ void launch_gemm(const Weight& weight, Tensor& out, Nvfp4W4a4Workspace workspace const Nvfp4ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), Geometry::kOutputRows}; const float alpha = 1.0F / (weight.input_scale_divisor * weight.weight_scale_divisor); - nvfp4_w4a4_mma_kernel<<>>( - activation, static_cast(weight.qdata), - static_cast(weight.scales), tokens, alpha, Nvfp4IdentityEpilogue{}, - output); + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kNvfp4W4a4StaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + nvfp4_w4a4_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + nvfp4_w4a4_mma_kernel + <<>>( + activation, static_cast(weight.qdata), + static_cast(weight.scales), tokens, alpha, + Nvfp4IdentityEpilogue{}, output); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/linear/nvfp4/nvfp4_w4a4_mma.cuh b/src/ops/linear/nvfp4/nvfp4_w4a4_mma.cuh index 675df5b992..8bdfe1f1d7 100644 --- a/src/ops/linear/nvfp4/nvfp4_w4a4_mma.cuh +++ b/src/ops/linear/nvfp4/nvfp4_w4a4_mma.cuh @@ -67,6 +67,17 @@ struct Nvfp4W4a4SharedStorage { std::uint8_t b_scales[Schedule::kStages][Schedule::kBlockN * Schedule::kK64PerStage * 4]; }; +// Toolchains before CUDA 13 cap statically allocated shared memory at 48 KiB. An instantiation whose +// staged storage exceeds that cap stages it in dynamic shared memory, so every launch site requests +// the size reported here and raises the block limit before launching. +inline constexpr std::size_t kNvfp4W4a4StaticSharedBytes = 48 * 1024; + +template +__host__ __device__ constexpr std::size_t nvfp4_w4a4_mma_dynamic_bytes() { + constexpr std::size_t kStorageBytes = sizeof(Nvfp4W4a4SharedStorage); + return kStorageBytes > kNvfp4W4a4StaticSharedBytes ? kStorageBytes : 0; +} + template __device__ __forceinline__ int nvfp4_w4a4_swizzled_byte(int row, int logical_byte) { static_assert((Schedule::kSegmentsPerRow & (Schedule::kSegmentsPerRow - 1)) == 0); @@ -213,7 +224,13 @@ __launch_bounds__(Schedule::kThreads, Schedule::kMinBlocksPerSm) void nvfp4_w4a4 static_assert(!PairRows || (Schedule::kBlockN % 2) == 0); static_assert(!PairRows || ((Geometry::kOutputRows / 2) % (Schedule::kBlockN / 2)) == 0); - __shared__ Nvfp4W4a4SharedStorage shared; + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + constexpr bool kDynamicShared = kDynamicBytes != 0; + __shared__ __align__(16) unsigned char + static_shared[kDynamicShared ? 1 : sizeof(Nvfp4W4a4SharedStorage)]; + extern __shared__ __align__(16) unsigned char dynamic_shared[]; + auto& shared = *reinterpret_cast*>( + kDynamicShared ? dynamic_shared : static_shared); const int token_begin = static_cast(blockIdx.y) * Schedule::kBlockM; constexpr int kRowsPerBlock = PairRows ? Schedule::kBlockN / 2 : Schedule::kBlockN; const int row_begin = static_cast(blockIdx.x) * kRowsPerBlock; diff --git a/src/ops/linear/w8/w8_config.h b/src/ops/linear/w8/w8_config.h index b4cc70aed6..8bda02f558 100644 --- a/src/ops/linear/w8/w8_config.h +++ b/src/ops/linear/w8/w8_config.h @@ -2,10 +2,16 @@ #include "ops/common/memory.cuh" +#include #include namespace ninfer::ops::detail { +// Toolchains before CUDA 13 cap statically allocated shared memory at 48 KiB. A W8 kernel whose +// staged storage exceeds that cap stages it in dynamic shared memory instead: the kernel and every +// launch site size it through the matching dynamic-bytes helper. +inline constexpr std::size_t kW8SmallTMmaStaticSharedBytes = 48 * 1024; + enum class W8SmallTMmaScaleAccess : std::uint8_t { Direct, Shared, diff --git a/src/ops/linear/w8/w8_feature.cu b/src/ops/linear/w8/w8_feature.cu index 670582a3f6..b45d41365d 100644 --- a/src/ops/linear/w8/w8_feature.cu +++ b/src/ops/linear/w8/w8_feature.cu @@ -16,9 +16,18 @@ void launch_small(const Tensor& x, const Weight& weight, Tensor& out, cudaStream using Schedule = W8SmallTMmaSchedule; const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), Geometry::kOutputRows}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, W8SmallTMmaStoreEpilogue{}, @@ -42,10 +51,19 @@ void launch(const Tensor& x, const Weight& weight, Tensor& out, cudaStream_t str using Schedule = W8RowSplitMmaGemmSchedule; const dim3 grid(weight.n / Rows, (x.ne[1] + 63) / 64); const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), weight.n}; - w8_rowsplit_gemm_mma_kernel<<>>( - static_cast(x.data), static_cast(weight.qdata), - static_cast(weight.scales), output, weight.n, weight.k, x.ne[1], - weight.padded_shape[1]); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + w8_rowsplit_gemm_mma_kernel + <<>>( + static_cast(x.data), + static_cast(weight.qdata), + static_cast(weight.scales), output, weight.n, weight.k, x.ne[1], + weight.padded_shape[1]); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/linear/w8/w8_rowsplit_gemm_medium_t_splitk.cuh b/src/ops/linear/w8/w8_rowsplit_gemm_medium_t_splitk.cuh index fe9c3aa023..dbc47379ee 100644 --- a/src/ops/linear/w8/w8_rowsplit_gemm_medium_t_splitk.cuh +++ b/src/ops/linear/w8/w8_rowsplit_gemm_medium_t_splitk.cuh @@ -13,6 +13,27 @@ namespace ninfer::ops::detail { +// Staged weight codes and activations for one medium-T split-K tile. The layout matches the separate +// shared arrays this core used while the staged storage fits the static cap; both the kernel and its +// launch sites size it through the helper below. +template +struct W8RowSplitMediumTSplitKSharedStorage { + static constexpr int kTileK = 64; + static constexpr int kMmaRows = 16; + static constexpr int kGroupK = KSplits * kTileK; + static constexpr int kWarpCols = TileCols / NGroups; + + alignas(16) std::uint8_t code_shared[kMmaRows][kGroupK]; + alignas(16) __nv_bfloat16 b_shared[KSplits * NGroups][kWarpCols * kTileK]; +}; + +template +__host__ __device__ constexpr std::size_t w8_rowsplit_medium_t_splitk_dynamic_bytes() { + constexpr std::size_t kStorageBytes = + sizeof(W8RowSplitMediumTSplitKSharedStorage); + return kStorageBytes > kW8SmallTMmaStaticSharedBytes ? kStorageBytes : 0; +} + template __global__ @@ -32,8 +53,15 @@ __launch_bounds__(KSplits* NGroups * 32, MinBlocks) void w8_rowsplit_medium_t_sp static_assert(TileCols % NGroups == 0 && kWarpCols % 8 == 0); static_assert(Hidden % kGroupK == 0 && kKernelWarps <= 32); - __shared__ __align__(16) std::uint8_t code_shared[kMmaRows][kGroupK]; - __shared__ __align__(16) __nv_bfloat16 b_shared[kKernelWarps][kWarpCols * kTileK]; + using SharedStorage = W8RowSplitMediumTSplitKSharedStorage; + constexpr std::size_t kStorageBytes = sizeof(SharedStorage); + constexpr bool kDynamicShared = kStorageBytes > kW8SmallTMmaStaticSharedBytes; + __shared__ __align__(16) unsigned char static_shared[kDynamicShared ? 1 : kStorageBytes]; + extern __shared__ __align__(16) unsigned char dynamic_shared[]; + auto& shared = + *reinterpret_cast(kDynamicShared ? dynamic_shared : static_shared); + auto& code_shared = shared.code_shared; + auto& b_shared = shared.b_shared; const int tid = static_cast(threadIdx.x); const int warp = tid >> 5; diff --git a/src/ops/linear/w8/w8_rowsplit_gemm_mma.cu b/src/ops/linear/w8/w8_rowsplit_gemm_mma.cu index 0d41ca96a3..ebdb145052 100644 --- a/src/ops/linear/w8/w8_rowsplit_gemm_mma.cu +++ b/src/ops/linear/w8/w8_rowsplit_gemm_mma.cu @@ -20,9 +20,17 @@ void launch_slice(const Tensor& x, const Weight& w, Tensor& out, cudaStream_t st const dim3 grid(static_cast(div_up(rows, Schedule::BM)), static_cast(div_up(cols, Schedule::BN)), 1u); const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), rows}; - w8_rowsplit_gemm_mma_kernel<<>>( - static_cast(x.data), static_cast(w.qdata), - static_cast(w.scales), output, rows, k, cols, padded_k); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + w8_rowsplit_gemm_mma_kernel + <<>>( + static_cast(x.data), static_cast(w.qdata), + static_cast(w.scales), output, rows, k, cols, padded_k); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/linear/w8/w8_rowsplit_gemm_mma.cuh b/src/ops/linear/w8/w8_rowsplit_gemm_mma.cuh index 13e555e1b1..e002a8ede6 100644 --- a/src/ops/linear/w8/w8_rowsplit_gemm_mma.cuh +++ b/src/ops/linear/w8/w8_rowsplit_gemm_mma.cuh @@ -10,6 +10,7 @@ #include "ops/common/mma.cuh" #include "ops/common/math.cuh" +#include "ops/linear/w8/w8_config.h" #include "ops/linear/w8/w8_rowsplit_output.cuh" #include @@ -62,6 +63,23 @@ __device__ __forceinline__ int w8g32_swz64(int row, int col) { return (((col >> 3) ^ (row & 7)) << 3) | (col & 7); } +// Staged activations, weight codes and scales for one row-split MMA tile. The layout matches the +// separate shared arrays this core used while the staged storage fits the static cap; both the +// kernel and its launch sites size it through the helper below. +template +struct W8RowSplitMmaGemmSharedStorage { + alignas(16) __nv_bfloat16 As[Cfg::BM * Cfg::BK]; + alignas(16) __nv_bfloat16 Bs[Cfg::ACTIVATION_STAGES][Cfg::BN * Cfg::BK]; + alignas(16) std::uint8_t Cr[Cfg::BM * Cfg::BK]; + alignas(16) std::uint8_t Sr[Cfg::BM * Cfg::SCALE_CACHE_BYTES]; +}; + +template +__host__ __device__ constexpr std::size_t w8_rowsplit_gemm_mma_dynamic_bytes() { + constexpr std::size_t kStorageBytes = sizeof(W8RowSplitMmaGemmSharedStorage); + return kStorageBytes > kW8SmallTMmaStaticSharedBytes ? kStorageBytes : 0; +} + template __global__ __launch_bounds__(Cfg::THREADS, Cfg::MIN_BLOCKS) void w8_rowsplit_gemm_mma_kernel( @@ -82,10 +100,17 @@ __global__ __launch_bounds__(Cfg::THREADS, Cfg::MIN_BLOCKS) void w8_rowsplit_gem static_assert(!kSwiGlu || Cfg::WARPS_M == 1 || Cfg::WARPS_M == 2, "SwiGLU supports warp-local or shared-memory row pairing"); - __shared__ __align__(16) __nv_bfloat16 As[BM * BK]; - __shared__ __align__(16) __nv_bfloat16 Bs[Cfg::ACTIVATION_STAGES][BN * BK]; - __shared__ __align__(16) std::uint8_t Cr[BM * BK]; - __shared__ __align__(16) std::uint8_t Sr[BM * Cfg::SCALE_CACHE_BYTES]; + using SharedStorage = W8RowSplitMmaGemmSharedStorage; + constexpr std::size_t kStorageBytes = sizeof(SharedStorage); + constexpr bool kDynamicShared = kStorageBytes > kW8SmallTMmaStaticSharedBytes; + __shared__ __align__(16) unsigned char static_shared[kDynamicShared ? 1 : kStorageBytes]; + extern __shared__ __align__(16) unsigned char dynamic_shared[]; + auto& shared = + *reinterpret_cast(kDynamicShared ? dynamic_shared : static_shared); + auto& As = shared.As; + auto& Bs = shared.Bs; + auto& Cr = shared.Cr; + auto& Sr = shared.Sr; const int tid = static_cast(threadIdx.x); const int warp = tid >> 5; diff --git a/src/ops/linear/w8/w8_rowsplit_gemm_splitk.cu b/src/ops/linear/w8/w8_rowsplit_gemm_splitk.cu index a9ee08ab48..a01c45889c 100644 --- a/src/ops/linear/w8/w8_rowsplit_gemm_splitk.cu +++ b/src/ops/linear/w8/w8_rowsplit_gemm_splitk.cu @@ -39,8 +39,17 @@ void launch_active_cols(const Tensor& x, const Weight& weight, Tensor& out, cuda using Schedule = W8SmallTMmaSchedule; static_assert((kRows % kRowsPerCta) == 0); const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), kRows}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output); @@ -65,8 +74,17 @@ void require_problem(const Tensor& x, const Weight& w, const Tensor& out) { template void launch_medium(const Tensor& x, const Weight& w, Tensor& out, cudaStream_t stream) { const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), kRows}; + constexpr std::size_t kDynamicBytes = + w8_rowsplit_medium_t_splitk_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_medium_t_splitk_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_medium_t_splitk_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(w.qdata), static_cast(w.scales), output, x.ne[1]); } diff --git a/src/ops/linear/w8/w8_small_t.cu b/src/ops/linear/w8/w8_small_t.cu index 22edb78ae5..d8ca5ef72f 100644 --- a/src/ops/linear/w8/w8_small_t.cu +++ b/src/ops/linear/w8/w8_small_t.cu @@ -20,8 +20,17 @@ void launch_exact(const Tensor& x, const Weight& weight, Tensor& out, cudaStream const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), Geometry::kOutputRows}; constexpr int kBlocks = Geometry::kOutputRows / Schedule::kRowsPerCta; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output); @@ -40,9 +49,18 @@ void launch_vocabulary_tile(const Tensor& x, const Weight& weight, Tensor& out, using Schedule = W8SmallTMmaSchedule; const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), Geometry::kOutputRows}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, W8SmallTMmaStoreEpilogue{}, diff --git a/src/ops/linear/w8/w8_small_t_mma.cuh b/src/ops/linear/w8/w8_small_t_mma.cuh index 3148271234..1dd404bbb8 100644 --- a/src/ops/linear/w8/w8_small_t_mma.cuh +++ b/src/ops/linear/w8/w8_small_t_mma.cuh @@ -68,6 +68,17 @@ union alignas(16) W8SmallTMmaSharedStorage { float partial[Schedule::kKWarps * (Schedule::kTileTokens / 8) * 32 * 4]; }; +// Toolchains before CUDA 13 cap statically allocated shared memory at 48 KiB. An instantiation whose +// staged storage exceeds that cap stages it in dynamic shared memory, so every launch site requests +// the size reported here and raises the block limit before launching. +template +__host__ __device__ constexpr std::size_t w8_small_t_mma_dynamic_bytes() { + constexpr std::size_t kStorageBytes = sizeof(W8SmallTMmaSharedStorage); + return kStorageBytes > kW8SmallTMmaStaticSharedBytes || (TiledColumns && ActiveCols > 64) + ? kStorageBytes + : 0; +} + struct W8SmallTMmaIdentityColumns { __device__ __forceinline__ int operator()(int column) const { return column; } }; @@ -99,9 +110,10 @@ w8_small_t_mma(const __nv_bfloat16* __restrict__ x, const std::uint8_t* __restri using SharedStorage = W8SmallTMmaSharedStorage; - constexpr bool kDynamicShared = TiledColumns && ActiveCols > 64; - __shared__ __align__( - 16) unsigned char static_shared[kDynamicShared ? 1 : sizeof(SharedStorage)]; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + constexpr bool kDynamicShared = kDynamicBytes != 0; + __shared__ __align__(16) unsigned char static_shared[kDynamicShared ? 1 : sizeof(SharedStorage)]; extern __shared__ __align__(16) unsigned char dynamic_shared[]; auto& shared = *reinterpret_cast(kDynamicShared ? dynamic_shared : static_shared); diff --git a/src/ops/linear_add/nvfp4/nvfp4_linear_add_w4a4.cu b/src/ops/linear_add/nvfp4/nvfp4_linear_add_w4a4.cu index 4dc44ba376..38cbf29291 100644 --- a/src/ops/linear_add/nvfp4/nvfp4_linear_add_w4a4.cu +++ b/src/ops/linear_add/nvfp4/nvfp4_linear_add_w4a4.cu @@ -26,11 +26,20 @@ void launch_gemm(const Weight& weight, Tensor& residual, Nvfp4W4a4Workspace work const Nvfp4W4a4MaterializedActivation activation{workspace.codes, workspace.scales}; auto* output = static_cast<__nv_bfloat16*>(residual.data); const float alpha = 1.0F / (weight.input_scale_divisor * weight.weight_scale_divisor); - nvfp4_w4a4_mma_kernel<<>>( - activation, static_cast(weight.qdata), - static_cast(weight.scales), tokens, alpha, - Nvfp4AddResidualEpilogue{output, Geometry::kOutputRows}, - Nvfp4ContiguousOutput{output, Geometry::kOutputRows}); + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kNvfp4W4a4StaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + nvfp4_w4a4_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } + nvfp4_w4a4_mma_kernel + <<>>( + activation, static_cast(weight.qdata), + static_cast(weight.scales), tokens, alpha, + Nvfp4AddResidualEpilogue{output, Geometry::kOutputRows}, + Nvfp4ContiguousOutput{output, Geometry::kOutputRows}); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/linear_add/w8/w8_linear_add_gemm_mma.cu b/src/ops/linear_add/w8/w8_linear_add_gemm_mma.cu index 7afa32fa42..732a1a1f1f 100644 --- a/src/ops/linear_add/w8/w8_linear_add_gemm_mma.cu +++ b/src/ops/linear_add/w8/w8_linear_add_gemm_mma.cu @@ -18,8 +18,15 @@ void launch_tt(const Tensor& x, const Weight& w, Tensor& residual_out, cudaStrea const dim3 grid(static_cast(div_up(rows, Schedule::BM)), static_cast(div_up(cols, Schedule::BN)), 1u); const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(residual_out.data), rows}; + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(w.qdata), static_cast(w.scales), output, rows, k, cols, padded_k); } diff --git a/src/ops/linear_add/w8/w8_linear_add_gemm_splitk.cu b/src/ops/linear_add/w8/w8_linear_add_gemm_splitk.cu index eb8df897b8..ef1c0131aa 100644 --- a/src/ops/linear_add/w8/w8_linear_add_gemm_splitk.cu +++ b/src/ops/linear_add/w8/w8_linear_add_gemm_splitk.cu @@ -40,9 +40,18 @@ void launch_active_cols(const Tensor& x, const Weight& weight, Tensor& residual_ static_assert((kRows % kRowsPerCta) == 0); auto* residual = static_cast<__nv_bfloat16*>(residual_out.data); const W8ContiguousOutput output{residual, kRows}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, W8SmallTMmaResidualEpilogue{}); @@ -63,9 +72,18 @@ template void launch_medium(const Tensor& x, Tensor& residual_out, const Weight& weight, cudaStream_t stream) { const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(residual_out.data), kRows}; + constexpr std::size_t kDynamicBytes = + w8_rowsplit_medium_t_splitk_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_medium_t_splitk_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_medium_t_splitk_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), output, x.ne[1]); diff --git a/src/ops/linear_pair/w8/w8_pair_gemm_concat.cu b/src/ops/linear_pair/w8/w8_pair_gemm_concat.cu index ed7e8a5617..447c20ffc0 100644 --- a/src/ops/linear_pair/w8/w8_pair_gemm_concat.cu +++ b/src/ops/linear_pair/w8/w8_pair_gemm_concat.cu @@ -53,8 +53,15 @@ void launch_variant(const Tensor& x, const Weight& first_weight, Tensor& first_o static_cast<__nv_bfloat16*>(second_out.data)}; const dim3 grid(static_cast(2 * div_up(kRows, Schedule::BM)), static_cast(div_up(x.ne[1], Schedule::BN)), 1u); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(first_weight.qdata), static_cast(first_weight.scales), output, 2 * kRows, kHidden, diff --git a/src/ops/linear_pair/w8/w8_pair_gemm_splitk.cu b/src/ops/linear_pair/w8/w8_pair_gemm_splitk.cu index b53a31b770..ebe051dfbb 100644 --- a/src/ops/linear_pair/w8/w8_pair_gemm_splitk.cu +++ b/src/ops/linear_pair/w8/w8_pair_gemm_splitk.cu @@ -74,10 +74,20 @@ void launch_active_cols(const Tensor& x, const Weight& first_weight, const Weigh const W8ContiguousOutput ignored{static_cast<__nv_bfloat16*>(first_out.data), kRows}; const W8PairExactTEpilogue epilogue{static_cast<__nv_bfloat16*>(first_out.data), static_cast<__nv_bfloat16*>(second_out.data)}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel<<>>( - static_cast(x.data), first_codes, first_scales, ignored, epilogue, - W8PairExactTRows{}); + W8PairExactTRows> + <<>>( + static_cast(x.data), first_codes, first_scales, ignored, + epilogue, W8PairExactTRows{}); } template @@ -101,8 +111,17 @@ void launch_medium(const Tensor& x, const Weight& first_weight, const Weight& se } const PairOutput output{static_cast<__nv_bfloat16*>(first_out.data), static_cast<__nv_bfloat16*>(second_out.data)}; + constexpr std::size_t kDynamicBytes = + w8_rowsplit_medium_t_splitk_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_medium_t_splitk_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_medium_t_splitk_kernel - <<<(2 * kRows) / 16, KSplits * NGroups * 32, 0, stream>>>( + <<<(2 * kRows) / 16, KSplits * NGroups * 32, kDynamicBytes, stream>>>( static_cast(x.data), first_codes, first_scales, output, x.ne[1]); } diff --git a/src/ops/linear_swiglu/nvfp4/nvfp4_linear_swiglu_w4a4.cu b/src/ops/linear_swiglu/nvfp4/nvfp4_linear_swiglu_w4a4.cu index 13e1b60528..a09ac8eb55 100644 --- a/src/ops/linear_swiglu/nvfp4/nvfp4_linear_swiglu_w4a4.cu +++ b/src/ops/linear_swiglu/nvfp4/nvfp4_linear_swiglu_w4a4.cu @@ -71,11 +71,20 @@ void launch_gemm(const Weight& weight, Tensor& out, Nvfp4W4a4Workspace workspace const Nvfp4SwiGluRows row_policy{}; const Nvfp4SwiGluOutput output{static_cast<__nv_bfloat16*>(out.data)}; const float alpha = 1.0F / (weight.input_scale_divisor * weight.weight_scale_divisor); + constexpr std::size_t kDynamicBytes = nvfp4_w4a4_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kNvfp4W4a4StaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + nvfp4_w4a4_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } nvfp4_w4a4_mma_kernel<<>>( - activation, static_cast(weight.qdata), - static_cast(weight.scales), tokens, alpha, Nvfp4IdentityEpilogue{}, - output, row_policy); + Nvfp4SwiGluRows, true> + <<>>( + activation, static_cast(weight.qdata), + static_cast(weight.scales), tokens, alpha, + Nvfp4IdentityEpilogue{}, output, row_policy); CUDA_CHECK(cudaGetLastError()); } diff --git a/src/ops/linear_swiglu/w8/w8_dflash2_linear_swiglu.cu b/src/ops/linear_swiglu/w8/w8_dflash2_linear_swiglu.cu index 34ac22a1b1..b140067291 100644 --- a/src/ops/linear_swiglu/w8/w8_dflash2_linear_swiglu.cu +++ b/src/ops/linear_swiglu/w8/w8_dflash2_linear_swiglu.cu @@ -32,9 +32,18 @@ void launch_tile(const Tensor& x, const Weight& weight, Tensor& out, cudaStream_ const W8SwiGluDirectEpilogue epilogue{static_cast<__nv_bfloat16*>(out.data), kIntermediate}; const RowPolicy row_policy{}; constexpr int kBlocks = kIntermediate / RowPolicy::kOutputRowsPerCta; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(weight.qdata), static_cast(weight.scales), ignored_output, epilogue, row_policy, x.ne[1]); diff --git a/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_mma.cu b/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_mma.cu index 9f2fc1f14f..882d9ddb64 100644 --- a/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_mma.cu +++ b/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_mma.cu @@ -12,11 +12,19 @@ void launch_variant(const Tensor& x, const Weight& w, Tensor& out, cudaStream_t const W8ContiguousOutput output{static_cast<__nv_bfloat16*>(out.data), out.ne[0]}; const dim3 grid(out.ne[0] / (Schedule::BM / 2), static_cast(div_up(x.ne[1], Schedule::BN)), 1u); + constexpr std::size_t kDynamicBytes = w8_rowsplit_gemm_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_rowsplit_gemm_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_rowsplit_gemm_mma_kernel - <<>>(static_cast(x.data), - static_cast(w.qdata), - static_cast(w.scales), output, - w.n, w.k, x.ne[1], w.padded_shape[1]); + <<>>( + static_cast(x.data), static_cast(w.qdata), + static_cast(w.scales), output, w.n, w.k, x.ne[1], + w.padded_shape[1]); } template diff --git a/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_splitk.cu b/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_splitk.cu index 6b8abbd820..8f71105f99 100644 --- a/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_splitk.cu +++ b/src/ops/linear_swiglu/w8/w8_linear_swiglu_gemm_splitk.cu @@ -37,9 +37,19 @@ void launch_active_cols(const Tensor& x, const Weight& w, Tensor& out, cudaStrea const W8ContiguousOutput ignored_output{static_cast<__nv_bfloat16*>(out.data), kIntermediate}; const W8SwiGluDirectEpilogue epilogue{static_cast<__nv_bfloat16*>(out.data), kIntermediate}; const RowPolicy row_policy{}; + constexpr std::size_t kDynamicBytes = + w8_small_t_mma_dynamic_bytes(); + if constexpr (kDynamicBytes > kW8SmallTMmaStaticSharedBytes) { + static const cudaError_t attribute = cudaFuncSetAttribute( + w8_small_t_mma_kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kDynamicBytes)); + CUDA_CHECK(attribute); + } w8_small_t_mma_kernel - <<>>( + <<>>( static_cast(x.data), static_cast(w.qdata), static_cast(w.scales), ignored_output, epilogue, row_policy); }