diff --git a/transformer_engine/common/cast/fused_group_requantize.cu b/transformer_engine/common/cast/fused_group_requantize.cu index 7efe629d11..aad83cb67c 100644 --- a/transformer_engine/common/cast/fused_group_requantize.cu +++ b/transformer_engine/common/cast/fused_group_requantize.cu @@ -13,10 +13,12 @@ #include #include +#include #include #include #include "../common.h" +#include "../util/cuda_runtime.h" #include "../util/ptx_arch_spec.cuh" #include "../utils.cuh" #include "mxfp8/swizzle.cuh" @@ -413,6 +415,172 @@ void launch_fused_group_requantize(const Tensor &input, Tensor *output, input_scale_stride); } +// --------------------------------------------------------------------------- +// Wide-tile variant for the fast-math case without a dequantized output. +// +// Same per-element arithmetic as fused_group_requantize_kernel, so the +// outputs are bit-identical; only the work decomposition differs. The kernel above runs one 32x128 +// tile per 128-thread CTA, thread 0 binary-searches the group offsets before the tile can finish, +// and both scale outputs are written one byte at a time. Here a CTA owns a 128x128 tile: warp w +// takes rows 32w..32w+31 and lane l columns 4l..4l+3, so the data moves as coalesced 128-byte row +// segments held in registers, and each of the two swizzled scale outputs is one contiguous +// 512-byte block per tile. The grid is persistent and walks the tiles in row-block order, so a +// CTA stops at the first tile past tensor_offsets[num_groups]; a capacity-sized buffer with a +// short live prefix then costs little more than the live rows. + +constexpr int kWideTile = 128; +constexpr int kWideThreads = 128; +constexpr int kWideMaxGroups = 128; +constexpr int kWideCtasPerSM = 16; + +template +__global__ void __launch_bounds__(kWideThreads) + fused_group_requantize_wide_kernel(const IType *const input, + const e8m0_t *const input_scale_inv, + const int input_scale_stride, fp8e4m3 *const colwise_out, + e8m0_t *const rowwise_scale_inv_swizzled, + e8m0_t *const colwise_scale_inv, + const int64_t *const tensor_offsets, const int num_groups, + const int num_cols, const int num_tiles) { +#if (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + static_assert(sizeof(IType) == 1); + __shared__ int64_t offsets[kWideMaxGroups + 1]; + __shared__ alignas(16) e8m0_t colwise_scale_block[kWideTile * 4]; + + const int tid = threadIdx.x; + const int warp = tid / THREADS_PER_WARP; + const int lane = tid % THREADS_PER_WARP; + const int tiles_x = num_cols / kWideTile; + const int words_per_row = num_cols / 4; + // This lane's 4 columns lie in the 32-column scaling block lane / 8 of the tile. + const int scale_shift = 8 * (lane / 8); + for (int i = tid; i <= num_groups; i += kWideThreads) { + offsets[i] = tensor_offsets[i]; + } + __syncthreads(); + const int64_t live_end = offsets[num_groups]; + + const uint32_t *const input_words = reinterpret_cast(input); + uint32_t *const output_words = reinterpret_cast(colwise_out); + for (int tile = blockIdx.x; tile < num_tiles; tile += gridDim.x) { + const int tile_x = tile % tiles_x; + const int tile_y = tile / tiles_x; + const int row_base = tile_y * kWideTile; + const int col_base = tile_x * kWideTile; + const int64_t row_element_base = static_cast(row_base) * num_cols; + // Rows at or past the live end are left untouched, as by the kernel above. Later tiles of + // this CTA lie further down, so the walk ends here. + if (row_element_base >= live_end) { + break; + } + + // 32 rows x 4 columns per thread, one 4-byte word per row. + const int row0 = row_base + 32 * warp; + const uint32_t *const src = + input_words + static_cast(row0) * words_per_row + col_base / 4 + lane; + uint32_t data[32]; +#pragma unroll + for (int r = 0; r < 32; ++r) { + data[r] = __ldg(src + static_cast(r) * words_per_row); + } + // Lane l holds the 4 rowwise scale codes (the tile's four 32-column blocks) of row row0 + l. + const uint32_t scale_word = __ldg(reinterpret_cast( + input_scale_inv + static_cast(row0 + lane) * input_scale_stride + + col_base / MXFP8_SCALE_DIM)); + + // The owning group, as found by the upper_bound lookup of the kernel above. 128-aligned group + // sizes mean a tile never straddles a group boundary. + int group = 0; + for (int i = 0; i + 1 < num_groups; ++i) { + group += (offsets[i + 1] <= row_element_base) ? 1 : 0; + } + const int group_start = static_cast(offsets[group] / num_cols); + const int group_rows = static_cast((offsets[group + 1] - offsets[group]) / num_cols); + + // GEMM-swizzled copy of the rowwise scales: the tile's 128 rows x 4 codes are one 512-byte + // block, in which row r of the tile starts at (r % 32) * 16 + (r / 32) * 4. + *reinterpret_cast(rowwise_scale_inv_swizzled + + (static_cast(tile_y) * tiles_x + tile_x) * 512 + + lane * 16 + warp * 4) = scale_word; + + // Dequantize to BF16 and take the amax of each column over the 32-row block. + ptx::bf16x2 dequantized01[32]; + ptx::bf16x2 dequantized23[32]; + ptx::bf16x2 amax01 = {static_cast(0.0f), static_cast(0.0f)}; + ptx::bf16x2 amax23 = {static_cast(0.0f), static_cast(0.0f)}; +#pragma unroll + for (int r = 0; r < 32; ++r) { + const e8m0_t scale_code = + static_cast((__shfl_sync(0xffffffff, scale_word, r) >> scale_shift) & 0xFF); + const uint16_t values01 = static_cast(data[r]); + const uint16_t values23 = static_cast(data[r] >> 16); + dequantized01[r] = + dequantize_mxfp8_2x(reinterpret_cast &>(values01), scale_code); + dequantized23[r] = + dequantize_mxfp8_2x(reinterpret_cast &>(values23), scale_code); + ptx::abs_max_2x(amax01, amax01, dequantized01[r]); + ptx::abs_max_2x(amax23, amax23, dequantized23[r]); + } + constexpr float kMaxNormRcp = Quantized_Limits::max_norm_rcp; + const e8m0_t exponent0 = ptx::float_to_e8m0(static_cast(amax01.x) * kMaxNormRcp); + const e8m0_t exponent1 = ptx::float_to_e8m0(static_cast(amax01.y) * kMaxNormRcp); + const e8m0_t exponent2 = ptx::float_to_e8m0(static_cast(amax23.x) * kMaxNormRcp); + const e8m0_t exponent3 = ptx::float_to_e8m0(static_cast(amax23.y) * kMaxNormRcp); + const ptx::bf16x2 multiplier01 = {ptx::exp2f_rcp(exponent0), + ptx::exp2f_rcp(exponent1)}; + const ptx::bf16x2 multiplier23 = {ptx::exp2f_rcp(exponent2), + ptx::exp2f_rcp(exponent3)}; + + uint32_t *const dst = + output_words + static_cast(row0) * words_per_row + col_base / 4 + lane; +#pragma unroll + for (int r = 0; r < 32; ++r) { + ptx::fp8e4m3x4 quantized; + ptx::mul_cvt_4x(quantized, dequantized01[r], multiplier01, dequantized23[r], multiplier23); + dst[static_cast(r) * words_per_row] = reinterpret_cast(quantized); + } + + // Columnwise scales: column c of the tile and 32-row block w go to (c % 32) * 16 + (c / 32) * 4 + // + w of the tile's 512-byte block in the group's swizzled layout. Stage the block in shared + // memory, then store it with one 4-byte word per thread. + const e8m0_t exponents[4] = {exponent0, exponent1, exponent2, exponent3}; +#pragma unroll + for (int j = 0; j < 4; ++j) { + const int c = 4 * lane + j; + colwise_scale_block[(c % 32) * 16 + (c / 32) * 4 + warp] = exponents[j]; + } + __syncthreads(); + const size_t colwise_scale_base = static_cast(group_start) / MXFP8_SCALE_DIM * num_cols; + const int tile_in_group = (row_base - group_start) / kWideTile; + const int colwise_scale_tiles_x = group_rows / kWideTile; + e8m0_t *const colwise_scale_dst = + colwise_scale_inv + colwise_scale_base + + (static_cast(tile_x) * colwise_scale_tiles_x + tile_in_group) * 512; + reinterpret_cast(colwise_scale_dst)[tid] = + reinterpret_cast(colwise_scale_block)[tid]; + __syncthreads(); // The next tile rewrites colwise_scale_block. + } +#else + NVTE_DEVICE_THREAD0_ERROR("Fused grouped requantization requires Blackwell (SM100+) hardware."); +#endif // (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) +} + +template +void launch_fused_group_requantize_wide(const Tensor &input, Tensor *output, + const Tensor &tensor_offsets, const int num_groups, + const int num_tiles, const int num_cols, + const int input_scale_stride, cudaStream_t stream) { + const int grid = static_cast( + std::min(num_tiles, static_cast(cuda::sm_count()) * kWideCtasPerSM)); + fused_group_requantize_wide_kernel<<>>( + reinterpret_cast(input.data.dptr), + reinterpret_cast(input.scale_inv.dptr), input_scale_stride, + reinterpret_cast(output->columnwise_data.dptr), + reinterpret_cast(output->scale_inv.dptr), + reinterpret_cast(output->columnwise_scale_inv.dptr), + reinterpret_cast(tensor_offsets.data.dptr), num_groups, num_cols, num_tiles); +} + void fused_group_requantize(const Tensor &input, Tensor *output, const Tensor &tensor_offsets, Tensor *dequantized, const QuantizationConfig *quant_config, cudaStream_t stream) { @@ -515,11 +683,30 @@ void fused_group_requantize(const Tensor &input, Tensor *output, const Tensor &t return; } + const bool use_fast_math = quant_config != nullptr && quant_config->use_fast_math; + + // The wide-tile kernel covers the fast-math case without a dequantized output. It loads the + // rowwise scales and stores both scale outputs as 4-byte words. + const int64_t num_wide_tiles = + static_cast(num_rows / kWideTile) * static_cast(num_cols / kWideTile); + const bool use_wide_kernel = + use_fast_math && !return_dequantized && num_groups <= kWideMaxGroups && + num_wide_tiles <= static_cast(kMaxInt) && input_scale_stride % 4 == 0 && + is_aligned_ptr(input.scale_inv.dptr, 4) && is_aligned_ptr(output->scale_inv.dptr, 4) && + is_aligned_ptr(output->columnwise_scale_inv.dptr, 4); + if (use_wide_kernel) { + TRANSFORMER_ENGINE_TYPE_SWITCH_FP8ONLY( + input.data.dtype, IType, + launch_fused_group_requantize_wide(input, output, tensor_offsets, num_groups, + static_cast(num_wide_tiles), num_cols, + input_scale_stride, stream);); // NOLINT(*) + NVTE_CHECK_CUDA(cudaGetLastError()); + return; + } + alignas(64) CUtensorMap output_tensor_map{}; create_2D_tensor_map(output_tensor_map, output->columnwise_data, num_rows, num_cols, kTileRows, kTileCols, num_cols, 0, typeToNumBits(output->columnwise_data.dtype)); - - const bool use_fast_math = quant_config != nullptr && quant_config->use_fast_math; TRANSFORMER_ENGINE_SWITCH_CONDITION( use_fast_math, USE_FAST_MATH, TRANSFORMER_ENGINE_TYPE_SWITCH_FP8ONLY(