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); }