Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
191 changes: 189 additions & 2 deletions transformer_engine/common/cast/fused_group_requantize.cu
Original file line number Diff line number Diff line change
Expand Up @@ -13,10 +13,12 @@
#include <cuda_runtime.h>
#include <transformer_engine/cast.h>

#include <algorithm>
#include <limits>
#include <type_traits>

#include "../common.h"
#include "../util/cuda_runtime.h"
#include "../util/ptx_arch_spec.cuh"
#include "../utils.cuh"
#include "mxfp8/swizzle.cuh"
Expand Down Expand Up @@ -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<IType, fp8e4m3, true, false>, 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 <typename IType>
__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<const uint32_t *>(input);
uint32_t *const output_words = reinterpret_cast<uint32_t *>(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<int64_t>(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<size_t>(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<size_t>(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<const uint32_t *>(
input_scale_inv + static_cast<size_t>(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<int>(offsets[group] / num_cols);
const int group_rows = static_cast<int>((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<uint32_t *>(rowwise_scale_inv_swizzled +
(static_cast<size_t>(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<bf16>(0.0f), static_cast<bf16>(0.0f)};
ptx::bf16x2 amax23 = {static_cast<bf16>(0.0f), static_cast<bf16>(0.0f)};
#pragma unroll
for (int r = 0; r < 32; ++r) {
const e8m0_t scale_code =
static_cast<e8m0_t>((__shfl_sync(0xffffffff, scale_word, r) >> scale_shift) & 0xFF);
const uint16_t values01 = static_cast<uint16_t>(data[r]);
const uint16_t values23 = static_cast<uint16_t>(data[r] >> 16);
dequantized01[r] =
dequantize_mxfp8_2x(reinterpret_cast<const ptx::FPx2<IType> &>(values01), scale_code);
dequantized23[r] =
dequantize_mxfp8_2x(reinterpret_cast<const ptx::FPx2<IType> &>(values23), scale_code);
ptx::abs_max_2x(amax01, amax01, dequantized01[r]);
ptx::abs_max_2x(amax23, amax23, dequantized23[r]);
}
constexpr float kMaxNormRcp = Quantized_Limits<fp8e4m3>::max_norm_rcp;
const e8m0_t exponent0 = ptx::float_to_e8m0(static_cast<float>(amax01.x) * kMaxNormRcp);
const e8m0_t exponent1 = ptx::float_to_e8m0(static_cast<float>(amax01.y) * kMaxNormRcp);
const e8m0_t exponent2 = ptx::float_to_e8m0(static_cast<float>(amax23.x) * kMaxNormRcp);
const e8m0_t exponent3 = ptx::float_to_e8m0(static_cast<float>(amax23.y) * kMaxNormRcp);
const ptx::bf16x2 multiplier01 = {ptx::exp2f_rcp<bf16>(exponent0),
ptx::exp2f_rcp<bf16>(exponent1)};
const ptx::bf16x2 multiplier23 = {ptx::exp2f_rcp<bf16>(exponent2),
ptx::exp2f_rcp<bf16>(exponent3)};

uint32_t *const dst =
output_words + static_cast<size_t>(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<size_t>(r) * words_per_row] = reinterpret_cast<const uint32_t &>(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<size_t>(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<size_t>(tile_x) * colwise_scale_tiles_x + tile_in_group) * 512;
reinterpret_cast<uint32_t *>(colwise_scale_dst)[tid] =
reinterpret_cast<const uint32_t *>(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 <typename IType>
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<int>(
std::min<int64_t>(num_tiles, static_cast<int64_t>(cuda::sm_count()) * kWideCtasPerSM));
fused_group_requantize_wide_kernel<IType><<<grid, kWideThreads, 0, stream>>>(
reinterpret_cast<const IType *>(input.data.dptr),
reinterpret_cast<const e8m0_t *>(input.scale_inv.dptr), input_scale_stride,
reinterpret_cast<fp8e4m3 *>(output->columnwise_data.dptr),
reinterpret_cast<e8m0_t *>(output->scale_inv.dptr),
reinterpret_cast<e8m0_t *>(output->columnwise_scale_inv.dptr),
reinterpret_cast<const int64_t *>(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) {
Expand Down Expand Up @@ -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<int64_t>(num_rows / kWideTile) * static_cast<int64_t>(num_cols / kWideTile);
const bool use_wide_kernel =
use_fast_math && !return_dequantized && num_groups <= kWideMaxGroups &&
num_wide_tiles <= static_cast<int64_t>(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<IType>(input, output, tensor_offsets, num_groups,
static_cast<int>(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(
Expand Down
Loading