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;