From 3a051da0d8d4c3e85564acbfd1cc50df813fa7a5 Mon Sep 17 00:00:00 2001 From: Ravi M Date: Wed, 30 Sep 2026 11:29:47 -0700 Subject: [PATCH 1/2] [Common] Fall back from gated TMA kernels when shared memory does not fit The delayed-scaling gated activation dispatch selected the TMA kernel (cast_fp8_gated_kernel) on every device with compute capability 10.0+, without checking its dynamic shared memory request against the device limit. On SM 12.0 (opt-in limit 99 KiB per block) FP32 configurations request up to 160 KiB, so cudaFuncSetAttribute fails with "invalid argument" in both forward and backward (issue #3299). Compute the kernel's shared memory requirement in one helper shared by dispatch and launch, compare it with the device's cached opt-in per-block limit, and use the existing non-TMA kernels when it does not fit. Configurations that fit (e.g. BF16/FP16 on SM 12.0, everything on SM 10.0) keep using the TMA kernel. Co-Authored-By: Claude Opus 5.5 Signed-off-by: Ravi M --- .../common/cast/dispatch/gated.cuh | 8 +++- .../common/cast/fp8/gated_fp8.cuh | 46 +++++++++++++------ .../common/util/cuda_runtime.cpp | 17 +++++++ transformer_engine/common/util/cuda_runtime.h | 12 +++++ 4 files changed, 68 insertions(+), 15 deletions(-) diff --git a/transformer_engine/common/cast/dispatch/gated.cuh b/transformer_engine/common/cast/dispatch/gated.cuh index 06e8f0e306c..e26d855cd63 100644 --- a/transformer_engine/common/cast/dispatch/gated.cuh +++ b/transformer_engine/common/cast/dispatch/gated.cuh @@ -46,7 +46,9 @@ void quantize_gated_fwd_helper(const NVTETensor nvte_input, NVTETensor nvte_outp switch (output->scaling_mode) { case NVTE_DELAYED_TENSOR_SCALING: { - const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100(); + const bool use_tma_kernels = + (cols % 32 == 0) && is_supported_by_CC_100() && + fp8::cast_gated_tma_fits_device(/*is_bwd=*/false, input.dtype(), output->dtype()); if (use_tma_kernels) { Tensor dummy_grad_tensor; fp8::cast_gated_tma(input, dummy_grad_tensor, @@ -137,7 +139,9 @@ void quantize_gated_bwd_helper(const NVTETensor nvte_grad, const NVTETensor nvte switch (output->scaling_mode) { case NVTE_DELAYED_TENSOR_SCALING: { - const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100(); + const bool use_tma_kernels = + (cols % 32 == 0) && is_supported_by_CC_100() && + fp8::cast_gated_tma_fits_device(/*is_bwd=*/true, gated_input.dtype(), output->dtype()); if (use_tma_kernels) { fp8::cast_gated_tma(gated_input, grad, output, p, stream); diff --git a/transformer_engine/common/cast/fp8/gated_fp8.cuh b/transformer_engine/common/cast/fp8/gated_fp8.cuh index 631143c6aef..875fbb7392f 100644 --- a/transformer_engine/common/cast/fp8/gated_fp8.cuh +++ b/transformer_engine/common/cast/fp8/gated_fp8.cuh @@ -17,6 +17,7 @@ #include #include "../../common.h" +#include "../../util/cuda_runtime.h" #include "../../util/math.h" #include "../../util/ptx.cuh" #include "../../util/vectorized_pointwise.h" @@ -42,6 +43,9 @@ constexpr size_t BUFFER_STAGES_NUM = BUFFER_DIM_Y / THREADS_PER_CHUNK_Y; // 8 constexpr size_t ITERATIONS = CHUNK_DIM_Y / BUFFER_DIM_Y; // 4 = 128 / 32 static_assert(ITERATIONS >= 1); +// Static shared memory of cast_fp8_gated_kernel (one mbarrier per iteration) +constexpr size_t STATIC_SHMEM_SIZE = ITERATIONS * sizeof(uint64_t); + template __global__ void __launch_bounds__(THREADS_PER_CHUNK) @@ -279,6 +283,33 @@ __global__ void __launch_bounds__(THREADS_PER_CHUNK) } } // namespace kernel +// Dynamic shared memory requested by cast_fp8_gated_kernel. Must match the +// buffer layout inside the kernel. +inline size_t cast_gated_tma_dynamic_shmem_size(const bool is_bwd, const DType itype, + const DType otype) { + using namespace kernel; + const size_t buff_elems_total = BUFFERS_NUM * SHMEM_DIM_Y * SHMEM_DIM_X; + const size_t buff_size_aligned_in = + DIVUP_TO_MULTIPLE(buff_elems_total * typeToNumBits(itype) / 8, TMA_SHMEM_ALIGNMENT); + const size_t buff_size_aligned_out = + DIVUP_TO_MULTIPLE(buff_elems_total * typeToNumBits(otype) / 8, TMA_SHMEM_ALIGNMENT); + const size_t grad_mem = (is_bwd ? buff_size_aligned_in : 0); + const size_t in_act_mem = buff_size_aligned_in; + const size_t in_gate_mem = buff_size_aligned_in; + const size_t out_act_mem = buff_size_aligned_out; + const size_t out_gate_mem = buff_size_aligned_out; + return grad_mem + (in_act_mem + in_gate_mem) + (out_act_mem + out_gate_mem) + TMA_SHMEM_ALIGNMENT; +} + +// Whether cast_fp8_gated_kernel fits in the shared memory of the current device. +// Devices with compute capability 10.0+ differ in shared memory per block +// (e.g. SM 12.0 has much less than SM 10.0), so FP32 configurations may not fit. +inline bool cast_gated_tma_fits_device(const bool is_bwd, const DType itype, const DType otype) { + const size_t required = + cast_gated_tma_dynamic_shmem_size(is_bwd, itype, otype) + kernel::STATIC_SHMEM_SIZE; + return required <= cuda::max_shared_memory_per_block_optin(); +} + template void cast_gated_tma(const Tensor &gated_input, const Tensor &grad, Tensor *output, ParamOP &p, @@ -329,19 +360,8 @@ void cast_gated_tma(const Tensor &gated_input, const Tensor &grad, Tensor *outpu SHMEM_DIM_X, tensor_stride_elems, cols, typeToNumBits(output->dtype())); - const size_t buff_elems_total = BUFFERS_NUM * SHMEM_DIM_Y * SHMEM_DIM_X; - const size_t buff_size_aligned_in = - DIVUP_TO_MULTIPLE(buff_elems_total * sizeof(IType), TMA_SHMEM_ALIGNMENT); - const size_t buff_size_aligned_out = - DIVUP_TO_MULTIPLE(buff_elems_total * sizeof(OType), TMA_SHMEM_ALIGNMENT); - const size_t grad_mem = (IS_BWD ? buff_size_aligned_in : 0); - const size_t in_act_mem = buff_size_aligned_in; - const size_t in_gate_mem = buff_size_aligned_in; - const size_t out_act_mem = buff_size_aligned_out; - const size_t out_gate_mem = buff_size_aligned_out; - - const size_t shmem_size = grad_mem + (in_act_mem + in_gate_mem) + - (out_act_mem + out_gate_mem) + TMA_SHMEM_ALIGNMENT; + const size_t shmem_size = + cast_gated_tma_dynamic_shmem_size(IS_BWD, gated_input.dtype(), output->dtype()); auto kernel = cast_fp8_gated_kernel; NVTE_CHECK_CUDA(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, diff --git a/transformer_engine/common/util/cuda_runtime.cpp b/transformer_engine/common/util/cuda_runtime.cpp index d23f11ccc6d..ca2206c4f95 100644 --- a/transformer_engine/common/util/cuda_runtime.cpp +++ b/transformer_engine/common/util/cuda_runtime.cpp @@ -143,6 +143,23 @@ int sm_count(int device_id) { return cache[device_id]; } +size_t max_shared_memory_per_block_optin(int device_id) { + static std::vector cache(num_devices(), 0); + static std::vector flags(num_devices()); + if (device_id < 0) { + device_id = current_device(); + } + NVTE_CHECK(0 <= device_id && device_id < num_devices(), "invalid CUDA device ID"); + auto init = [&]() { + int value; + NVTE_CHECK_CUDA( + cudaDeviceGetAttribute(&value, cudaDevAttrMaxSharedMemoryPerBlockOptin, device_id)); + cache[device_id] = static_cast(value); + }; + std::call_once(flags[device_id], init); + return cache[device_id]; +} + void stream_priority_range(int *low_priority, int *high_priority, int device_id) { static std::vector> cache(num_devices()); static std::vector flags(num_devices()); diff --git a/transformer_engine/common/util/cuda_runtime.h b/transformer_engine/common/util/cuda_runtime.h index 0f355940010..b4315729a34 100644 --- a/transformer_engine/common/util/cuda_runtime.h +++ b/transformer_engine/common/util/cuda_runtime.h @@ -38,6 +38,18 @@ int sm_arch(int device_id = -1); */ int sm_count(int device_id = -1); +/* \brief Maximum shared memory per block available with opt-in + * + * This is the upper bound for the static plus dynamic shared memory + * of a single thread block, after raising the kernel's + * cudaFuncAttributeMaxDynamicSharedMemorySize attribute. + * + * \param[in] device_id CUDA device (default is current device) + * + * \return Shared memory size in bytes + */ +size_t max_shared_memory_per_block_optin(int device_id = -1); + /* \brief Minimum and maximum stream priorities supported on device * * \param[in] device_id CUDA device (default is current device) From fbad621cfa2e0b5efa22142bdcd41d8a2de9d76d Mon Sep 17 00:00:00 2001 From: Ravi M Date: Sat, 3 Oct 2026 17:45:24 -0700 Subject: [PATCH 2/2] [Common] Read gated TMA kernel static shared memory from the compiled kernel The fit check hard-coded the static shared memory of cast_fp8_gated_kernel as its mbarrier array (32 B) and missed the staging array declared in reduce_max (64 B). Query the compiled kernel instance instead via cudaFuncGetAttributes, exposed as cuda::static_shared_memory_size() and cached per kernel and device, so __shared__ arrays in called device functions are always included. Also replace the file-local get_max_dynamic_smem() in swizzle.cu with cuda::max_shared_memory_per_block_optin(). Co-Authored-By: Claude Opus 5.5 Signed-off-by: Ravi M --- .../common/cast/dispatch/gated.cuh | 6 +++-- .../common/cast/fp8/gated_fp8.cuh | 19 ++++++++++---- transformer_engine/common/swizzle/swizzle.cu | 26 ++++++------------- .../common/util/cuda_runtime.cpp | 16 ++++++++++++ transformer_engine/common/util/cuda_runtime.h | 12 +++++++++ 5 files changed, 54 insertions(+), 25 deletions(-) diff --git a/transformer_engine/common/cast/dispatch/gated.cuh b/transformer_engine/common/cast/dispatch/gated.cuh index e26d855cd63..0c1761264a1 100644 --- a/transformer_engine/common/cast/dispatch/gated.cuh +++ b/transformer_engine/common/cast/dispatch/gated.cuh @@ -48,7 +48,8 @@ void quantize_gated_fwd_helper(const NVTETensor nvte_input, NVTETensor nvte_outp case NVTE_DELAYED_TENSOR_SCALING: { const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100() && - fp8::cast_gated_tma_fits_device(/*is_bwd=*/false, input.dtype(), output->dtype()); + fp8::cast_gated_tma_fits_device( + input.dtype(), output->dtype()); if (use_tma_kernels) { Tensor dummy_grad_tensor; fp8::cast_gated_tma(input, dummy_grad_tensor, @@ -141,7 +142,8 @@ void quantize_gated_bwd_helper(const NVTETensor nvte_grad, const NVTETensor nvte case NVTE_DELAYED_TENSOR_SCALING: { const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100() && - fp8::cast_gated_tma_fits_device(/*is_bwd=*/true, gated_input.dtype(), output->dtype()); + fp8::cast_gated_tma_fits_device( + gated_input.dtype(), output->dtype()); if (use_tma_kernels) { fp8::cast_gated_tma(gated_input, grad, output, p, stream); diff --git a/transformer_engine/common/cast/fp8/gated_fp8.cuh b/transformer_engine/common/cast/fp8/gated_fp8.cuh index 875fbb7392f..124ec48b91b 100644 --- a/transformer_engine/common/cast/fp8/gated_fp8.cuh +++ b/transformer_engine/common/cast/fp8/gated_fp8.cuh @@ -43,9 +43,6 @@ constexpr size_t BUFFER_STAGES_NUM = BUFFER_DIM_Y / THREADS_PER_CHUNK_Y; // 8 constexpr size_t ITERATIONS = CHUNK_DIM_Y / BUFFER_DIM_Y; // 4 = 128 / 32 static_assert(ITERATIONS >= 1); -// Static shared memory of cast_fp8_gated_kernel (one mbarrier per iteration) -constexpr size_t STATIC_SHMEM_SIZE = ITERATIONS * sizeof(uint64_t); - template __global__ void __launch_bounds__(THREADS_PER_CHUNK) @@ -304,9 +301,21 @@ inline size_t cast_gated_tma_dynamic_shmem_size(const bool is_bwd, const DType i // Whether cast_fp8_gated_kernel fits in the shared memory of the current device. // Devices with compute capability 10.0+ differ in shared memory per block // (e.g. SM 12.0 has much less than SM 10.0), so FP32 configurations may not fit. -inline bool cast_gated_tma_fits_device(const bool is_bwd, const DType itype, const DType otype) { +// The static shared memory is read from the compiled kernel so that __shared__ +// arrays in the device functions it calls (e.g. reduce_max) are included. +template +bool cast_gated_tma_fits_device(const DType itype, const DType otype) { + using namespace kernel; + size_t static_shmem_size = 0; + TRANSFORMER_ENGINE_TYPE_SWITCH_INPUT( + itype, IType, + TRANSFORMER_ENGINE_TYPE_SWITCH_OUTPUT( + otype, OType, + static_shmem_size = cuda::static_shared_memory_size(reinterpret_cast( + &cast_fp8_gated_kernel)););); const size_t required = - cast_gated_tma_dynamic_shmem_size(is_bwd, itype, otype) + kernel::STATIC_SHMEM_SIZE; + cast_gated_tma_dynamic_shmem_size(IS_BWD, itype, otype) + static_shmem_size; return required <= cuda::max_shared_memory_per_block_optin(); } diff --git a/transformer_engine/common/swizzle/swizzle.cu b/transformer_engine/common/swizzle/swizzle.cu index 38b526360e6..69ba2bc2d4a 100644 --- a/transformer_engine/common/swizzle/swizzle.cu +++ b/transformer_engine/common/swizzle/swizzle.cu @@ -24,20 +24,6 @@ namespace { constexpr int MXFP8_BLOCK_SIZE = 32; constexpr int NVFP4_BLOCK_SIZE = 16; -int get_max_dynamic_smem(int device_id = -1) { - static std::vector cache(cuda::num_devices(), -1); - static std::vector flags(cuda::num_devices()); - if (device_id < 0) { - device_id = cuda::current_device(); - } - auto init = [&]() { - NVTE_CHECK_CUDA(cudaDeviceGetAttribute(&cache[device_id], - cudaDevAttrMaxSharedMemoryPerBlockOptin, device_id)); - }; - std::call_once(flags[device_id], init); - return cache[device_id]; -} - constexpr __device__ __host__ int TB_DIM = 32; constexpr __device__ __host__ int NEW_SF_TILE_DIM_K = 16; constexpr __device__ __host__ int N_SF_PER_TD_PER_TILE = 4; @@ -1109,7 +1095,8 @@ void swizzle_scaling_factors(const Tensor* input, Tensor* output, cudaStream_t s const int narrow_k_slm_size = TB_DIM * num_tiles_k * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast(sizeof(int8_t)); - if (num_tiles_k < TB_DIM && narrow_k_slm_size <= get_max_dynamic_smem()) { + if (num_tiles_k < TB_DIM && + static_cast(narrow_k_slm_size) <= cuda::max_shared_memory_per_block_optin()) { // Narrow-K: batch TB_DIM M-tiles per block, fully utilizing all threads. dim3 num_blocks_narrow(DIVUP(num_tiles_m, TB_DIM)); NVTE_CHECK_CUDA( @@ -1166,7 +1153,8 @@ void swizzle_scaling_factors(const Tensor* input, Tensor* output, cudaStream_t s const int narrow_m_slm_size = TB_DIM * num_tiles_m * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast(sizeof(int8_t)); - if (num_tiles_m < TB_DIM && narrow_m_slm_size <= get_max_dynamic_smem()) { + if (num_tiles_m < TB_DIM && + static_cast(narrow_m_slm_size) <= cuda::max_shared_memory_per_block_optin()) { // Narrow-M: batch TB_DIM K-tiles per block, fully utilizing all threads. dim3 num_blocks_narrow(DIVUP(num_tiles_k, TB_DIM)); NVTE_CHECK_CUDA( @@ -1534,7 +1522,8 @@ void multi_tensor_swizzle_scaling_factors(const std::vector& input, const int narrow_k_slm = TB_DIM * num_tiles_k * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast(sizeof(int8_t)); all_narrow_k = - all_narrow_k && (num_tiles_k < TB_DIM) && (narrow_k_slm <= get_max_dynamic_smem()); + all_narrow_k && (num_tiles_k < TB_DIM) && + (static_cast(narrow_k_slm) <= cuda::max_shared_memory_per_block_optin()); int vec_load_size_i = (num_tiles_k - 1) % 4 + 1; // We use the minimum vec_load_size across all tensors. // TODO(zhongbo): fix vec_load_size for NVFP4 @@ -1606,7 +1595,8 @@ void multi_tensor_swizzle_scaling_factors(const std::vector& input, const int narrow_m_slm = TB_DIM * num_tiles_m * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast(sizeof(int8_t)); all_narrow_m = - all_narrow_m && (num_tiles_m < TB_DIM) && (narrow_m_slm <= get_max_dynamic_smem()); + all_narrow_m && (num_tiles_m < TB_DIM) && + (static_cast(narrow_m_slm) <= cuda::max_shared_memory_per_block_optin()); int vec_load_size_i = (num_tiles_k - 1) % 4 + 1; // We use the minimum vec_load_size across all tensors. vec_load_size = std::min(vec_load_size, vec_load_size_i); diff --git a/transformer_engine/common/util/cuda_runtime.cpp b/transformer_engine/common/util/cuda_runtime.cpp index ca2206c4f95..17e5026bf09 100644 --- a/transformer_engine/common/util/cuda_runtime.cpp +++ b/transformer_engine/common/util/cuda_runtime.cpp @@ -10,7 +10,9 @@ #include #include +#include #include +#include #include "../common.h" #include "../util/cuda_driver.h" @@ -160,6 +162,20 @@ size_t max_shared_memory_per_block_optin(int device_id) { return cache[device_id]; } +size_t static_shared_memory_size(const void *kernel) { + static std::map, size_t> cache; + static std::mutex mutex; + const auto key = std::make_pair(kernel, current_device()); + const std::lock_guard lock(mutex); + auto it = cache.find(key); + if (it == cache.end()) { + cudaFuncAttributes attr; + NVTE_CHECK_CUDA(cudaFuncGetAttributes(&attr, kernel)); + it = cache.emplace(key, attr.sharedSizeBytes).first; + } + return it->second; +} + void stream_priority_range(int *low_priority, int *high_priority, int device_id) { static std::vector> cache(num_devices()); static std::vector flags(num_devices()); diff --git a/transformer_engine/common/util/cuda_runtime.h b/transformer_engine/common/util/cuda_runtime.h index b4315729a34..02f95c08f1e 100644 --- a/transformer_engine/common/util/cuda_runtime.h +++ b/transformer_engine/common/util/cuda_runtime.h @@ -50,6 +50,18 @@ int sm_count(int device_id = -1); */ size_t max_shared_memory_per_block_optin(int device_id = -1); +/* \brief Static shared memory used by a compiled kernel on the current device + * + * Size of the kernel's fixed-size __shared__ arrays, including those + * declared in device functions it calls, as reported by + * cudaFuncGetAttributes. The result is cached per kernel and device. + * + * \param[in] kernel Pointer to the __global__ function + * + * \return Shared memory size in bytes + */ +size_t static_shared_memory_size(const void *kernel); + /* \brief Minimum and maximum stream priorities supported on device * * \param[in] device_id CUDA device (default is current device)