From 8ea14e248951757238872eb6685ca062ff481623 Mon Sep 17 00:00:00 2001 From: Ravi M Date: Sat, 3 Oct 2026 18:37:22 -0700 Subject: [PATCH] [Common] Add launch bounds to multi-tensor swizzle kernels The multi-tensor swizzle and unswizzle kernels launch with TB_DIM x TB_DIM (1024) threads but, unlike the other swizzle kernels, had no __launch_bounds__. ptxas could then use more than 64 registers per thread (up to 99 for the int4 variants), which makes the launch fail with "too many resources requested for launch". Which variants fail depends on the target arch and CUDA version. Add __launch_bounds__(TB_DIM * TB_DIM) to the four kernels and add test shapes that reach the rowwise int2 and columnwise int4 variants. Fixes #3621 Co-Authored-By: Claude Opus 5.5 Signed-off-by: Ravi M --- tests/cpp/operator/test_multi_swizzle.cu | 4 ++++ transformer_engine/common/swizzle/swizzle.cu | 12 ++++++++---- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/cpp/operator/test_multi_swizzle.cu b/tests/cpp/operator/test_multi_swizzle.cu index 4984b7783b3..d8b8211e5b7 100644 --- a/tests/cpp/operator/test_multi_swizzle.cu +++ b/tests/cpp/operator/test_multi_swizzle.cu @@ -377,6 +377,10 @@ std::vector> multi_tensor_test_cases = { {3, 256, 4096, false}, {2, 128, 8192, true}, {2, 128, 8192, false}, + // Rowwise num_tiles_k = 34 selects vec_load_size = 2 (int2 kernel) + {2, 128, 4352, true}, + // Colwise num_tiles_k = 512 / 32 / 4 = 4 selects vec_load_size = 4 (int4 kernel) + {2, 512, 4096, false}, }; } // namespace diff --git a/transformer_engine/common/swizzle/swizzle.cu b/transformer_engine/common/swizzle/swizzle.cu index 38b526360e6..f402f175cf3 100644 --- a/transformer_engine/common/swizzle/swizzle.cu +++ b/transformer_engine/common/swizzle/swizzle.cu @@ -791,7 +791,8 @@ __global__ void __launch_bounds__(TB_DIM* TB_DIM) } template -__global__ void multi_tensor_unswizzle_row_scaling_kernel(MultiSwizzleArgs kernel_args) { +__global__ void __launch_bounds__(TB_DIM* TB_DIM) + multi_tensor_unswizzle_row_scaling_kernel(MultiSwizzleArgs kernel_args) { const int bid = blockIdx.x; int tensor_id = 0; while (kernel_args.block_range[tensor_id + 1] <= bid) { @@ -818,7 +819,8 @@ __global__ void multi_tensor_unswizzle_row_scaling_kernel(MultiSwizzleArgs kerne } template -__global__ void multi_tensor_unswizzle_col_scaling_kernel(MultiSwizzleArgs kernel_args) { +__global__ void __launch_bounds__(TB_DIM* TB_DIM) + multi_tensor_unswizzle_col_scaling_kernel(MultiSwizzleArgs kernel_args) { const int bid = blockIdx.x; int tensor_id = 0; while (kernel_args.block_range[tensor_id + 1] <= bid) { @@ -844,7 +846,8 @@ __global__ void multi_tensor_unswizzle_col_scaling_kernel(MultiSwizzleArgs kerne } template -__global__ void multi_tensor_swizzle_row_scaling_kernel(MultiSwizzleArgs kernel_args) { +__global__ void __launch_bounds__(TB_DIM* TB_DIM) + multi_tensor_swizzle_row_scaling_kernel(MultiSwizzleArgs kernel_args) { // Find tensor corresponding to block const int bid = blockIdx.x; int tensor_id = 0; @@ -878,7 +881,8 @@ __global__ void multi_tensor_swizzle_row_scaling_kernel(MultiSwizzleArgs kernel_ } template -__global__ void multi_tensor_swizzle_col_scaling_kernel(MultiSwizzleArgs kernel_args) { +__global__ void __launch_bounds__(TB_DIM* TB_DIM) + multi_tensor_swizzle_col_scaling_kernel(MultiSwizzleArgs kernel_args) { // Find tensor corresponding to block const int bid = blockIdx.x; int tensor_id = 0;