diff --git a/cpp/kernels/fmha_v2/fmha_test.py b/cpp/kernels/fmha_v2/fmha_test.py index b79bef940dc7..4e5d653277dc 100644 --- a/cpp/kernels/fmha_v2/fmha_test.py +++ b/cpp/kernels/fmha_v2/fmha_test.py @@ -155,21 +155,8 @@ def test_trtllm_flash_attention_fmha(d, s, dtype, flag, tiled_kernel): check=True) -# The test cases for sage attention. -@pytest.mark.parametrize('d', [80, 128], ids=["head-size-80", "head-size-128"]) -@pytest.mark.parametrize('s', [1024, 4096], ids=["seqlen-1024", "seqlen-4096"]) -def test_trtllm_sage_attention_fmha(d, s): - sm_version = getSMVersion() - if sm_version != 89 and sm_version != 90: - pytest.skip("Sage attention only supports sm89 and sm90 currently.") - - # Ada. - if sm_version == 89: - subprocess.run( - f"bin/fmha.exe -v 0 -runs 1 -min-s 1024 -s {s} -b 16 -h 8 -d {d} -bf16 \ - -sage-block-q 64 -sage-block-k 32 -sage-block-v 32 -force-non-tiled", - shell=True, - check=True) +# SageAttention ships as a pre-built cubin (see SAGE_CUBIN_TRAITS in setup.py), so bin/fmha.exe +# contains no sage kernel to exercise. It is covered by the runtime integration tests. # The test cases for mla attention. diff --git a/cpp/kernels/fmha_v2/setup.py b/cpp/kernels/fmha_v2/setup.py index 82cd51ed345b..dcd6fa12493a 100644 --- a/cpp/kernels/fmha_v2/setup.py +++ b/cpp/kernels/fmha_v2/setup.py @@ -2233,6 +2233,7 @@ def selected_mask_types(kspec): def get_kernel_code(kspec, kname, lname): + min_cuda_version = 0 # no restriction # The architecture that determines the instruction. @@ -3333,7 +3334,11 @@ def use_cubin_header(sm, dtype, output_dtype=None, enable_skip_softmax=False, - attention_mask_type=None): + attention_mask_type=None, + sage_block_sizes=None): + # SageAttention is warp-specialized + TMA only. + if sage_block_sizes: + return True if enable_skip_softmax: return False if 'e4m3' in dtype and output_dtype in ['bf16', 'fp16']: @@ -3355,7 +3360,7 @@ def get_cubin_header(kernel_traits, specs_names): if '_bidirectional_sliding_window' in kname else None if generate_cu_trtllm and not use_cubin_header( kspec.sm, kspec.head_size, kspec.dtype, kspec.output_dtype, - kspec.enable_skip_softmax, mask_type): + kspec.enable_skip_softmax, mask_type, kspec.sage_block_sizes): continue name = fname.replace('.', '_') # No `extern "C"` -- the build-time INCBIN aggregator @@ -3429,8 +3434,10 @@ def get_cubin_header(kernel_traits, specs_names): toks.pop(-5) toks.pop(-4) toks.pop(-3) + has_sage = True else: sage_block_sizes = (0, 0, 0) + has_sage = False head_size = toks[-3] if 'x' in head_size: (head_size, head_size_v) = head_size.split('x') @@ -3468,6 +3475,12 @@ def get_cubin_header(kernel_traits, specs_names): if output_prec is None: output_prec = prec + # SageAttention runs int8 Q/K with an e4m3 PV. The metadata carries a single input data + # type, so record the composite one: it still selects the kernel, and it tells a reader + # this is not a plain fp8 kernel. + if has_sage: + prec = 'KV_INT8_E4M3' + is_il = pythonBoolean2cpp['_il' in kname] attention_mask_type = AttentionMaskType.PADDING is_tiled = pythonBoolean2cpp['_tiled' in kname] @@ -3535,7 +3548,7 @@ def get_lname_from_kname(kname: str) -> str: if use_cubin_header(int(sm), int(head_size), prec.lower(), output_prec.lower(), enable_skip_softmax, - attention_mask_type): + attention_mask_type, has_sage): return 'nullptr' lname = kname.replace('_kernel', '') mask_types = [ @@ -3558,7 +3571,8 @@ def get_lname_from_kname(kname: str) -> str: {is_alibi_supported}, {is_tiled}, {has_softcapping_scale}, {return_softmax_stats_flag}, {enable_skip_softmax_flag}, {lname}}}\ '''.format(**locals()) if use_cubin_header( int(sm), int(head_size), prec.lower(), output_prec.lower(), - enable_skip_softmax, attention_mask_type) else '''\ + enable_skip_softmax, attention_mask_type, + has_sage) else '''\ {{ DATA_TYPE_{prec}, DATA_TYPE_{output_prec}, {seq_len}, {q_step}, {kv_step}, {head_size}, {head_size_v}, \ {sage_block_sizes[0]}, {sage_block_sizes[1]}, {sage_block_sizes[2]}, kSM_{sm}, nullptr, \ 0, \"{kname}\", {smem}, {threads}, {meta_unroll_step}, {attention_mask_type_value}, \ @@ -3848,6 +3862,8 @@ class Launch_params; constexpr int32_t kSM_121 = 121; // FIXME: These are duplicated declarations, we should remove them in the future. +// The enumerator order has to match kernels/multiHeadAttentionCommon.h, since both describe the +// same values in one binary. enum Data_type {{ DATA_TYPE_BOOL, @@ -3859,7 +3875,11 @@ class Launch_params; DATA_TYPE_BF16, DATA_TYPE_E2M1, DATA_TYPE_E4M3, - DATA_TYPE_E5M2 + DATA_TYPE_E5M2, + // Composite kv data types + DATA_TYPE_KV_FP16_E4M3, + DATA_TYPE_KV_BF16_E4M3, + DATA_TYPE_KV_INT8_E4M3 }}; struct FusedMultiHeadAttentionKernelMetaInfoV2 @@ -3965,6 +3985,10 @@ def generate_files(specs_names): kfiles = [] valid_specs_names = [] + # SageAttention ships as pre-built cubins, so its specs contribute only the cubin extern + # declarations and the metadata row emitted by get_cubin_header, and are threaded straight to + # it alongside the generated kernels. + for kspec, fname, lname, kname in specs_names: code = get_kernel_code(kspec, kname, lname) # some kernels are skipped when generating cubins for trt-llm. @@ -6711,12 +6735,13 @@ def enumerate_kernels(): output_dtype="bf16", enable_skip_softmax=enable_skip_softmax) - # For now SageAttention only needs BF16 - # block_size_q should be divisible by 64 - # block_size_k should be divisible by 8 - # block_size_v should be divisible by 32 - for sage_block_sizes in [(64, 64, 64), (64, 64, 128), (64, 64, 256), - (64, 128, 64), (64, 128, 128), (64, 128, 256)]: + # SM90 SageAttention scales the tokens a single thread owns, so exactly one combination is + # supported. All three entries are block sizes, each counted along the axis quantization runs + # over: + # block_q = 2 the 2 query rows a thread owns (strided: rows r and r + 8) + # block_k = 16 the 16 keys of one of a thread's four sub-fragments (strided) + # block_v = 1 one channel, since V is quantized along the channel axis + for sage_block_sizes in [(2, 16, 1)]: enumerate_qgmma_flash_warpspec_kernels( specs, sm=90, @@ -6747,17 +6772,6 @@ def enumerate_kernels(): dtype='e4m3_fp32', head_sizes=[192, 576], output_dtype="bf16") - # Sage Attention on Ada only supports block_size = (64, 32, 32) - enumerate_qmma_flash_kernels(specs, - sm=89, - dtype='e4m3_fp32', - sage_block_sizes=(64, 32, 32), - output_dtype="bf16") - enumerate_qmma_flash_kernels(specs, - sm=89, - dtype='e4m3_fp32', - sage_block_sizes=(64, 32, 32), - output_dtype="fp16") enumerate_imma_kernels(specs, sm=89) enumerate_hmma_kernels(specs, sm=89, dtype='fp16') @@ -6939,27 +6953,19 @@ def enumerate_kernels(): and ((kspec.warp_specialization == True and kspec.alibi == False) # sm90 or (kspec.warp_specialization == False and kspec.tiled == True)) # non-sm90 and kspec.enable_attn_logit_softcapping == False) - # SageAttention (warp_spec, head_size in (80, 128), packed QKV, padding mask) + # SageAttention (warp_spec, separate Q/K/V, padding mask). SageAttention + # quantizes Q, K and V separately, so separate Q/K/V is the layout it is exported + # for; the head sizes listed are the ones we ship a cubin for. or (kspec.sm == 90 - and kspec.head_size in [80, 128] + and kspec.head_size in [128] and kspec.version == 2 - and kspec.sage_block_sizes in [(64, 64, 256)] + and kspec.sage_block_sizes in [(2, 16, 1)] and kspec.cross_mha == False and kspec.flash_attention == True and kspec.warp_specialization == True - and kspec.input_layout == InputLayout.PACKED_QKV + and kspec.input_layout == InputLayout.SEPARATE_Q_K_V and kspec.alibi == False - and kspec.enable_attn_logit_softcapping == False) - # SageAttention on Ada (head_size in (80, 128), packed QKV, padding mask) - or (kspec.sm == 89 - and kspec.head_size in [80, 128] - and kspec.sage_block_sizes in [(64, 32, 32)] - and kspec.output_dtype in ['fp16', 'bf16'] - and kspec.version == 2 - and kspec.cross_mha == False - and kspec.flash_attention == True - and kspec.warp_specialization == False - and kspec.input_layout == InputLayout.PACKED_QKV)) + and kspec.enable_attn_logit_softcapping == False)) # only generate head_size = 128/256 for attn_logit_softcapping operation. and (kspec.head_size == 128 or kspec.head_size == 256 or not kspec.enable_attn_logit_softcapping)] # yapf: enable diff --git a/cpp/kernels/fmha_v2/src/fmha/warpspec/compute.h b/cpp/kernels/fmha_v2/src/fmha/warpspec/compute.h index f1d65f861f8d..48814603c57b 100644 --- a/cpp/kernels/fmha_v2/src/fmha/warpspec/compute.h +++ b/cpp/kernels/fmha_v2/src/fmha/warpspec/compute.h @@ -148,25 +148,39 @@ struct Compute static constexpr int SAGE_BLOCK_SIZE_Q = Kernel_traits::SAGE_BLOCK_SIZE_Q; - // sanitize 0 to -1, avoid DIV BY ZERO below - static constexpr int SAGE_BLOCK_SIZE_K - = Kernel_traits::SAGE_BLOCK_SIZE_K > 0 ? Kernel_traits::SAGE_BLOCK_SIZE_K : -1; + static constexpr int SAGE_BLOCK_SIZE_K = Kernel_traits::SAGE_BLOCK_SIZE_K; - static constexpr int SAGE_BLOCK_SIZE_V - = Kernel_traits::SAGE_BLOCK_SIZE_V > 0 ? Kernel_traits::SAGE_BLOCK_SIZE_V : -1; + static constexpr int SAGE_BLOCK_SIZE_V = Kernel_traits::SAGE_BLOCK_SIZE_V; - // BLOCK_SIZE_Q should be multiply of STEP_Q (usually 64) so that q scale can be fused into scale_bmm1 - static_assert(SAGE_BLOCK_SIZE_Q < 0 || SAGE_BLOCK_SIZE_Q % STEP_Q == 0); - static_assert(SAGE_BLOCK_SIZE_K < 0 || SAGE_BLOCK_SIZE_K % 8 == 0); // 8 = columns of a gmma CORE - static_assert(SAGE_BLOCK_SIZE_V < 0 || SAGE_BLOCK_SIZE_V % 32 == 0); // 32 = K dimension of a qgmma + // SageAttention groups tokens the way the BMM1 fragment distributes them across threads, so a + // thread never needs more than one scale per axis. Q: the 2 rows a thread owns (quad_row and + // quad_row + 8, strided). K: one of the SAGE_K_CHUNKS sub-fragments of a thread's key columns. + static constexpr int SAGE_Q_TOKENS_PER_SCALE = 2; - // SAGE_BLOCKS_PER_STEP_X is used to declare scale buffer like `float scales_k[SAGE_BLOCKS_PER_STEP_K];` - // if SAGE_BLOCKS_PER_STEP_X == 0, you will get `zero-sized variable is not allowed in device code` - // error from nvcc, so the minimal value have to be 1. But don't worry, unused local variables will - // be optimized out by compiler. - static constexpr int SAGE_BLOCKS_PER_STEP_K = std::max(STEP_KV / SAGE_BLOCK_SIZE_K, 1); + static constexpr int SAGE_K_TOKENS_PER_SCALE = STEP_KV / (4 * Kernel_traits::SAGE_K_CHUNKS); - static constexpr int SAGE_BLOCKS_PER_STEP_V = std::max(STEP_KV / SAGE_BLOCK_SIZE_V, 1); + static constexpr int SAGE_Q_SCALES_PER_TILE = STEP_Q / SAGE_Q_TOKENS_PER_SCALE; + + static constexpr int SAGE_K_SCALES_PER_TILE = 4 * Kernel_traits::SAGE_K_CHUNKS; + + // The scale buffers are laid out as (H, scales for the whole batch) rather than + // (B, H, max_nblock): the caller sizes its workspace when only the batch size and the total + // token count are known, not max_seqlen. A sequence's scales begin at a base derived from its + // cu_seqlens entry, and each sequence reserves one spare tile so that the scales of a partial + // final tile -- which the kernel indexes in full -- cannot run into the next sequence. + static constexpr int SAGE_Q_SLOTS_PER_SEQ = SAGE_Q_SCALES_PER_TILE; + + // K reserves 4 more so its base can be rounded up to a multiple of 4, which keeps the float4 + // scale load aligned. Rounding up costs at most 3, and one spare tile plus 4 covers that. + static constexpr int SAGE_K_SLOTS_PER_SEQ = SAGE_K_SCALES_PER_TILE + 4; + + static_assert(SAGE_K_SLOTS_PER_SEQ % 4 == 0, "the K scale base must stay 16B-aligned"); + static_assert(SAGE_BLOCK_SIZE_Q <= 0 || SAGE_BLOCK_SIZE_Q == SAGE_Q_TOKENS_PER_SCALE, + "sage block_q must be 2: Q is quantized per thread (2 strided rows)"); + static_assert(SAGE_BLOCK_SIZE_K <= 0 || SAGE_BLOCK_SIZE_K == SAGE_K_TOKENS_PER_SCALE, + "sage block_k must be 16: K is quantized per thread sub-fragment"); + static_assert(SAGE_BLOCK_SIZE_V <= 0 || Kernel_traits::SAGE_V_PER_CHANNEL, + "V is quantized along the channel axis; sage block_v must be 1 or 0"); #define K_TILE_WAIT() \ int ready_k = cbr_k.peek(); \ @@ -186,7 +200,7 @@ struct Compute actual_kv_seqlen, alibi_head_scale, \ USE_CUSTOM_MASK ? (head_info.mask_sum_s + q_step_idx * STEP_Q + local_q_tile_offset) \ : (q_step_idx * STEP_Q + head_info.q_tile_offset), \ - kv_step_idx * STEP_KV, sage_scale_row, cbr, cbr_v, mutex_accessor, \ + kv_step_idx * STEP_KV, sage_scale_base_kv, cbr, cbr_v, mutex_accessor, \ &shared->skip_softmax_votes[kv_step_idx & 1][warpgroup_id], kv_step_idx == kv_idx_end - 1); //////////////////////////////////////////////////////////////////////////////////////////////// @@ -313,11 +327,21 @@ struct Compute // Calculate the alibi head_scaling_factor. float alibi_head_scale = APPLY_ALIBI ? get_alibi_head_scaling_factor(head_info.bidh, params.alibi_params) : 0.f; - // pre-compute the row of the scale for reuse - int sage_scale_row; + // pre-compute where this (sequence, head) pair's scales start, for reuse below. Q is + // quantized per query head, K per key/value head: in MQA/GQA several query heads share + // one KV head. max_nblock is the per-head stride of the whole batch's scales. + int sage_scale_base = 0; + int sage_scale_base_kv = 0; if constexpr (Kernel_traits::SAGE_ATTENTION) { - sage_scale_row = head_info.bidb * params.h + head_info.bidh; + int const bidb = head_info.bidb; + sage_scale_base = head_info.bidh * params.sage.q.max_nblock + + params.cu_q_seqlens[bidb] / SAGE_Q_TOKENS_PER_SCALE + bidb * SAGE_Q_SLOTS_PER_SEQ; + // Round the dense part up so the base is a multiple of 4; SAGE_K_SLOTS_PER_SEQ is + // one too, so every sequence's base stays 16B-aligned. + int const k_dense = params.cu_kv_seqlens[bidb] / SAGE_K_TOKENS_PER_SCALE; + sage_scale_base_kv = (head_info.bidh / params.h_q_per_kv) * params.sage.k.max_nblock + + ((k_dense + 3) & ~3) + bidb * SAGE_K_SLOTS_PER_SEQ; } // BMM2 epilogue @@ -339,7 +363,12 @@ struct Compute // to avoid frequent `__ldg`. But experiment shows that the current one is faster. // A bit counterintuitive. auto const scale_bmm1 = params.scale_bmm1_d ? __ldg(params.scale_bmm1_d) : params.scale_bmm1; - int const idx = sage_scale_row * params.sage.q.max_nblock + q_offset / SAGE_BLOCK_SIZE_Q; + // One scale per thread: it covers exactly the two rows this thread owns. All + // four lanes of a quad share quad_row and therefore read the same scale, which + // keeps the row-max shuffle reduction consistent. Derived from tidx because + // quad_row_ is only initialised for causal/sliding kernels. + int const q_pair = (tidx / 32) * 8 + (tidx % 32) / 4; + int const idx = sage_scale_base + (q_offset / STEP_Q) * SAGE_Q_SCALES_PER_TILE + q_pair; *(float*) (&softmax.scale_bmm1_) = reinterpret_cast(scale_bmm1) * __ldg(¶ms.sage.q.scales[idx]); } @@ -444,6 +473,51 @@ struct Compute } if (valid_run) { + // The per-channel V scale is constant along the BMM2 reduction dimension, so it never has + // to interrupt the QGMMA pipeline: it is applied once here, per output column. The + // accumulator N index is permuted relative to the head dimension; the column mapping below + // mirrors Gmem_tile_o_qgmma_fp32_16bits::store(). + if constexpr (Kernel_traits::SAGE_V_PER_CHANNEL) + { + // Layout is (H_kv, D): the amax is reduced over every token of every sequence, so + // unlike the Q/K scales it carries no batch dimension. + float const* v_scales_ch + = params.sage.v.scales + (size_t) (head_info.bidh / params.h_q_per_kv) * params.dv; + int const lane_quad = tidx % 4; +#pragma unroll + for (int mma_ni = 0; mma_ni < Mma_tile_o::MMAS_N; mma_ni++) + { + // Each even/odd core pair covers 4 consecutive output columns starting at col_base, + // so the scales come in as one 16B load. The (ni, ei) -> column order is exactly + // the (x, y, z, w) order of store(). +#pragma unroll + for (int ni = 0; ni < Mma_tile_o::CORES_N; ni += 2) + { + int const col_base = mma_ni * Mma_tile_o::CORES_N * 8 + ni * 8 + lane_quad * 4; + float4 scale4 = {1.f, 1.f, 1.f, 1.f}; + if (col_base + 4 <= params.dv) + { + scale4 = __ldg(reinterpret_cast(v_scales_ch + col_base)); + } + float const scale_ch[2][2] = {{scale4.x, scale4.z}, {scale4.y, scale4.w}}; +#pragma unroll + for (int nj = 0; nj < 2; nj++) + { +#pragma unroll + for (int ei = 0; ei < 2; ei++) + { +#pragma unroll + for (int mi = 0; mi < Mma_tile_o::CORES_M; mi++) + { + ctile_o.acc_[0][mma_ni].elt( + 2 * (ni + nj) * Mma_tile_o::CORES_M + 2 * mi + ei) + *= scale_ch[nj][ei]; + } + } + } + } + } + } // Final step's update. tile_o_epilogue.scale(ctile_o, p_max, p_sum); // Store o_tile to gmem. @@ -478,8 +552,9 @@ struct Compute inline __device__ void compute_single_tile(Params params, Compute_tile_p& ctile_p, Softmax& softmax, Compute_tile_o& ctile_o, float (&p_max)[Mma_tile_p::CORES_M], float (&p_sum)[Mma_tile_p::CORES_M], int const tidx, int const actual_kv_seqlen, float const alibi_head_scale, int const row_offset, - int const col_offset, int const sage_scale_row, Circular_buffer_q_reader& cbr, Circular_buffer_kv_reader& cbr_v, - OrderedMutexAccessor& mutex, uint32_t* skip_softmax_vote, bool complete = false) + int const col_offset, int const sage_scale_base_kv, Circular_buffer_q_reader& cbr, + Circular_buffer_kv_reader& cbr_v, OrderedMutexAccessor& mutex, uint32_t* skip_softmax_vote, + bool complete = false) { // Skip-softmax vote initialization @@ -489,27 +564,36 @@ struct Compute *skip_softmax_vote = 1; } // load the scales of K/V from global memory -#define LOAD_SCALES_KV(dst, which, blocks_per_step, block_size) \ - if constexpr (block_size > 0) \ +// Load this thread's K scales for one STEP_KV tile. They are stored swizzled as +// [tile][quad_col][chunk], so the SAGE_K_CHUNKS scales a thread needs are contiguous and come in as +// a single 128-bit access. +#define LOAD_SCALES_K(dst) \ + if constexpr (SAGE_BLOCK_SIZE_K > 0) \ { \ - const int _start = col_offset / block_size; \ - const float* _src = params.sage.which.scales + sage_scale_row * params.sage.which.max_nblock + _start; \ - const int _end = params.sage.which.max_nblock - _start; \ - _Pragma("unroll") for (int _i = 0; _i < blocks_per_step; _i++) \ + const float* _src = params.sage.k.scales + sage_scale_base_kv \ + + (col_offset / STEP_KV) * SAGE_K_SCALES_PER_TILE + (tidx % 4) * Kernel_traits::SAGE_K_CHUNKS; \ + if constexpr (Kernel_traits::SAGE_K_CHUNKS == 4) \ { \ - dst[_i] = _i < _end ? _src[_i] : 1.0f; \ + const float4 _s4 = __ldg(reinterpret_cast(_src)); \ + dst[0] = _s4.x; \ + dst[1] = _s4.y; \ + dst[2] = _s4.z; \ + dst[3] = _s4.w; \ + } \ + else \ + { \ + _Pragma("unroll") for (int _g = 0; _g < Kernel_traits::SAGE_K_CHUNKS; _g++) \ + { \ + dst[_g] = __ldg(_src + _g); \ + } \ } \ } -#define LOAD_SCALES_K(scales) LOAD_SCALES_KV(scales, k, SAGE_BLOCKS_PER_STEP_K, SAGE_BLOCK_SIZE_K) - -#define LOAD_SCALES_V(scales) LOAD_SCALES_KV(scales, v, SAGE_BLOCKS_PER_STEP_V, SAGE_BLOCK_SIZE_V) - // Load the needed packed masks. softmax.load_packed_mask(row_offset, col_offset); // experiments show that here is the best place to load scales of K - float scales_k[SAGE_BLOCKS_PER_STEP_K]; + float scales_k[Kernel_traits::SAGE_K_CHUNKS]; LOAD_SCALES_K(scales_k) // Wait until another warpgroup has already executed HGMMA. @@ -574,18 +658,29 @@ struct Compute // Unpack the elements from bmm1 output to floats. softmax.unpack(ctile_p); - // apply the scales of K before softmax + // apply the scales of K before softmax. A chunk is exactly the set of keys sharing one K + // scale, i.e. one of the SAGE_K_CHUNKS sub-fragments of this thread's key columns. + // + // NOTE: the shipped cubin folds this into the exp2f of the softmax instead, which costs a + // few multiplies per tile rather than one per element. This tree only has to agree on the + // ABI, the thread shape and the shared-memory budget, so it takes the simple route. if constexpr (SAGE_BLOCK_SIZE_K > 0) { + constexpr int ELTS_PER_CHUNK = Mma_tile_p::CORES_N * 2 / Kernel_traits::SAGE_K_CHUNKS; + static_assert( + Mma_tile_p::CORES_N * 2 % Kernel_traits::SAGE_K_CHUNKS == 0, "chunk must divide the fragment"); #pragma unroll - for (int ni = 0; ni < Mma_tile_p::CORES_N; ni++) + for (int g = 0; g < Kernel_traits::SAGE_K_CHUNKS; g++) { - float const scale_k = scales_k[SAGE_BLOCKS_PER_STEP_K * ni / Mma_tile_p::CORES_N]; + float const scale_k = scales_k[g]; #pragma unroll - for (int mi = 0; mi < Mma_tile_p::CORES_M; mi++) + for (int j = 0; j < ELTS_PER_CHUNK; j++) { - softmax.elt_[mi][2 * ni] *= scale_k; - softmax.elt_[mi][2 * ni + 1] *= scale_k; +#pragma unroll + for (int mi = 0; mi < Mma_tile_p::CORES_M; mi++) + { + softmax.elt_[mi][g * ELTS_PER_CHUNK + j] *= scale_k; + } } } } @@ -620,10 +715,6 @@ struct Compute return; } - // experiments show that here is the best place to load scales of V - float scales_v[SAGE_BLOCKS_PER_STEP_V]; - LOAD_SCALES_V(scales_v) - // Update flash attention scales and pack it for BMM2 softmax.pack(ctile_o, frag_p); @@ -642,42 +733,6 @@ struct Compute warpgroup_arrive(); - float last_scale_v; - -// Apply the scale of V to partial result. -// Note 2 points: -// 1. Because the matrix V is quantized along the inner dimension, it is necessary to interrupt -// the MMA workflow after processing each BLOCKS_SIZE_V rows of V and scale the intermediate -// results once. For example, STEP_KV=256, qgmma.K=32, then 256/32=8 MMAs are needs, -// so mma_ki = [0,1,2, ..., 7]. If the BLOCK_SIZE_V=64, then after each 2 qgmmas we should scale -// ctile_o. -// 2. The ctile_o is all zero at the beginning. if we directly apply the scale of V after each 2 -// qgmmas, let's see what happens: -// ctile_o = [0] -// ctile_o = (ctile_o + P0 x V0) * s0 = P0 x V0 * s0 -// ctile_o = (ctile_o + P1 x V1) * s1 = P0 x V0 * s0 * s1 + P1 x V1 * s1 -// ctile_o = (ctile_o + P2 x V2) * s2 = P0 x V0 * s0 * s1 * s2 + P1 x V1 * s1 * s2 + P2 x V2 * s2 -// ... -// As you see, the actual scale of a V block is the cumulative product of the scales of all -// later blocks. To solve this, we have to preprocess the scale s[i] of block[i] to s[i]/s[i+1], -// and the final block uses the actual scale. -// But to fetch the next scale in next STEP leads to bad performance. So we apply s[i-1]/s[i] to -// current partial result BEFORE each V block. -#define APPLY_SCALE_V(mma_ki) \ - if constexpr (SAGE_BLOCK_SIZE_V > 0) \ - { \ - if (mma_ki % (Mma_tile_o::MMAS_K / SAGE_BLOCKS_PER_STEP_V) == 0) \ - { \ - float _scale_v = scales_v[SAGE_BLOCKS_PER_STEP_V * mma_ki / Mma_tile_o::MMAS_K]; \ - if (mma_ki != 0) \ - { \ - warpgroup_commit(); \ - warpgroup_wait<0>(); \ - } \ - last_scale_v = _scale_v; \ - } \ - } - // BMM2 (S * V). #pragma unroll for (int kbi = 0; kbi < BMM2_MMAS_K_GROUPS - 1; kbi++) @@ -686,7 +741,6 @@ struct Compute for (int ki = 0; ki < BMM2_MMAS_K_PER_GROUP; ++ki) { int const mma_ki = kbi * BMM2_MMAS_K_PER_GROUP + ki; - APPLY_SCALE_V(mma_ki) ctile_o.fill_frag_a(frag_p[mma_ki]); ctile_o.compute(ki, false, ki == BMM2_MMAS_K_PER_GROUP - 1); } @@ -697,12 +751,10 @@ struct Compute for (int ki = 0; ki < BMM2_MMAS_K_PER_GROUP - 1; ++ki) { int const mma_ki = (BMM2_MMAS_K_GROUPS - 1) * BMM2_MMAS_K_PER_GROUP + ki; - APPLY_SCALE_V(mma_ki) ctile_o.fill_frag_a(frag_p[mma_ki]); ctile_o.compute(ki); } - APPLY_SCALE_V((Mma_tile_o::MMAS_K - 1)) ctile_o.fill_frag_a(frag_p[Mma_tile_o::MMAS_K - 1]); ctile_o.compute(Mma_tile_o::MMAS_K - 1, true, true); diff --git a/cpp/kernels/fmha_v2/src/fmha/warpspec/dma.h b/cpp/kernels/fmha_v2/src/fmha/warpspec/dma.h index 1d8440592267..eb32a4006ad7 100644 --- a/cpp/kernels/fmha_v2/src/fmha/warpspec/dma.h +++ b/cpp/kernels/fmha_v2/src/fmha/warpspec/dma.h @@ -557,7 +557,7 @@ struct DMA \ int v_barrier_id; \ void* v_barrier_ptr; \ - typename Kernel_traits::Element_data_type* v_smem; \ + typename Kernel_traits::Element_data_type_o* v_smem; \ \ if constexpr (DMA_GROUP_TRANSPOSE_V) \ { \ diff --git a/cpp/kernels/fmha_v2/src/fmha/warpspec/epilogue.h b/cpp/kernels/fmha_v2/src/fmha/warpspec/epilogue.h index 1c9e786cd80b..55054e279043 100644 --- a/cpp/kernels/fmha_v2/src/fmha/warpspec/epilogue.h +++ b/cpp/kernels/fmha_v2/src/fmha/warpspec/epilogue.h @@ -1093,10 +1093,14 @@ struct Softmax float tmp_17 = this->elt_[1][8 * ni + 7]; // +25 // Pack to 4 registers. - frag_p[ni].reg(0) = fmha::float4_to_fp8x4(tmp_00, tmp_01, tmp_02, tmp_03); - frag_p[ni].reg(1) = fmha::float4_to_fp8x4(tmp_10, tmp_11, tmp_12, tmp_13); - frag_p[ni].reg(2) = fmha::float4_to_fp8x4(tmp_04, tmp_05, tmp_06, tmp_07); - frag_p[ni].reg(3) = fmha::float4_to_fp8x4(tmp_14, tmp_15, tmp_16, tmp_17); + frag_p[ni].reg(0) + = fmha::float4_to_fp8x4(tmp_00, tmp_01, tmp_02, tmp_03); + frag_p[ni].reg(1) + = fmha::float4_to_fp8x4(tmp_10, tmp_11, tmp_12, tmp_13); + frag_p[ni].reg(2) + = fmha::float4_to_fp8x4(tmp_04, tmp_05, tmp_06, tmp_07); + frag_p[ni].reg(3) + = fmha::float4_to_fp8x4(tmp_14, tmp_15, tmp_16, tmp_17); } if (!IS_FIRST_COL) diff --git a/cpp/kernels/fmha_v2/src/fmha/warpspec/kernel_traits.h b/cpp/kernels/fmha_v2/src/fmha/warpspec/kernel_traits.h index 93a8b999ce67..52bbb3b50c99 100644 --- a/cpp/kernels/fmha_v2/src/fmha/warpspec/kernel_traits.h +++ b/cpp/kernels/fmha_v2/src/fmha/warpspec/kernel_traits.h @@ -32,6 +32,29 @@ namespace ws //////////////////////////////////////////////////////////////////////////////////////////////////// +// Picks the BMM1 instruction traits for a warp-specialized kernel whose BMM2 is QGMMA e4m3. +// SageAttention is INT8 QK with e4m3 PV: BMM1 is IGMMA int8 -> int32. Plain FP8 attention runs +// both GEMMs in e4m3. A template-template argument cannot be chosen with std::conditional_t, +// hence this selector. +template +struct Bmm1_traits_selector; + +template <> +struct Bmm1_traits_selector +{ + template + using type = Hopper_qgmma_e4m3_fp32_traits; +}; + +template <> +struct Bmm1_traits_selector +{ + template + using type = Hopper_igmma_int8_int32_traits; +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + template < // The instruction trait template for initializing BMM1 and BMM2 traits. template class Instruction_traits, @@ -77,10 +100,17 @@ template < // The output type (only used by fp8 kernels). typename OutputType = typename Instruction_traits::A_type, // The sage attention block size for Q, K and V - int SAGE_BLOCK_SIZE_Q_ = 0, int SAGE_BLOCK_SIZE_K_ = 0, int SAGE_BLOCK_SIZE_V_ = 0> + int SAGE_BLOCK_SIZE_Q_ = 0, int SAGE_BLOCK_SIZE_K_ = 0, int SAGE_BLOCK_SIZE_V_ = 0, + // The BMM2 instruction traits. Defaults to the BMM1 ones; SageAttention overrides BMM1 to + // IGMMA int8 while BMM2 stays QGMMA e4m3, so the two differ there. + template class Instruction_traits_o = Instruction_traits> struct Kernel_traits { + // The BMM2 input element type. Every V/O shared-memory buffer is sized from this rather than + // from the BMM1 element type, which SageAttention makes int8. + using Element_data_type_o = typename Instruction_traits_o::A_type; + // The step size in query sequence dimension (M of BMM1 and BMM2). static constexpr int STEP_Q = STEP_Q_; @@ -135,9 +165,18 @@ struct Kernel_traits static constexpr int SAGE_BLOCK_SIZE_V = SAGE_BLOCK_SIZE_V_; + // SageAttention quantizes with the groups the BMM1 fragment hands each thread. A thread owns + // STEP_KV / 4 keys, so a block of SAGE_BLOCK_SIZE_K_ keys splits them into this many chunks. + static constexpr int SAGE_K_CHUNKS = SAGE_BLOCK_SIZE_K_ > 0 ? STEP_KV_ / (4 * SAGE_BLOCK_SIZE_K_) : 1; + + // SAGE_BLOCK_SIZE_V_ counts channels, not tokens: 1 means one scale per channel. V is + // quantized along the channel (head-dim) axis, which is the N dimension of BMM2, so the scale + // is constant along the BMM2 reduction and applies once after the wgmma chain. + static constexpr int SAGE_V_PER_CHANNEL = SAGE_BLOCK_SIZE_V_ == 1; + // Whether the dma group transposes the v tile explicitly. - static constexpr int DMA_GROUP_TRANSPOSE_V = (std::is_same::value - || std::is_same::value); + static constexpr int DMA_GROUP_TRANSPOSE_V = (std::is_same::value + || std::is_same::value); // The number of smem scratch buffer for staging V transpose for Hopper QGMMA static constexpr int V_SCRATCH_BUFFERS = DMA_GROUP_TRANSPOSE_V ? 1 : 0; @@ -262,7 +301,7 @@ struct Kernel_traits // The instruction traits for the BMM2. // FP16/BF16 K = 16, FP8 K = 32. - using Traits_o = Instruction_traits; + using Traits_o = Instruction_traits_o; // The CTA description for BMM1. using Cta_tile_p = @@ -313,9 +352,9 @@ struct Kernel_traits // The q, k, v tile buffer. using Buffer_q_t = cuda::std::array; using Buffer_k_t = cuda::std::array; - using Buffer_v_t = cuda::std::array; + using Buffer_v_t = cuda::std::array; // We need one kv buffer to explicitly transose fp8 smem_tile. - using Buffer_v_scratch_t = cuda::std::array; + using Buffer_v_scratch_t = cuda::std::array; // The smem bytes of q, k, v tiles. static constexpr int SMEM_BYTES_Q = sizeof(Buffer_q_t); @@ -454,18 +493,21 @@ template < // The step size in query sequence dimension (M of BMM1 and BMM2). // The sage attention block size for Q, K and V int SAGE_BLOCK_SIZE_Q_ = 0, int SAGE_BLOCK_SIZE_K_ = 0, int SAGE_BLOCK_SIZE_V_ = 0> struct Kernel_traits_Hopper_qgmma_e4m3_fp32 - : public Kernel_traits + : public Kernel_traits 0 || SAGE_BLOCK_SIZE_K_ > 0 + || SAGE_BLOCK_SIZE_V_ > 0)>::template type, + STEP_Q_, STEP_KV_, D_, DV_, Q_BUFFERS_, KV_BUFFERS_, NUM_COMPUTE_GROUPS_, DMA2COMPUTE_DEPTH_, + ATTENTION_MASK_TYPE_, HEADS_INTERLEAVED_, APPLY_ALIBI_, ENABLE_MUTEX_, SCHEDULING_MODE_, INPUT_LAYOUT_, + USE_TMA_STORE_, ENABLE_BMM1_SOFTCAPPING_SCALE_, RETURN_SOFTMAX_STATS_, ENABLE_SKIP_SOFTMAX_, OutputType, + SAGE_BLOCK_SIZE_Q_, SAGE_BLOCK_SIZE_K_, SAGE_BLOCK_SIZE_V_, Hopper_qgmma_e4m3_fp32_traits> { // Base class. - using Base = Kernel_traits; + using Base = Kernel_traits 0 || SAGE_BLOCK_SIZE_K_ > 0 + || SAGE_BLOCK_SIZE_V_ > 0)>::template type, + STEP_Q_, STEP_KV_, D_, DV_, Q_BUFFERS_, KV_BUFFERS_, NUM_COMPUTE_GROUPS_, DMA2COMPUTE_DEPTH_, + ATTENTION_MASK_TYPE_, HEADS_INTERLEAVED_, APPLY_ALIBI_, ENABLE_MUTEX_, SCHEDULING_MODE_, INPUT_LAYOUT_, + USE_TMA_STORE_, ENABLE_BMM1_SOFTCAPPING_SCALE_, RETURN_SOFTMAX_STATS_, ENABLE_SKIP_SOFTMAX_, OutputType, + SAGE_BLOCK_SIZE_Q_, SAGE_BLOCK_SIZE_K_, SAGE_BLOCK_SIZE_V_, Hopper_qgmma_e4m3_fp32_traits>; static constexpr int USE_TMA_STORE = USE_TMA_STORE_; @@ -500,7 +542,7 @@ struct Kernel_traits_Hopper_qgmma_e4m3_fp32 using Buffer_v_scratch_t = typename Base::Buffer_v_scratch_t; // Extra O buffer if TMA is used for epilogue using Element_data_type = typename Base::Element_data_type; - using Buffer_o_t = cuda::std::array; + using Buffer_o_t = cuda::std::array; // The struct of shared memory buffers. struct __align__(128) Shared diff --git a/cpp/kernels/fmha_v2/src/fused_multihead_attention.cpp b/cpp/kernels/fmha_v2/src/fused_multihead_attention.cpp index f4ec62cd032e..def12de99bbd 100644 --- a/cpp/kernels/fmha_v2/src/fused_multihead_attention.cpp +++ b/cpp/kernels/fmha_v2/src/fused_multihead_attention.cpp @@ -21,6 +21,7 @@ #include #include #include +#include #include #include #include @@ -80,15 +81,6 @@ void run_conversion_fp32_to_e4m3(void* dst, void const* src, int s, int b, int h //////////////////////////////////////////////////////////////////////////////////////////////////// -void run_sage_quant(unsigned int batch_size, unsigned int head_num, unsigned int head_size, unsigned int max_seq_len, - // device var - void const* q, void const* k, void const* v, int stride_q, int stride_k, int stride_v, int const* cu_seqlens_q, - int const* cu_seqlens_kv, int block_size_q, int block_size_k, int block_size_v, - // output - void* quant_q, void* quant_k, void* quant_v, float* scales_q, float* scales_k, float* scales_v); - -//////////////////////////////////////////////////////////////////////////////////////////////////// - void ground_truth(RefBMM& bmm1, RefBMM& bmm2, const Data_type data_type, const Data_type acc_type, float const scale_bmm1, float const scale_softmax, float const scale_bmm2, float const softcapping_scale_bmm1, void* qkv_d, void* vt_d, void* mask_d, void* attention_sinks_d, void* p_d, void* s_d, void* tmp_d, void* o_d, @@ -1818,63 +1810,163 @@ int main(int argc, char** argv) if (sage_block_size_q > 0 || sage_block_size_k > 0 || sage_block_size_v > 0) { - assert(input_layout == Attention_input_layout::PACKED_QKV && "for now this test only supports PACKED_QKV"); + // SageAttention targets separate Q/K/V, where the q and kv sequence lengths may differ. + // Packed QKV is still supported; it forces s_q == s_kv because Q, K and V share a token. + bool const separate_qkv = input_layout == Attention_input_layout::SEPARATE_Q_K_V; + assert((separate_qkv || input_layout == Attention_input_layout::PACKED_QKV) + && "SageAttention supports the separate Q/K/V and packed QKV layouts"); + assert((separate_qkv || s_q == s) && "packed QKV shares one token between Q and KV, so it needs s_q == s"); assert(d == dv && "for now SageAttention doesn't support different QKV dims"); - assert(((sm == 90 && !force_non_warp_specialization) || (sm == 89)) - && "only hopper and ada kernels support SageAttention"); - fmha::e4m3_t* quant_qkv; - FMHA_CHECK_CUDA(cudaMalloc((void**) &quant_qkv, qkv_packed_size)); - params_v2.sage.q.block_size = sage_block_size_q; - params_v2.sage.q.max_nblock = (s + sage_block_size_q - 1) / sage_block_size_q; - FMHA_CHECK_CUDA( - cudaMalloc((void**) ¶ms_v2.sage.q.scales, params_v2.sage.q.max_nblock * h * b * sizeof(float))); - params_v2.sage.k.block_size = sage_block_size_k; - params_v2.sage.k.max_nblock = (s + sage_block_size_k - 1) / sage_block_size_k; - FMHA_CHECK_CUDA( - cudaMalloc((void**) ¶ms_v2.sage.k.scales, params_v2.sage.k.max_nblock * h * b * sizeof(float))); - params_v2.sage.v.block_size = sage_block_size_v; - params_v2.sage.v.max_nblock = (s + sage_block_size_v - 1) / sage_block_size_v; - FMHA_CHECK_CUDA( - cudaMalloc((void**) ¶ms_v2.sage.v.scales, params_v2.sage.v.max_nblock * h * b * sizeof(float))); -#if 1 + assert(sm == 90 && !force_non_warp_specialization + && "only the hopper warp-specialized kernels support SageAttention"); + assert(attention_mask_type == Attention_mask_type::PADDING + && "SageAttention supports the padding mask only (non-causal)"); + + sage_quant::Geometry const sage_geo(sage_block_size_k); + assert(sage_block_size_q == sage_geo.q_tokens_per_scale + && "-sage-block-q must be 2: Q is quantized per thread (2 rows)"); + assert(sage_block_size_k == sage_geo.k_tokens_per_scale + && "-sage-block-k must be 16: K is quantized per thread sub-fragment (16 keys)"); + assert((sage_block_size_v == 0 || sage_block_size_v == 1) + && "V is quantized along the channel axis; -sage-block-v must be 1 or 0"); + + sage_quant::Input sage_in; + sage_in.b = b; + sage_in.h = h; + sage_in.h_kv = h_kv; + sage_in.d = d; + sage_in.dv = dv; + sage_in.cu_q_seqlens = cu_q_seqlens.data(); + sage_in.cu_kv_seqlens = cu_seqlens.data(); + sage_in.block_size_q = sage_block_size_q; + sage_in.block_size_k = sage_block_size_k; + sage_in.block_size_v = sage_block_size_v; + + // Quantized bytes. Separate Q/K/V gets three buffers matching the layout the kernel reads + // (q [token][h][d], k [token][h_kv][d], v [token][h_kv][dv]); packed QKV gets one, with the + // three tensors aliasing it at different head offsets. + size_t const total_q = cu_q_seqlens.back(), total_kv = cu_seqlens.back(); + std::vector q_f, k_f, v_f; + std::vector q_b, k_b, v_b, packed_b; + + if (separate_qkv) + { + // Split the packed fp32 source the same way store_q_and_contiguous_kv_cache splits it + // for the non-sage layouts, including Q being right-aligned in the kv sequence when + // s_q < s_kv, so that both paths see identical values. + size_t const packed_hs = 2 * d + dv, packed_ts = h * packed_hs; + size_t const h_q_per_kv = h / h_kv; + q_f.assign(total_q * h * d, 0.f); + k_f.assign(total_kv * h_kv * d, 0.f); + v_f.assign(total_kv * h_kv * dv, 0.f); + for (size_t bi = 0; bi < b; bi++) + { + int const q_len = cu_q_seqlens[bi + 1] - cu_q_seqlens[bi]; + int const kv_len = cu_seqlens[bi + 1] - cu_seqlens[bi]; + for (int si = 0; si < q_len; si++) + { + size_t const src_t = cu_seqlens[bi] + kv_len - q_len + si; + size_t const dst_t = cu_q_seqlens[bi] + si; + for (size_t hi = 0; hi < h; hi++) + { + for (size_t di = 0; di < d; di++) + { + q_f[dst_t * h * d + hi * d + di] = qkv_packed_h[src_t * packed_ts + hi * packed_hs + di]; + } + } + } + for (int si = 0; si < kv_len; si++) + { + size_t const ti = cu_seqlens[bi] + si; + for (size_t hi = 0; hi < h_kv; hi++) + { + size_t const src = ti * packed_ts + hi * h_q_per_kv * packed_hs; + for (size_t di = 0; di < d; di++) + { + k_f[ti * h_kv * d + hi * d + di] = qkv_packed_h[src + d + di]; + } + for (size_t di = 0; di < dv; di++) + { + v_f[ti * h_kv * dv + hi * dv + di] = qkv_packed_h[src + 2 * d + di]; + } + } + } + } + q_b.assign(q_f.size(), 0); + k_b.assign(k_f.size(), 0); + v_b.assign(v_f.size(), 0); + sage_in.q = {q_f.data(), q_b.data(), h * d, d}; + sage_in.k = {k_f.data(), k_b.data(), h_kv * d, d}; + sage_in.v = {v_f.data(), v_b.data(), h_kv * dv, dv}; + } + else { - // simple test, all scales are the same - constexpr float const_scale = 0.618f; - fmha::e4m3_t* quant_qkv_h = (fmha::e4m3_t*) malloc(qkv_packed_size); - for (size_t i = 0; i < qkv_packed_size; i++) + // MQA/GQA and MHA pack QKV differently, so alias whichever buffer the kernel reads. + std::vector const& src = multi_query_attention ? mqa_qkv_packed_h : qkv_packed_h; + packed_b.assign(multi_query_attention ? mqa_qkv_packed_size : qkv_packed_size, 0); + if (multi_query_attention) { - quant_qkv_h[i] = fmha::e4m3_t(qkv_packed_h[i] / const_scale); + // [token][h + 2 * h_kv][d], all q heads first, then k heads, then v heads. + size_t const ts = (h + 2 * h_kv) * d; + sage_in.q = {src.data(), packed_b.data(), ts, d}; + sage_in.k = {src.data() + h * d, packed_b.data() + h * d, ts, d}; + sage_in.v = {src.data() + (h + h_kv) * d, packed_b.data() + (h + h_kv) * d, ts, d}; } - FMHA_CHECK_CUDA(cudaMemcpy(quant_qkv, quant_qkv_h, qkv_packed_size, cudaMemcpyHostToDevice)); - free(quant_qkv_h); - auto init_scales = [&](bert::Fused_multihead_attention_params_v2::SageAttention::Scales& x) + else { - std::vector scales(x.max_nblock * h * b, const_scale); - FMHA_CHECK_CUDA( - cudaMemcpy(x.scales, scales.data(), sizeof(float) * scales.size(), cudaMemcpyHostToDevice)); - }; - init_scales(params_v2.sage.q); - init_scales(params_v2.sage.k); - init_scales(params_v2.sage.v); + // [token][h][2 * d + dv], q at 0, k at d, v at 2 * d within each head. + size_t const hs = 2 * d + dv, ts = h * hs; + sage_in.q = {src.data(), packed_b.data(), ts, hs}; + sage_in.k = {src.data() + d, packed_b.data() + d, ts, hs}; + sage_in.v = {src.data() + 2 * d, packed_b.data() + 2 * d, ts, hs}; + } } -#else + + sage_quant::Output const sage_out = sage_quant::quantize(sage_in); + + params_v2.sage.q.block_size = sage_block_size_q; + params_v2.sage.q.max_nblock = sage_out.q_stride; + params_v2.sage.k.block_size = sage_block_size_k; + params_v2.sage.k.max_nblock = sage_out.k_stride; + params_v2.sage.v.block_size = sage_block_size_v; + + auto upload_bytes = [](void* dst, std::vector const& src) + { FMHA_CHECK_CUDA(cudaMemcpy(dst, src.data(), src.size(), cudaMemcpyHostToDevice)); }; + if (separate_qkv) { - // use external quant kernel - run_sage_quant(b, h, d, s, params_v2.qkv_ptr, - (char*) params_v2.qkv_ptr + get_size_in_bytes(h * d, data_type), - (char*) params_v2.qkv_ptr + get_size_in_bytes(2 * h * d, data_type, - params_v2.q_stride_in_bytes, - params_v2.k_stride_in_bytes, - params_v2.v_stride_in_bytes, - params_v2.cu_q_seqlens, params_v2.cu_kv_seqlens, sage_block_size_q, sage_block_size_k, - sage_block_size_v, quant_qkv, quant_qkv + h * d, quant_qkv + 2 * h * d, params_v2.sage.q.scales, - params_v2.sage.k.scales, params_v2.sage.v.scales); + upload_bytes(q_d, q_b); + upload_bytes(k_d, k_b); + upload_bytes(v_d, v_b); + params_v2.q_ptr = q_d; + params_v2.k_ptr = k_d; + params_v2.v_ptr = v_d; + params_v2.q_stride_in_bytes = get_size_in_bytes(h * d, DATA_TYPE_E4M3); + params_v2.k_stride_in_bytes = get_size_in_bytes(h_kv * d, DATA_TYPE_E4M3); + params_v2.v_stride_in_bytes = get_size_in_bytes(h_kv * dv, DATA_TYPE_E4M3); } -#endif - // no need to free old params_v2.qkv_ptr, it will be released in the end - params_v2.qkv_ptr = quant_qkv; - params_v2.q_stride_in_bytes = params_v2.k_stride_in_bytes = params_v2.v_stride_in_bytes - = get_size_in_bytes((h + 2 * h_kv) * d, DATA_TYPE_E4M3); + else + { + void* quant_qkv; + FMHA_CHECK_CUDA(cudaMalloc(&quant_qkv, packed_b.size())); + upload_bytes(quant_qkv, packed_b); + // no need to free old params_v2.qkv_ptr, it will be released in the end + params_v2.qkv_ptr = quant_qkv; + params_v2.q_stride_in_bytes = params_v2.k_stride_in_bytes = params_v2.v_stride_in_bytes + = get_size_in_bytes(sage_in.q.token_stride, DATA_TYPE_E4M3); + } + + auto upload = [](float** dst, std::vector const& src) + { + if (src.empty()) + { + return; + } + FMHA_CHECK_CUDA(cudaMalloc((void**) dst, sizeof(float) * src.size())); + FMHA_CHECK_CUDA(cudaMemcpy(*dst, src.data(), sizeof(float) * src.size(), cudaMemcpyHostToDevice)); + }; + upload(¶ms_v2.sage.q.scales, sage_out.scales_q); + upload(¶ms_v2.sage.k.scales, sage_out.scales_k); + upload(¶ms_v2.sage.v.scales, sage_out.scales_v); } #if defined(DEBUG_HAS_PRINT_BUFFER) diff --git a/cpp/kernels/fmha_v2/src/fused_multihead_attention_sage_utils.h b/cpp/kernels/fmha_v2/src/fused_multihead_attention_sage_utils.h new file mode 100644 index 000000000000..3d92765069df --- /dev/null +++ b/cpp/kernels/fmha_v2/src/fused_multihead_attention_sage_utils.h @@ -0,0 +1,376 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2011-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: NVIDIA TensorRT Source Code License Agreement + * + * NVIDIA CORPORATION, its affiliates and licensors retain all intellectual + * property and proprietary rights in and to this material, related + * documentation and any modifications thereto. Any use, reproduction, + * disclosure or distribution of this material and related documentation + * without an express license agreement from NVIDIA CORPORATION or + * its affiliates is strictly prohibited. + */ + +#pragma once + +#include +#include +#include +#include +#include + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Reference host-side quantizer for SageAttention, used by bin/fmha.exe to feed the SM90 kernels. +// It is a correctness reference, not a fast path: production callers quantize on the device. +// +// SageAttention is INT8 QK with an e4m3 PV. Q and K are quantized along the sequence axis and V +// along the channel axis, with the groups chosen so a thread never needs more than one scale per +// axis. See fmha/warpspec/compute.h for the matching kernel-side index arithmetic. + +namespace sage_quant +{ + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Geometry of the SM90 warp-specialized sage kernels, mirrored on the host so the scale layout +// built here is the one the kernel indexes into. +struct Geometry +{ + // The Q and KV tile sizes of the kernel. + static constexpr int step_q = 64; + static constexpr int step_kv = 256; + // Q is quantized per thread: a thread owns the 2 strided rows quad_row and quad_row + 8. + static constexpr int q_tokens_per_scale = 2; + + // Scales per q tile, laid out as [q_tile][warp * 8 + lane / 4]. + static constexpr int q_scales_per_tile = step_q / q_tokens_per_scale; + // K sub-fragments per thread. A thread owns 64 key columns; splitting them by ni-range gives + // this many groups, and a thread loads all of its scales with one 128-bit load. + int k_chunks; + // Scales per kv tile, laid out as [kv_tile][quad][chunk]. + int k_scales_per_tile; + // Keys covered by one K scale, i.e. the K block size the kernel was compiled for. + int k_tokens_per_scale; + + // The scale buffers are laid out as (H, scales for the whole batch), so a sequence's scales + // start at a base derived from its cu_seqlens entry. Each sequence reserves one spare tile + // because the kernel indexes a partial final tile in full; K reserves 4 more so its base can be + // rounded up to a multiple of 4 and keep the kernel's float4 scale load aligned. + static constexpr int q_slots_per_seq = q_scales_per_tile; + int k_slots_per_seq; + + explicit Geometry(int block_size_k) + : k_chunks(step_kv / (4 * std::max(1, block_size_k))) + , k_scales_per_tile(4 * k_chunks) + , k_tokens_per_scale(step_kv / k_scales_per_tile) + , k_slots_per_seq(k_scales_per_tile + 4) + { + } + + // Where sequence bi's scales start, given its first token. Must match warpspec/compute.h. + int q_base(int cu_seqlen, size_t bi) const + { + return cu_seqlen / q_tokens_per_scale + (int) bi * q_slots_per_seq; + } + + int k_base(int cu_seqlen, size_t bi) const + { + return ((cu_seqlen / k_tokens_per_scale + 3) & ~3) + (int) bi * k_slots_per_seq; + } + + // Scales per head, i.e. what the kernel reads as max_nblock. Depends only on the batch size and + // the total token count, both of which the caller knows before it sees any sequence length. + int q_stride(int total_tokens, size_t b) const + { + return q_base(total_tokens, b); + } + + int k_stride(int total_tokens, size_t b) const + { + return k_base(total_tokens, b); + } +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// One of Q, K or V: where to read fp32 from, where to write the quantized byte to, and how the +// tensor is addressed. Element (token ti, head hi, channel di) sits at +// ti * token_stride + hi * head_stride + di in both buffers, so one stride pair serves both. +// +// Strides rather than a layout enum because that covers every layout the harness feeds in with no +// branching in the inner loops: +// MHA packed [token][h][3][d] token_stride = h * 3 * d, head_stride = 3 * d +// MQA packed [token][h + 2 * h_kv][d] token_stride = (h + 2 * h_kv) * d, head_stride = d +// separate [token][h][d] token_stride = h * d, head_stride = d +// For a packed layout the three Tensors alias one buffer at different offsets; for separate Q/K/V +// they are three distinct buffers, and the q and kv sequence lengths may differ. +struct Tensor +{ + float const* src; + uint8_t* dst; + size_t token_stride, head_stride; +}; + +// What to quantize. +struct Input +{ + Tensor q, k, v; + + // The dimensions. + size_t b, h, h_kv, d, dv; + // Prefix sums of the actual sequence lengths, b + 1 entries each. They are the same array for + // a packed layout; separate Q/K/V allows them to differ. + int const *cu_q_seqlens, *cu_kv_seqlens; + + // The block sizes requested on the command line, counted along the axis each tensor is + // quantized over: 2 rows for Q, 16 keys for K, 1 channel for V (0 leaves V unquantized). + int block_size_q, block_size_k, block_size_v; +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// The scales. The quantized values themselves go straight into the caller's Tensor::dst buffers; +// Q and K are INT8 and V is e4m3, both one byte per element, so every stride is the input's. +struct Output +{ + // Layout (h, q_stride), a sequence's tiles swizzled as [q_tile][warp * 8 + lane / 4]. + std::vector scales_q; + // Layout (h_kv, k_stride), a sequence's tiles swizzled as [kv_tile][quad][chunk]. + std::vector scales_k; + // Layout (h_kv, dv), one scale per channel. Empty when V is not quantized. The amax reduces + // over every token of every sequence, so there is no batch dimension. + std::vector scales_v; + // Scales per head, what the kernel reads as max_nblock. + int q_stride, k_stride; +}; + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +namespace detail +{ + +// (token, head, channel) -> flat index into a Tensor's src and dst buffers. +inline size_t at(Tensor const& t, size_t ti, size_t hi, size_t di) +{ + return ti * t.token_stride + hi * t.head_stride + di; +} + +// A group with no valid tokens has amax 0. Give it a representative scale rather than an epsilon: +// the kernel masks with INT_MIN and relies on that surviving the multiply by this scale (see +// MASKED_ACC in fmha/warpspec/epilogue.h). +inline float make_scale(float amax, float div) +{ + return amax > 0.f ? amax / div : 1.f; +} + +// With random inputs the amax of every group comes out nearly identical, which would make the +// scales effectively constant and hide any mis-indexed scale. Spread them instead. +// +// All three axes are mixed so the factor depends on every one of them. A plain `idx % 4` over +// `quad * k_chunks + g` would give every quad the same pattern (4 * quad is 0 mod 4) and a +// wrong-quad read would be invisible; likewise `pair % 4` over `warp * 8 + lane / 4` hides a +// wrong-warp read. The tile index has to be in the mix as well: without it every tile gets an +// identical pattern, so reading a neighbouring tile's scales -- which is exactly what a wrong +// per-sequence base does -- changes nothing, and with iid inputs the amax of every tile is nearly +// the same, so the error stays at noise. +// +// Q and K are fixed-point INT8, so inflating a scale genuinely throws away bits (8x would cost +// three of them); keep the four factors distinct but gentle. +inline float qk_spread(size_t minor, size_t major, size_t tile) +{ + return 1.f + 0.25f * ((minor + 3 * major + tile) % 4); +} + +// V is e4m3, where a power-of-two factor is free: fp8 rounding is scale-invariant, so it leaves the +// error of a correct implementation unchanged and any extra error is a genuine indexing bug. +inline float v_spread(size_t idx) +{ + return float(1u << (idx % 4)); +} + +// Symmetric INT8, saturating. +inline void store_int8(Tensor const& t, size_t idx, float value, float scale) +{ + int const q = (int) lrintf(value / scale); + reinterpret_cast(t.dst)[idx] = (int8_t) std::min(127, std::max(-127, q)); +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Per-thread Q scales. The BMM1 fragment gives the thread at (warp w, lane l) the two query rows +// quad_row = w * 16 + l / 4 and quad_row + 8, so a group is 2 strided rows and there are +// step_q / 2 of them per q tile. +inline void quant_q_per_thread(Input const& in, Geometry const& geo, Output& out) +{ + out.q_stride = geo.q_stride(in.cu_q_seqlens[in.b], in.b); + out.scales_q.assign((size_t) out.q_stride * in.h, 1.f); + + for (size_t bi = 0; bi < in.b; bi++) + { + int const seq_beg = in.cu_q_seqlens[bi], seq_end = in.cu_q_seqlens[bi + 1]; + int const num_tiles = (seq_end - seq_beg + geo.step_q - 1) / geo.step_q; + int const base = geo.q_base(seq_beg, bi); + for (size_t hi = 0; hi < in.h; hi++) + { + for (int tile = 0; tile < num_tiles; tile++) + { + for (int pair = 0; pair < geo.q_scales_per_tile; pair++) + { + // pair == w * 8 + l / 4 -> quad_row = (pair / 8) * 16 + pair % 8 + int const quad_row = (pair / 8) * 16 + (pair % 8); + int const rows[2] = {tile * geo.step_q + quad_row, tile * geo.step_q + quad_row + 8}; + + float amax = 0.f; + for (int r = 0; r < 2; r++) + { + if (seq_beg + rows[r] >= seq_end) + { + continue; + } + size_t const ti = seq_beg + rows[r]; + for (size_t di = 0; di < in.d; di++) + { + amax = std::max(amax, fabsf(in.q.src[at(in.q, ti, hi, di)])); + } + } + + float const scale = make_scale(amax, 127.f) * qk_spread(pair % 8, pair / 8, tile); + out.scales_q[hi * out.q_stride + base + tile * geo.q_scales_per_tile + pair] = scale; + + for (int r = 0; r < 2; r++) + { + if (seq_beg + rows[r] >= seq_end) + { + continue; + } + size_t const ti = seq_beg + rows[r]; + for (size_t di = 0; di < in.d; di++) + { + size_t const idx = at(in.q, ti, hi, di); + store_int8(in.q, idx, in.q.src[idx], scale); + } + } + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Per-thread K scales. A thread with lane % 4 == quad owns the key columns {8 * ni + 2 * quad + e}, +// ni in [0, 32), e in {0, 1}. Splitting that fragment into k_chunks sub-fragments by ni-range gives +// k_tokens_per_scale keys per scale. +inline void quant_k_per_thread(Input const& in, Geometry const& geo, Output& out) +{ + int const ni_per_chunk = geo.step_kv / 8 / geo.k_chunks; + out.k_stride = geo.k_stride(in.cu_kv_seqlens[in.b], in.b); + out.scales_k.assign((size_t) out.k_stride * in.h_kv, 1.f); + + for (size_t bi = 0; bi < in.b; bi++) + { + int const seq_beg = in.cu_kv_seqlens[bi], seq_end = in.cu_kv_seqlens[bi + 1]; + int const num_tiles = (seq_end - seq_beg + geo.step_kv - 1) / geo.step_kv; + int const base = geo.k_base(seq_beg, bi); + for (size_t hi = 0; hi < in.h_kv; hi++) + { + for (int tile = 0; tile < num_tiles; tile++) + { + for (int quad = 0; quad < 4; quad++) + { + for (int g = 0; g < geo.k_chunks; g++) + { + std::vector toks; + for (int ni = g * ni_per_chunk; ni < (g + 1) * ni_per_chunk; ni++) + { + for (int e = 0; e < 2; e++) + { + int const local = tile * geo.step_kv + 8 * ni + 2 * quad + e; + if (seq_beg + local < seq_end) + { + toks.push_back(seq_beg + local); + } + } + } + + float amax = 0.f; + for (int ti : toks) + { + for (size_t di = 0; di < in.d; di++) + { + amax = std::max(amax, fabsf(in.k.src[at(in.k, ti, hi, di)])); + } + } + + float const scale = make_scale(amax, 127.f) * qk_spread(g, quad, tile); + out.scales_k[hi * out.k_stride + base + tile * geo.k_scales_per_tile + quad * geo.k_chunks + g] + = scale; + + for (int ti : toks) + { + for (size_t di = 0; di < in.d; di++) + { + size_t const idx = at(in.k, ti, hi, di); + store_int8(in.k, idx, in.k.src[idx], scale); + } + } + } + } + } + } + } +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +// Per-(head, channel) V scales: the amax reduces over the sequence, keeping the head dimension. +// There is a single set for the whole batch, so the reduction spans every token of every sequence. +inline void quant_v_per_channel(Input const& in, Output& out) +{ + int const total_tokens = in.cu_kv_seqlens[in.b]; + out.scales_v.assign(in.h_kv * in.dv, 1.f); + + for (size_t hi = 0; hi < in.h_kv; hi++) + { + for (size_t di = 0; di < in.dv; di++) + { + float amax = 0.f; + for (int ti = 0; ti < total_tokens; ti++) + { + amax = std::max(amax, fabsf(in.v.src[at(in.v, ti, hi, di)])); + } + + float const scale = make_scale(amax, 448.f) * v_spread(di); + out.scales_v[hi * in.dv + di] = scale; + + for (int ti = 0; ti < total_tokens; ti++) + { + size_t const idx = at(in.v, ti, hi, di); + reinterpret_cast(in.v.dst)[idx] = fmha::e4m3_t(in.v.src[idx] / scale); + } + } + } +} + +} // namespace detail + +inline Output quantize(Input const& in) +{ + Geometry const geo(in.block_size_k); + + Output out; + out.q_stride = out.k_stride = 0; + + detail::quant_q_per_thread(in, geo, out); + detail::quant_k_per_thread(in, geo, out); + if (in.block_size_v == 1) + { + detail::quant_v_per_channel(in, out); + } + + return out; +} + +//////////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace sage_quant diff --git a/cpp/tensorrt_llm/common/attentionOp.cpp b/cpp/tensorrt_llm/common/attentionOp.cpp index 56f64b161c69..2ce34dfe3761 100644 --- a/cpp/tensorrt_llm/common/attentionOp.cpp +++ b/cpp/tensorrt_llm/common/attentionOp.cpp @@ -846,8 +846,7 @@ size_t AttentionOp::getWorkspaceSizeForContext(tensorrt_llm::DataType type, int3 = mNumAttnHeads * dim_k_per_head; // Assuming effective num_kv_heads = head_num for layout int const total_v_dim_all_heads = mNumAttnHeads * dim_v_per_head; // Assuming effective num_kv_heads = head_num for layout - bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mIsMLAEnabled && mFmhaDispatcher->isSeparateQAndKvInput() - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); + bool const useSageAttnSeparateQkv = mEnableContextFMHA && useSageAttn() && mFmhaDispatcher->isSeparateQAndKvInput(); // Packed fp8 qkv buffer size for normal fp8 context FMHA size_t fp8_qkv_buffer_size = mFP8ContextFMHA && mEnableContextFMHA && !mFmhaDispatcher->isSeparateQAndKvInput() @@ -892,15 +891,17 @@ size_t AttentionOp::getWorkspaceSizeForContext(tensorrt_llm::DataType type, int3 fp8_v_buf_size = total_kv_len * static_cast(local_hidden_units_kv); } - int32_t const q_max_n_blk - = mSageAttnNumEltsPerBlkQ > 0 ? tc::divUp(max_num_tokens, mSageAttnNumEltsPerBlkQ) + batch_size - 1 : 0; - int32_t const k_max_n_blk - = mSageAttnNumEltsPerBlkK > 0 ? tc::divUp(total_kv_len, mSageAttnNumEltsPerBlkK) + batch_size - 1 : 0; + bool const hopperSage = useHopperSageAttn(); + int32_t const q_max_n_blk = tc::getSageScaleHeadStride( + tc::getSageQPartition(hopperSage), mSageAttnNumEltsPerBlkQ, max_num_tokens, batch_size); + int32_t const k_max_n_blk = tc::getSageScaleHeadStride( + tc::getSageKPartition(hopperSage), mSageAttnNumEltsPerBlkK, total_kv_len, batch_size); + int32_t const v_max_n_blk + = mSageAttnNumEltsPerBlkV > 0 ? tc::divUp(local_hidden_units_kv, mSageAttnNumEltsPerBlkV) : 0; size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); - size_t const sage_v_sfs_buffer_size = mSageAttnNumEltsPerBlkV > 0 - ? sizeof(float) * tc::divUp(local_hidden_units_kv, std::max(1, mSageAttnNumEltsPerBlkV)) - : 0; + size_t const sage_v_sfs_buffer_size = sizeof(float) * static_cast(v_max_n_blk); + size_t const sage_k_mean_buffer_size = mSageAttnSmoothK ? sizeof(float) * local_hidden_units_kv : 0; size_t const padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * padded_num_tokens; size_t const encoder_padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * padded_kv_tokens; @@ -943,6 +944,7 @@ size_t AttentionOp::getWorkspaceSizeForContext(tensorrt_llm::DataType type, int3 workspaceSizes.sageQScale = sage_q_sfs_buffer_size; workspaceSizes.sageKScale = sage_k_sfs_buffer_size; workspaceSizes.sageVScale = sage_v_sfs_buffer_size; + workspaceSizes.sageKMean = sage_k_mean_buffer_size; workspaceSizes.cpWorkspace = cpWorkspaceSize; workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; context_workspace_size = AttentionWorkspaceManager::buildContextLayout(workspaceSizes).totalSize; @@ -1584,8 +1586,14 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea size_t fp8_q_buf_size = 0; size_t fp8_k_buf_size = 0; size_t fp8_v_buf_size = 0; - bool const useSageAttnSeparateQkv = mEnableContextFMHA && !mIsMLAEnabled && mFmhaDispatcher->isSeparateQAndKvInput() - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); + // SageAttention has no unfused fallback. Falling back would silently run unquantized + // attention, silently giving wrong results. Need to report it upfront. + TLLM_CHECK_WITH_INFO(!sageAttnRequested() || useSageAttn(), + "SageAttention requires the FP8 context FMHA path and does not apply to MLA."); + TLLM_CHECK_WITH_INFO(!useSageAttn() || mEnableContextFMHA, + "Sage Attention requires contextFMHA with no unfused fallback, but the supplied configuration is unsupported."); + + bool const useSageAttnSeparateQkv = mEnableContextFMHA && useSageAttn() && mFmhaDispatcher->isSeparateQAndKvInput(); if (mEnableContextFMHA && mFP8ContextMLA && mFmhaDispatcher->isSeparateQAndKvInput()) { fp8_q_buf_size = params.num_tokens * static_cast(total_q_dim_all_heads); @@ -1610,17 +1618,17 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea fp8_v_buf_size = params.total_kv_len * static_cast(local_hidden_units_kv); } - int32_t const q_max_n_blk = mSageAttnNumEltsPerBlkQ > 0 - ? tc::divUp(params.num_tokens, mSageAttnNumEltsPerBlkQ) + params.batch_size - 1 - : 0; - int32_t const k_max_n_blk = mSageAttnNumEltsPerBlkK > 0 - ? tc::divUp(params.total_kv_len, mSageAttnNumEltsPerBlkK) + params.batch_size - 1 - : 0; + bool const hopperSage = useHopperSageAttn(); + int32_t const q_max_n_blk = tc::getSageScaleHeadStride( + tc::getSageQPartition(hopperSage), mSageAttnNumEltsPerBlkQ, params.num_tokens, params.batch_size); + int32_t const k_max_n_blk = tc::getSageScaleHeadStride( + tc::getSageKPartition(hopperSage), mSageAttnNumEltsPerBlkK, params.total_kv_len, params.batch_size); int32_t const v_max_n_blk = mSageAttnNumEltsPerBlkV > 0 ? tc::divUp(local_hidden_units_kv, mSageAttnNumEltsPerBlkV) : 0; size_t const sage_q_sfs_buffer_size = sizeof(float) * mNumAttnHeads * static_cast(q_max_n_blk); size_t const sage_k_sfs_buffer_size = sizeof(float) * mNumAttnKVHeads * static_cast(k_max_n_blk); - size_t const sage_v_sfs_buffer_size = sizeof(float) * v_max_n_blk; + size_t const sage_v_sfs_buffer_size = sizeof(float) * static_cast(v_max_n_blk); + size_t const sage_k_mean_buffer_size = mSageAttnSmoothK ? sizeof(float) * local_hidden_units_kv : 0; size_t const padding_offset_size = mEnableContextFMHA ? 0 : sizeof(int) * params.batch_size * params.input_seq_length; @@ -1666,6 +1674,7 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea workspaceSizes.sageQScale = sage_q_sfs_buffer_size; workspaceSizes.sageKScale = sage_k_sfs_buffer_size; workspaceSizes.sageVScale = sage_v_sfs_buffer_size; + workspaceSizes.sageKMean = sage_k_mean_buffer_size; workspaceSizes.cpWorkspace = cpWorkspaceSize; workspaceSizes.fmhaMultiCtasKvScratch = fmha_multi_ctas_kv_scratch_size; auto const workspaceLayout = AttentionWorkspaceManager::buildContextLayout(workspaceSizes); @@ -1941,16 +1950,26 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea } else if (useSageAttnSeparateQkv) { - TLLM_CHECK_WITH_INFO(mFP8ContextFMHA, "SageAttention kernel runs under mFP8ContextFMHA option."); - TLLM_CHECK_WITH_INFO(mFmhaDispatcher->isSupported(), "SageAttention has no unfused fallback implemented."); TLLM_CHECK_WITH_INFO(mMaskType == AttentionMaskType::PADDING, "SageAttention only supports dense (padding) mask, got mask type %d.", static_cast(mMaskType)); TLLM_CHECK_WITH_INFO( mSageAttnNumEltsPerBlkQ > 0 && mSageAttnNumEltsPerBlkK > 0 && mSageAttnNumEltsPerBlkV == 1, "SageQuant requires positive block sizes for Q and K while the block size for V must be 1."); + // A zero stride means the quantizer has no kernel for this block size, which would + // otherwise surface as an empty scale buffer deep inside invokeSageQuant(). + TLLM_CHECK_WITH_INFO(q_max_n_blk > 0 && k_max_n_blk > 0, + "No SageQuant kernel on sm_%d for block sizes (q, k, v) = (%d, %d, %d) with qk_int8=%s. SM90 " + "requires (2, 16, 1) together with INT8 Q/K; SM100 requires q/k block sizes of 1, 4 or 16.", + mSM, mSageAttnNumEltsPerBlkQ, mSageAttnNumEltsPerBlkK, mSageAttnNumEltsPerBlkV, + mSageAttnQkInt8 ? "true" : "false"); TLLM_CHECK_WITH_INFO(!params.kv_scale_quant_orig, "SageAttention disregards the configured params.kv_scale_quant_orig, invalidating the result."); + // Reduction buffers must be initialized prior to invokeSageQuant(). check_cuda_error(cudaMemsetAsync(workspaceViews.sageVScale, 0, sage_v_sfs_buffer_size, stream)); + if (mSageAttnSmoothK) + { + check_cuda_error(cudaMemsetAsync(workspaceViews.sageKMean, 0, sage_k_mean_buffer_size, stream)); + } // Common params for sageQuant tc::SageQuantParams sageQuantParams{}; @@ -1960,16 +1979,20 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea sageQuantParams.vStage = 0; sageQuantParams.sumSeqLensV = params.total_kv_len; sageQuantParams.numHeadsV = mNumAttnKVHeads; + sageQuantParams.kSmooth = mSageAttnSmoothK; sageQuantParams.ptrV = params.v_ptr; sageQuantParams.ptrVQuant = workspaceViews.fp8VBuf; sageQuantParams.ptrVScale = workspaceViews.sageVScale; + sageQuantParams.ptrKForMean = params.k_ptr; + sageQuantParams.ptrKMean = workspaceViews.sageKMean; sageQuantParams.smCount = mMultiProcessorCount; sageQuantParams.stream = stream; - // Quantize into Fp8Q, SfsQ, SfsV + // Quantize into Q, SfsQ, SfsV sageQuantParams.sumSeqLensQk = params.num_tokens; sageQuantParams.batchSize = params.batch_size; sageQuantParams.numHeads = mNumAttnHeads; + sageQuantParams.partition = tc::getSageQPartition(hopperSage); sageQuantParams.tokenBlockSize = mSageAttnNumEltsPerBlkQ; sageQuantParams.ptrCuSeqLensQk = contextCuQSeqlens; sageQuantParams.ptrQk = attention_input; @@ -1978,10 +2001,11 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea sageQuantParams.vStage = 1; tc::invokeSageQuant(sageQuantParams); - // Quantize into Fp8K, SfsK, Fp8V + // Quantize into K, SfsK, V sageQuantParams.sumSeqLensQk = params.total_kv_len; sageQuantParams.batchSize = params.batch_size; sageQuantParams.numHeads = mNumAttnKVHeads; + sageQuantParams.partition = tc::getSageKPartition(hopperSage); sageQuantParams.tokenBlockSize = mSageAttnNumEltsPerBlkK; sageQuantParams.ptrCuSeqLensQk = contextCuKvSeqlens; sageQuantParams.ptrQk = params.k_ptr; @@ -2080,10 +2104,6 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea fmhaParams.qPtr = reinterpret_cast(workspaceViews.fp8QBuf); fmhaParams.kPtr = reinterpret_cast(workspaceViews.fp8KBuf); fmhaParams.vPtr = reinterpret_cast(workspaceViews.fp8VBuf); - // Set sage attention scaling factor pointers. - fmhaParams.qScalePtr = workspaceViews.sageQScale; - fmhaParams.kScalePtr = workspaceViews.sageKScale; - fmhaParams.vScalePtr = workspaceViews.sageVScale; } else { @@ -2091,6 +2111,17 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea : reinterpret_cast(attention_input); fmhaParams.qPtr = reinterpret_cast(workspaceViews.qBuf); } + + if (useSageAttnSeparateQkv) + { + // SageAttention scaling factors. + fmhaParams.qScalePtr = workspaceViews.sageQScale; + fmhaParams.kScalePtr = workspaceViews.sageKScale; + fmhaParams.vScalePtr = workspaceViews.sageVScale; + fmhaParams.qMaxNBlock = q_max_n_blk; + fmhaParams.kMaxNBlock = k_max_n_blk; + fmhaParams.vMaxNBlock = v_max_n_blk; + } // TODO: add contiguous kv buffer (cross-attention). fmhaParams.kvPtr = nullptr; if (isCrossAttention() && !useKVCache()) @@ -2987,8 +3018,7 @@ int AttentionOp::initialize() noexcept // Construct the fmha runner. MHARunnerFixedParams fmhaParams{}; - bool const useSageAttn = mFP8ContextFMHA && !mIsMLAEnabled - && (mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0); + bool const useSageAttn = this->useSageAttn(); // Pre-checked during constructing. Data_type data_type, data_type_kv; diff --git a/cpp/tensorrt_llm/common/attentionOp.h b/cpp/tensorrt_llm/common/attentionOp.h index bbc066b24c64..1a7177e8bf3f 100644 --- a/cpp/tensorrt_llm/common/attentionOp.h +++ b/cpp/tensorrt_llm/common/attentionOp.h @@ -428,6 +428,22 @@ class AttentionOp return useSparseMLA() || (mUseSparseAttention && mUseTllmGen && mUseTllmGenSparseAttention); } + [[nodiscard]] bool sageAttnRequested() const + { + return mSageAttnNumEltsPerBlkQ > 0 || mSageAttnNumEltsPerBlkK > 0 || mSageAttnNumEltsPerBlkV > 0; + } + + [[nodiscard]] bool useSageAttn() const + { + return sageAttnRequested() && mFP8ContextFMHA && !mIsMLAEnabled && !useKVCache(); + } + + [[nodiscard]] bool useHopperSageAttn() const + { + return useSageAttn() && mSM == kernels::kSM_90 && mSageAttnQkInt8 && mSageAttnNumEltsPerBlkQ == 2 + && mSageAttnNumEltsPerBlkK == 16 && mSageAttnNumEltsPerBlkV == 1; + } + [[nodiscard]] int smVersion() const { return mSM; @@ -561,12 +577,12 @@ class AttentionOp float mSkipCorrectionThreshold = 0; // Use spcompress (context phase, SM107 only). bool mUsesSpcompress = false; - // Optional SageAttention block sizes. - // Currently, these are only consumed by the TllmGen backend path. + // Optional SageAttention block sizes, in tokens per scale for q/k and channels per scale for v. int mSageAttnNumEltsPerBlkQ = 0; int mSageAttnNumEltsPerBlkK = 0; int mSageAttnNumEltsPerBlkV = 0; bool mSageAttnQkInt8 = false; + bool mSageAttnSmoothK = false; #ifdef SKIP_SOFTMAX_STAT uint32_t* mSkipSoftmaxTotalBlocks; uint32_t* mSkipSoftmaxSkippedBlocks; @@ -592,7 +608,7 @@ class AttentionOp mSkipAttn, mFuseFp4Quant, mFusesDsv4InvRopeFp8Quant, mNbMultiBlockSemaphores, mAttentionChunkSize.value_or(-1), mSkipSoftmaxThresholdScaleFactorPrefill, mSkipSoftmaxThresholdScaleFactorDecode, mSkipCorrectionThreshold, mUsesSpcompress, mSageAttnNumEltsPerBlkQ, - mSageAttnNumEltsPerBlkK, mSageAttnNumEltsPerBlkV, mSageAttnQkInt8); + mSageAttnNumEltsPerBlkK, mSageAttnNumEltsPerBlkV, mSageAttnQkInt8, mSageAttnSmoothK); }; private: diff --git a/cpp/tensorrt_llm/common/attentionWorkspace.h b/cpp/tensorrt_llm/common/attentionWorkspace.h index 3ba53da3ff34..774235881aa0 100644 --- a/cpp/tensorrt_llm/common/attentionWorkspace.h +++ b/cpp/tensorrt_llm/common/attentionWorkspace.h @@ -65,6 +65,7 @@ struct AttentionContextWorkspaceSizes size_t sageQScale{}; size_t sageKScale{}; size_t sageVScale{}; + size_t sageKMean{}; size_t cpWorkspace{}; size_t fmhaMultiCtasKvScratch{}; }; @@ -96,6 +97,7 @@ struct AttentionContextWorkspaceLayout WorkspaceSlice sageQScale{}; WorkspaceSlice sageKScale{}; WorkspaceSlice sageVScale{}; + WorkspaceSlice sageKMean{}; WorkspaceSlice cpWorkspace{}; WorkspaceSlice fmhaMultiCtasKvScratch{}; size_t totalSize{}; @@ -129,6 +131,7 @@ struct AttentionContextWorkspaceViews float* sageQScale{}; float* sageKScale{}; float* sageVScale{}; + float* sageKMean{}; T* gatherInBuffer{}; T* gatherOutBuffer{}; int* cuCpPartialSeqlens{}; @@ -254,6 +257,7 @@ class AttentionWorkspaceManager layout.sageQScale = nextSlice(offset, sizes.sageQScale, alignment); layout.sageKScale = nextSlice(offset, sizes.sageKScale, alignment); layout.sageVScale = nextSlice(offset, sizes.sageVScale, alignment); + layout.sageKMean = nextSlice(offset, sizes.sageKMean, alignment); layout.cpWorkspace = nextSlice(offset, sizes.cpWorkspace, alignment); layout.fmhaMultiCtasKvScratch = nextSlice(offset, sizes.fmhaMultiCtasKvScratch, alignment); layout.totalSize = offset; @@ -291,6 +295,7 @@ class AttentionWorkspaceManager views.sageQScale = ptr(workspace, layout.sageQScale); views.sageKScale = ptr(workspace, layout.sageKScale); views.sageVScale = ptr(workspace, layout.sageVScale); + views.sageKMean = ptr(workspace, layout.sageKMean); views.gatherInBuffer = ptr(workspace, layout.cpWorkspace); if (views.gatherInBuffer != nullptr) diff --git a/cpp/tensorrt_llm/common/sagePartition.h b/cpp/tensorrt_llm/common/sagePartition.h new file mode 100644 index 000000000000..5effb2dfc0af --- /dev/null +++ b/cpp/tensorrt_llm/common/sagePartition.h @@ -0,0 +1,145 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 1993-2026 NVIDIA CORPORATION & + * AFFILIATES. All rights reserved. SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#pragma once + +#include + +namespace tensorrt_llm::common +{ + +// How SageAttention groups tokens into quantization scales, as a CuTe layout. +// +// A partition is a bijective layout ((intra...), (group...)) -> token offset inside one *tile*: +// mode 0 enumerates the tokens that share one scale, +// mode 1 enumerates the scale groups of a tile, +// size(layout) is the tile size in tokens. +// +// This is what differs between GPUs. SM100 wants each scale to cover a run of contiguous tokens, +// so its tile is one group. SM90 scales per thread, and the BMM1 fragment hands each thread a +// *strided* set of rows/columns, so its groups interleave and the group index inside a tile is +// swizzled. Expressing both as a layout lets one quantizer serve both: it just evaluates +// partition(i, g). +// +// ScaleAlign is the alignment, in scales, that the consumer needs a sequence's scale base to +// have. The SM90 K path loads its four scales as one float4, so it needs 4; everything else +// needs 1. +// +// HeadStrideCountsBlocks selects how the consumer derives the head stride of a scale buffer. +// The two conventions agree on where a sequence's scales start and differ only in the total, +// so a partition has to declare the one its consumer was built with. + +// ------------------------------------------------------------------------------------------- +// SM100 / contiguous: a tile is one group of BlockSize consecutive tokens. +template +struct ContiguousPartition +{ + using Layout = cute::Layout>, cute::Shape>, + cute::Stride, cute::Stride>>>; + static constexpr int ScaleAlign = 1; + // Head stride: started blocks over the whole batch, plus one spare per extra sequence. + static constexpr bool HeadStrideCountsBlocks = true; +}; + +// ------------------------------------------------------------------------------------------- +// SM90 Q: STEP_Q = 64 rows per tile, 2 rows per scale, 32 scales per tile. +// +// The BMM1 fragment gives the thread at (warp w, lane l) the query rows quad_row and quad_row + 8, +// with quad_row = 16 * w + l / 4. Numbering the groups by pair = w * 8 + l / 4 (which is how +// fmha/warpspec/compute.h indexes them), group pair = p8 + 8 * w owns rows 16 * w + p8 + 8 * r: +// intra : (r) stride 8 +// group : (p8, w) strides (1, 16), linear index p8 + 8 * w == pair +struct HopperQPartition +{ + using Layout = cute::Layout, cute::Shape>, + cute::Stride, cute::Stride>>; + static constexpr int ScaleAlign = 1; + // Head stride: the sequence base function evaluated at (sumSeqLens, batchSize). + static constexpr bool HeadStrideCountsBlocks = false; +}; + +// ------------------------------------------------------------------------------------------- +// SM90 K: STEP_KV = 256 keys per tile, 16 keys per scale, 16 scales per tile. +// +// A thread with lane % 4 == quad owns the key columns 8 * ni + 2 * quad + e, ni in [0, 32), +// e in {0, 1}. Splitting that fragment into 4 sub-fragments by ni-range gives 16 keys per scale; +// writing ni = 8 * g + n, a key is 64 * g + 8 * n + 2 * quad + e: +// intra : (e, n) strides (1, 8) +// group : (g, quad) strides (64, 2), linear index g + 4 * quad == quad * K_CHUNKS + g +struct HopperKPartition +{ + using Layout = cute::Layout, cute::Shape>, + cute::Stride, cute::Stride>>; + // The consumer loads a thread's four chunk scales with one 128-bit access. + static constexpr int ScaleAlign = 4; + // Head stride: the sequence base function evaluated at (sumSeqLens, batchSize). + static constexpr bool HeadStrideCountsBlocks = false; +}; + +// ------------------------------------------------------------------------------------------- +// Derived geometry, shared by the kernel and the host. +template +struct PartitionTraits +{ + using Layout = typename Partition::Layout; + + // Tokens sharing one scale. + static constexpr int TokensPerScale = cute::size<0>(Layout{}); + // Scales in one tile. + static constexpr int ScalesPerTile = cute::size<1>(Layout{}); + // Tokens in one tile. + static constexpr int TileTokens = cute::size(Layout{}); + static constexpr int ScaleAlign = Partition::ScaleAlign; + + static_assert(TileTokens == TokensPerScale * ScalesPerTile, "partition must be a bijection"); + static_assert(cute::cosize(Layout{}) == TileTokens, "partition must tile densely"); + static_assert(ScaleAlign > 0 && (ScaleAlign & (ScaleAlign - 1)) == 0, "ScaleAlign must be 2^k"); + + // A sequence reserves one spare tile's worth of scales, because a partial final tile still + // indexes all ScalesPerTile of them, plus ScaleAlign - 1 to absorb the base round-up. Rounded + // up to a multiple of ScaleAlign so that every base stays aligned. + static constexpr int SlotsPerSeq = ((ScalesPerTile + ScaleAlign - 1) + ScaleAlign - 1) / ScaleAlign * ScaleAlign; + + // Where sequence seqIdx's scales start within a head. + CUTE_HOST_DEVICE static constexpr int scaleBase(int cuSeqLen, int seqIdx) + { + return ((cuSeqLen / TokensPerScale + ScaleAlign - 1) / ScaleAlign) * ScaleAlign + seqIdx * SlotsPerSeq; + } + + // Scales per head, i.e. the head stride of the scale buffer. Depends only on the batch size and + // the total token count, never on the longest sequence. The consumer kernels are not passed + // this value; they recompute it, so it has to match the convention they were built with. + CUTE_HOST_DEVICE static constexpr int scaleHeadStride(int sumSeqLens, int batchSize) + { + if constexpr (Partition::HeadStrideCountsBlocks) + { + return (sumSeqLens + TokensPerScale - 1) / TokensPerScale + batchSize - 1; + } + else + { + return scaleBase(sumSeqLens, batchSize); + } + } + + // Token index inside a tile for intra-group element i of group g. + CUTE_HOST_DEVICE static constexpr int tokenInTile(int i, int g) + { + return Layout{}(i, g); + } +}; + +} // namespace tensorrt_llm::common diff --git a/cpp/tensorrt_llm/common/sageQuant.cu b/cpp/tensorrt_llm/common/sageQuant.cu index be822f970e1e..305f39d86704 100644 --- a/cpp/tensorrt_llm/common/sageQuant.cu +++ b/cpp/tensorrt_llm/common/sageQuant.cu @@ -17,6 +17,7 @@ #include "sageQuant.h" +#include "sagePartition.h" #include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/cudaUtils.h" #include "tensorrt_llm/common/logger.h" @@ -31,40 +32,44 @@ namespace tensorrt_llm::common { -/// @brief SageAttn quantization kernel for Q and K +using tensorrt_llm::kernels::DATA_TYPE_BF16; +using tensorrt_llm::kernels::DATA_TYPE_E4M3; +using tensorrt_llm::kernels::DATA_TYPE_FP16; +using tensorrt_llm::kernels::DATA_TYPE_INT8; -// Quantization kernel for SageAttention doing 2 tasks per invocation: -// 1. Performs per-token-block quantization for Q **or** K, depending on the actual pointers passed in. -// 2. [Optional when gridDim.z>=2] Performs either per-channel Sfs gathering or per-channel quantization for V. +// SageAttention quantization kernel. Tensors are interpreted as column-major [D, H, S], which is +// the same physical layout as contiguous PyTorch [S, H, D]. Each launch quantizes Q or K. It can +// simultaneously collect V scales (VStage=1) or quantize V using scales collected by an earlier +// launch (VStage=2). When SmoothK is enabled, the VStage=1 launch also collects a per-head-per- +// channel K mean and the VStage=2 launch uses it to perform K-smoothing before quantization. // -// NOTE: all tensors in this file are treated as column-major: [D, H, S]. +// Which tokens share a scale is a template parameter: Partition is a CuTe layout +// ((intra...), (group...)) -> token offset inside a tile (see sagePartition.h). Contiguous blocks +// and swizzled per-thread groups are both instances, so this kernel serves both without knowing +// which it was given. // -// Q/K Sfs layout, matching the FMHA kernel: head stride is ceil(sumSeqLensQk / TokenPerScale) + batchSize - 1 and the -// local token `t` of sequence `b` uses `cumSeqLens[b] / TokenPerScale + b + t / TokenPerScale`, the `+ b` padding -// keeping a trailing partial block from being shared with the next sequence. -template +// Q/K scale layout: a sequence's scales start at PartitionTraits::scaleBase(cuSeqLens[b], b) and +// the head stride is scaleHeadStride(sumSeqLensQk, batchSize). Every sequence reserves a whole +// spare tile, so a trailing partial tile -- whose scales the consumer still indexes in full -- +// cannot run into the next sequence. +template __global__ void sageQuantQkvKernel(int sumSeqLensQk, int batchSize, int const* ptrCuSeqLensQk, void const* ptrQk, - void* ptrQkQuant, float* ptrQkScale, float* ptrKMean, int sumSeqLensV, int numHeadsV, void const* ptrV, - void* ptrVQuant, float* ptrVScale) + void* ptrQkQuant, float* ptrQkScale, void const* ptrKForMean, float* ptrKMean, int sumSeqLensV, int numHeadsV, + void const* ptrV, void* ptrVQuant, float* ptrVScale) { using namespace cute; using namespace cutlass; - static_assert(!KSmooth, "K-smoothing not implemented yet"); -#ifdef ENABLE_FP8 + using Traits = PartitionTraits; + constexpr int TokensPerScale = Traits::TokensPerScale; + constexpr int ScalesPerTile = Traits::ScalesPerTile; + constexpr int TileTokens = Traits::TileTokens; + static_assert(!SmoothK || VStage != 0, "K smoothing requires V staging"); static_assert(std::is_same_v || std::is_same_v, "Unrecognized target dtype for quantization"); constexpr float TypeMax = cute::is_same_v ? 448.0f : static_cast(126.9f); -#else - static_assert( - std::is_same_v, "Only int8 quantization is available without ENABLE_FP8."); - constexpr float TypeMax = static_cast(126.9f); -#endif constexpr int BestVL = 128 / sizeof_bits_v; using VL = Int; - // Silence currently-unused argument until K-smoothing support is added. - (void) ptrKMean; - int const numWarpsPerCta = blockDim.x / 32; int const numWarps = gridDim.x * numWarpsPerCta; int const warpId = blockIdx.x * numWarpsPerCta + threadIdx.x / 32; @@ -72,26 +77,25 @@ __global__ void sageQuantQkvKernel(int sumSeqLensQk, int batchSize, int const* p if (blockIdx.z == 0) { - // Qk task -- per-token-block quantization. blockIdx.y maps to (headIdx * batchSize + seqIdx). + // blockIdx.y maps to (headIdx * batchSize + seqIdx). int const numHeads = gridDim.y / batchSize; int const headIdx = blockIdx.y / batchSize; int const seqIdx = blockIdx.y % batchSize; - // Threads count per token block + // threadsPerScale threads cooperate on one scale group: each owns a BestVL-wide slice of + // the head dim and all TokensPerScale tokens of the group, and they reduce over the head + // dim with a shuffle. This is independent of how the group's tokens are laid out, which is + // why the partition only has to change the addresses, not the thread mapping. constexpr int threadsPerScale = HeadDim / BestVL; static_assert(HeadDim % BestVL == 0, "VL must divide HeadDim"); static_assert(threadsPerScale <= 32, "One token block should never exceed warp scope"); int const numScalesPerWarp = 32 / threadsPerScale; int const numScalesPerWave = numWarps * numScalesPerWarp; - - // Thread coordinates - int const tokBlkIdxInWave = warpId * numScalesPerWarp + thrId / threadsPerScale; + int const scaleIdxInWave = warpId * numScalesPerWarp + thrId / threadsPerScale; int const threadInScaleIdx = thrId % threadsPerScale; - // Lanes of this token block: a trailing block is taken by one group only, so never reduce over the warp. constexpr uint32_t scaleMask = threadsPerScale == 32 ? ~0u : ((1u << threadsPerScale) - 1u); uint32_t const laneMask = scaleMask << (thrId / threadsPerScale * threadsPerScale); - // This head and this sequence int const seqBegin = ptrCuSeqLensQk[seqIdx]; int const seqLen = ptrCuSeqLensQk[seqIdx + 1] - seqBegin; if (seqLen <= 0) @@ -99,138 +103,146 @@ __global__ void sageQuantQkvKernel(int sumSeqLensQk, int batchSize, int const* p return; } - // Sfs of this head, see the layout note above. float* ptrQkScaleHead - = ptrQkScale + static_cast(headIdx) * (ceil_div(sumSeqLensQk, TokenPerScale) + batchSize - 1); + = ptrQkScale + static_cast(headIdx) * Traits::scaleHeadStride(sumSeqLensQk, batchSize); int const tokenStride = numHeads * HeadDim; - - // IO tensors of this sequence and head int64_t const seqHeadOffset = static_cast(seqBegin) * tokenStride + headIdx * HeadDim; - Tensor gQkSeq = make_tensor(reinterpret_cast(ptrQk) + seqHeadOffset, - make_shape(Int{}, seqLen), make_stride(_1{}, tokenStride)); - Tensor gQkSeqQuant = make_tensor(reinterpret_cast(ptrQkQuant) + seqHeadOffset, - make_shape(Int{}, seqLen), make_stride(_1{}, tokenStride)); - Tensor gQkSeqScale = make_tensor( - ptrQkScaleHead + seqBegin / TokenPerScale + seqIdx, make_shape(ceil_div(seqLen, TokenPerScale))); - - // Tiling - Tensor gQkVecs = tiled_divide(gQkSeq, Shape>{}); - Tensor gQkVecsQuant = tiled_divide(gQkSeqQuant, Shape>{}); - - // Quantize one token block. A partial trailing block is zero-filled so it does not affect the Sfs. - auto quantizeTokBlk = [&](auto isFullBlk, int tokBlkIdx, int numValidTokens) + // This thread's BestVL-wide channel slice of token 0 of the sequence. A token is reached by + // adding tokenIdx * tokenStride, so the partition only contributes an index, never a + // stride. + Element const* ptrQkThread + = reinterpret_cast(ptrQk) + seqHeadOffset + threadInScaleIdx * BestVL; + ElementQuantized* ptrQkQuantThread + = reinterpret_cast(ptrQkQuant) + seqHeadOffset + threadInScaleIdx * BestVL; + float* ptrQkSeqScale = ptrQkScaleHead + Traits::scaleBase(seqBegin, seqIdx); + int const numTiles = ceil_div(seqLen, TileTokens); + + // tileIdx/grpIdx identify a scale; isFullTile is a compile-time fast path for tiles that + // lie entirely inside the sequence. In a partial tile validity has to be tested per token + // rather than as a prefix count, because a group's tokens are strided in general -- for + // SM90 Q the two rows of a group are 8 apart, so one can be in range while the other is + // not. + auto quantizeScale = [&](auto isFullTile, int tileIdx, int grpIdx) { - constexpr bool IsFullBlk = decltype(isFullBlk)::value; - - // Register buffers - Tensor rQk = make_tensor(Shape>{}); - Tensor rQkQuant = make_tensor(Shape>{}); - Tensor rQkCompute = make_tensor(Shape>{}); + constexpr bool IsFullTile = decltype(isFullTile)::value; + int const tileBase = tileIdx * TileTokens; + auto tokenOf = [&](int i) { return tileBase + Traits::tokenInTile(i, grpIdx); }; - // Compute tensors + Tensor rQk = make_tensor(Shape>{}); + Tensor rQkQuant = make_tensor(Shape>{}); + Tensor rQkCompute = make_tensor(Shape>{}); Tensor rQk_x2 = recast>(rQk); Tensor rQkCompute_x2 = recast>(rQkCompute); - // Conversion tensors Tensor rQk_x4 = recast>(rQk); Tensor rQkQuant_x4 = recast>(rQkQuant); - // Load input - if constexpr (IsFullBlk) + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < TokensPerScale; ++i) { - cute::copy(AutoVectorizingCopy{}, gQkVecs(_, threadInScaleIdx, tokBlkIdx), rQk); + int const tokenIdx = tokenOf(i); + if (IsFullTile || tokenIdx < seqLen) + { + Tensor gSrc = make_tensor( + make_gmem_ptr(ptrQkThread + static_cast(tokenIdx) * tokenStride), Shape{}); + cute::copy(AutoVectorizingCopy{}, gSrc, rQk(_, i)); + } + else + { + CUTLASS_PRAGMA_UNROLL + for (int j = 0; j < BestVL; ++j) + { + rQk(j, i) = static_cast(0); + } + } } - else + cute::transform(rQk_x2, rQkCompute_x2, NumericArrayConverter::convert); + + if constexpr (SmoothK && VStage == 2) { + float const* ptrKMeanHead = ptrKMean + headIdx * HeadDim; CUTLASS_PRAGMA_UNROLL - for (int i = 0; i < size<1>(rQk); ++i) + for (int tokenIdx = 0; tokenIdx < TokensPerScale; ++tokenIdx) { - if (i < numValidTokens) - { - cute::copy( - AutoVectorizingCopy{}, gQkVecs(make_tuple(_, i), threadInScaleIdx, tokBlkIdx), rQk(_, i)); - } - else + if (IsFullTile || tokenOf(tokenIdx) < seqLen) { CUTLASS_PRAGMA_UNROLL - for (int j = 0; j < BestVL; ++j) + for (int vecIdx = 0; vecIdx < BestVL; ++vecIdx) { - rQk(j, i) = static_cast(0); + rQkCompute(vecIdx, tokenIdx) -= ptrKMeanHead[threadInScaleIdx * BestVL + vecIdx]; } } } } - cute::transform(rQk_x2, rQkCompute_x2, NumericArrayConverter::convert); - // Intra-thread reduction float maxScale = 1e-3f; CUTLASS_PRAGMA_UNROLL for (int i = 0; i < size(rQk); ++i) { maxScale = ::fmaxf(maxScale, ::fabsf(rQkCompute(i))); } - // Intra-warp reduction CUTLASS_PRAGMA_UNROLL for (int delta = 1; delta < threadsPerScale; delta <<= 1) { maxScale = ::fmaxf(maxScale, __shfl_xor_sync(laneMask, maxScale, delta)); } - // Rescale to TypeMax - maxScale = maxScale / TypeMax; - // Store maxScale - gQkSeqScale(tokBlkIdx) = maxScale; - - // 1/maxScale - Array scaleQuant - = NumericArrayConverter::convert(Array{maxScale, maxScale}); - scaleQuant = cutlass::reciprocal_approximate>{}(scaleQuant); - cutlass::multiplies> scaleQuantOp; - // Qk /= maxScale - cute::transform(rQk_x2, rQk_x2, [&](auto& x) { return scaleQuantOp(x, scaleQuant); }); - // Convert to target quant type - cute::transform(rQk_x4, rQkQuant_x4, NumericArrayConverter::convert); - - // Store quantized output - if constexpr (IsFullBlk) + maxScale /= TypeMax; + // Every group of every tile gets a scale, including groups of a partial tile whose + // tokens are all out of range: the consumer indexes a partial tile in full, and masks + // with a large negative sentinel that must stay negative after being multiplied by this + // scale. The 1e-3 floor above is what guarantees that -- do not lower it to an epsilon. + ptrQkSeqScale[tileIdx * ScalesPerTile + grpIdx] = maxScale; + float const invScale = 1.0f / maxScale; + if constexpr (SmoothK && VStage == 2) { - cute::copy(AutoVectorizingCopy{}, rQkQuant, gQkVecsQuant(_, threadInScaleIdx, tokBlkIdx)); + // rQkCompute holds the result value + Tensor rQkCompute_x4 = recast>(rQkCompute); + cute::transform(rQkCompute, rQkCompute, [&](float const& x) { return x * invScale; }); + cute::transform(rQkCompute_x4, rQkQuant_x4, NumericArrayConverter::convert); } else { - CUTLASS_PRAGMA_UNROLL - for (int i = 0; i < size<1>(rQk); ++i) + // rQk holds the result value + Array scaleQuant + = NumericArrayConverter::convert(Array{invScale, invScale}); + cutlass::multiplies> scaleQuantOp; + cute::transform(rQk_x2, rQk_x2, [&](auto& x) { return scaleQuantOp(x, scaleQuant); }); + cute::transform(rQk_x4, rQkQuant_x4, NumericArrayConverter::convert); + } + + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < TokensPerScale; ++i) + { + int const tokenIdx = tokenOf(i); + if (IsFullTile || tokenIdx < seqLen) { - if (i < numValidTokens) - { - cute::copy(AutoVectorizingCopy{}, rQkQuant(_, i), - gQkVecsQuant(make_tuple(_, i), threadInScaleIdx, tokBlkIdx)); - } + Tensor gDst = make_tensor( + make_gmem_ptr(ptrQkQuantThread + static_cast(tokenIdx) * tokenStride), Shape{}); + cute::copy(AutoVectorizingCopy{}, rQkQuant(_, i), gDst); } } }; - int const numWholeScales = seqLen / TokenPerScale; - - // Unpredicated iterations - for (int tokBlkIdx = tokBlkIdxInWave; tokBlkIdx < numWholeScales; tokBlkIdx += numScalesPerWave) + // Work items are (tile, group) pairs. Whole tiles take the fast path; only the last tile of + // a sequence can be partial, and all of its groups still need a scale written. + int const numWholeTiles = seqLen / TileTokens; + int const numWholeScales = numWholeTiles * ScalesPerTile; + for (int workIdx = scaleIdxInWave; workIdx < numWholeScales; workIdx += numScalesPerWave) { - quantizeTokBlk(cute::true_type{}, tokBlkIdx, TokenPerScale); + quantizeScale(cute::true_type{}, workIdx / ScalesPerTile, workIdx % ScalesPerTile); } - - // Predicated iteration, taken by the group owning the trailing block - int const numTailTokens = seqLen - numWholeScales * TokenPerScale; - if (numTailTokens > 0 && tokBlkIdxInWave == numWholeScales % numScalesPerWave) + if (numWholeTiles < numTiles) { - quantizeTokBlk(cute::false_type{}, numWholeScales, numTailTokens); + for (int grpIdx = scaleIdxInWave; grpIdx < ScalesPerTile; grpIdx += numScalesPerWave) + { + quantizeScale(cute::false_type{}, numWholeTiles, grpIdx); + } } } else if (blockIdx.z == 1) { - // V task -- per-channel (all tokens) 2-stage task. blockIdx.y maps to headIdx. int const headIdx = blockIdx.y; using ElementQuantizedV = cutlass::float_e4m3_t; - - // IO tensors constexpr int threadsPerHead = HeadDim / BestVL; static_assert(HeadDim % BestVL == 0, "VL must divide HeadDim"); static_assert(threadsPerHead <= 32, "One token block should never exceed warp scope"); @@ -240,78 +252,54 @@ __global__ void sageQuantQkvKernel(int sumSeqLensQk, int batchSize, int const* p make_shape(VL{}, Int{}, numHeadsV, sumSeqLensV)); Tensor gVScale = make_tensor(ptrVScale, make_shape(VL{}, Int{}, numHeadsV)); - // Register buffers Tensor rV = make_tensor(Shape{}); Tensor rVMax = make_tensor(Shape{}); Tensor rVQuant = make_tensor(Shape{}); Tensor rVScale = make_tensor(Shape{}); Tensor rVCompute = make_tensor(Shape{}); - - // Compute tensors Tensor rV_x2 = recast>(rV); Tensor rVMax_x2 = recast>(rVMax); Tensor rVScale_x2 = recast>(rVScale); Tensor rVCompute_x2 = recast>(rVCompute); - - // Conversion tensors Tensor rVCompute_x4 = recast>(rVCompute); Tensor rVQuant_x4 = recast>(rVQuant); - // If the parallel on-going task is handling Q, numHeads inferred from gridDim.y could be larger than numHeadsKv if (headIdx < numHeadsV) { - // Thread coordinates int const numToksPerWarp = 32 / threadsPerHead; int tokIdx = warpId * numToksPerWarp + thrId / threadsPerHead; int const threadInTokIdx = thrId % threadsPerHead; - - // Thread-local tensors Tensor gVSeq = gV(_, threadInTokIdx, headIdx, _); Tensor gVSeqQuant = gVQuant(_, threadInTokIdx, headIdx, _); Tensor gVSeqScale = gVScale(_, threadInTokIdx, headIdx); if constexpr (VStage == 1) { - // Stage 1: reduction to obtain the Sfs - - // Avoid heavy atomics: limit the number of warps. int const numWarpsToUse = cutlass::fast_min(numWarps, 256); int const numToksPerWave = numWarpsToUse * numToksPerWarp; if (warpId >= numWarpsToUse) { return; } - - // Initialize CUTLASS_PRAGMA_UNROLL for (int i = 0; i < size(rVScale); ++i) { rVScale(i) = 1e-3f; } cute::transform(rVScale_x2, rVMax_x2, cutlass::NumericArrayConverter::convert); - - // Loop over all tokens for (; tokIdx < sumSeqLensV; tokIdx += numToksPerWave) { - // Load inputs cute::copy(AutoVectorizingCopy{}, gVSeq(_, tokIdx), rV); - // Compute abs-max cute::transform(rV_x2, rV_x2, cutlass::absolute_value_op>{}); cute::transform(rV_x2, rVMax_x2, rVMax_x2, cutlass::maximum>{}); } - - // Transform max to Sfs cute::transform(rVMax_x2, rVScale_x2, cutlass::NumericArrayConverter::convert); cute::transform(rVScale_x2, rVScale_x2, cutlass::scale>{1 / 448.0f}); - - // Intra-warp reduction. for (int delta = threadsPerHead; delta < 32; delta <<= 1) { cute::transform(rVScale, rVScale, [&](auto const& x) { return ::fmaxf(x, __shfl_xor_sync(0xffffffffu, x, delta)); }); } - - // Atomic reduction into global memory. if (threadInTokIdx == thrId) { CUTLASS_PRAGMA_UNROLL @@ -324,34 +312,84 @@ __global__ void sageQuantQkvKernel(int sumSeqLensQk, int batchSize, int const* p } else if constexpr (VStage == 2) { - // Stage 2: scale according to the Sfs - - // Full waves. int const numToksPerWave = numWarps * numToksPerWarp; - - // Load Sfs cute::copy(AutoVectorizingCopy{}, gVSeqScale, rVScale); - // Take reciprocal cute::transform(rVScale_x2, rVScale_x2, cutlass::reciprocal_approximate>{}); - - // Loop over all tokens for (; tokIdx < sumSeqLensV; tokIdx += numToksPerWave) { - // Load inputs cute::copy(AutoVectorizingCopy{}, gVSeq(_, tokIdx), rV); - // Convert up cute::transform(rV_x2, rVCompute_x2, cutlass::NumericArrayConverter::convert); - // Scale cute::transform(rVCompute_x2, rVScale_x2, rVCompute_x2, cutlass::multiplies>{}); - // Convert (quantize) cute::transform( rVCompute_x4, rVQuant_x4, cutlass::NumericArrayConverter::convert); - // Write output cute::copy(AutoVectorizingCopy{}, rVQuant, gVSeqQuant(_, tokIdx)); } } } } + else if (blockIdx.z == 2) + { + if constexpr (SmoothK && VStage == 1) + { + int const headIdx = blockIdx.y; + constexpr int threadsPerHead = HeadDim / BestVL; + static_assert(HeadDim % BestVL == 0, "VL must divide HeadDim"); + static_assert(threadsPerHead <= 32, "One token block should never exceed warp scope"); + + if (headIdx < numHeadsV) + { + Tensor gK = make_tensor(reinterpret_cast(ptrKForMean), + make_shape(VL{}, Int{}, numHeadsV, sumSeqLensV)); + Tensor gKMean = make_tensor(ptrKMean, make_shape(VL{}, Int{}, numHeadsV)); + Tensor rK = make_tensor(Shape{}); + Tensor rKCompute = make_tensor(Shape{}); + Tensor rKSum = make_tensor(Shape{}); + Tensor rK_x2 = recast>(rK); + Tensor rKCompute_x2 = recast>(rKCompute); + + int const numToksPerWarp = 32 / threadsPerHead; + int const numWarpsToUse = cutlass::fast_min(numWarps, 256); + int const numToksPerWave = numWarpsToUse * numToksPerWarp; + if (warpId >= numWarpsToUse) + { + return; + } + int tokIdx = warpId * numToksPerWarp + thrId / threadsPerHead; + int const threadInTokIdx = thrId % threadsPerHead; + Tensor gKSeq = gK(_, threadInTokIdx, headIdx, _); + Tensor gKMeanHead = gKMean(_, threadInTokIdx, headIdx); + + clear(rKSum); + for (; tokIdx < sumSeqLensV; tokIdx += numToksPerWave) + { + cute::copy(AutoVectorizingCopy{}, gKSeq(_, tokIdx), rK); + cute::transform(rK_x2, rKCompute_x2, NumericArrayConverter::convert); + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < BestVL; ++i) + { + rKSum(i) += rKCompute(i); + } + } + for (int delta = threadsPerHead; delta < 32; delta <<= 1) + { + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < BestVL; ++i) + { + rKSum(i) += __shfl_xor_sync(0xffffffffu, rKSum(i), delta); + } + } + if (threadInTokIdx == thrId) + { + float const invNumTokens = 1.0f / sumSeqLensV; + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < BestVL; ++i) + { + atomicAdd(&gKMeanHead(i), rKSum(i) * invNumTokens); + } + } + } + } + } } template @@ -359,118 +397,164 @@ void invokeSageQuantQkvImpl(SageQuantParams const& params) { using namespace cute; TLLM_CHECK_WITH_INFO(params.sumSeqLensQk > 0 && params.batchSize > 0 && params.ptrCuSeqLensQk != nullptr - && params.numHeads > 0 && params.headDim > 0 && params.tokenBlockSize > 0 && params.ptrQk != nullptr - && params.ptrQkQuant != nullptr && params.ptrQkScale != nullptr && params.smCount > 0, - "Invalid SageQuantQk parameters."); + && params.numHeads > 0 && params.headDim > 0 + && (params.partition != SageScalePartition::Contiguous || params.tokenBlockSize > 0) + && params.ptrQk != nullptr && params.ptrQkQuant != nullptr && params.ptrQkScale != nullptr + && params.smCount > 0, + "Invalid SageQuantQk parameters"); TLLM_CHECK_WITH_INFO(params.vStage == 0 || (params.sumSeqLensV > 0 && params.numHeadsV > 0 && params.ptrV != nullptr && params.ptrVQuant != nullptr && params.ptrVScale != nullptr), - "Invalid SageQuantV parameters."); - TLLM_CHECK_WITH_INFO(!params.kSmooth, "SageQuantQk K-smoothing is not supported yet."); + "Invalid SageQuantV parameters"); + TLLM_CHECK_WITH_INFO(!params.kSmooth || params.vStage != 0, "SageQuant K smoothing requires V staging"); + TLLM_CHECK_WITH_INFO( + !params.kSmooth || (params.ptrKMean != nullptr && (params.vStage != 1 || params.ptrKForMean != nullptr)), + "Invalid SageQuant K-smoothing parameters"); - auto invokeKernel = [&](auto headDimStatic, auto tokenBlockSizeStatic) + auto invokeKernel = [&](auto headDimStatic, auto partitionStatic) { constexpr int HeadDim_ = headDimStatic; - constexpr int TokenBlockSize_ = tokenBlockSizeStatic; - + using Partition_ = decltype(partitionStatic); SageQuantParams kernelParams = params; - void* kernelArgs[] - = {&kernelParams.sumSeqLensQk, &kernelParams.batchSize, &kernelParams.ptrCuSeqLensQk, &kernelParams.ptrQk, - &kernelParams.ptrQkQuant, &kernelParams.ptrQkScale, &kernelParams.ptrKMean, &kernelParams.sumSeqLensV, - &kernelParams.numHeadsV, &kernelParams.ptrV, &kernelParams.ptrVQuant, &kernelParams.ptrVScale}; + void* kernelArgs[] = {&kernelParams.sumSeqLensQk, &kernelParams.batchSize, &kernelParams.ptrCuSeqLensQk, + &kernelParams.ptrQk, &kernelParams.ptrQkQuant, &kernelParams.ptrQkScale, &kernelParams.ptrKForMean, + &kernelParams.ptrKMean, &kernelParams.sumSeqLensV, &kernelParams.numHeadsV, &kernelParams.ptrV, + &kernelParams.ptrVQuant, &kernelParams.ptrVScale}; - auto launchWithVStage = [&](auto vStageStatic) + auto launchKernel = [&](auto smoothKStatic, auto vStageStatic) { + constexpr bool SmoothK_ = decltype(smoothKStatic)::value; constexpr int VStage_ = vStageStatic; void const* kernelFunc = nullptr; - if (params.quantType == kernels::DATA_TYPE_E4M3) + if (params.quantType == DATA_TYPE_E4M3) { -#ifdef ENABLE_FP8 kernelFunc = reinterpret_cast( - sageQuantQkvKernel); -#else - TLLM_THROW("SageQuantQk FP8 quantization requires ENABLE_FP8."); -#endif + sageQuantQkvKernel); } - else if (params.quantType == kernels::DATA_TYPE_INT8) + else if (params.quantType == DATA_TYPE_INT8) { kernelFunc = reinterpret_cast( - sageQuantQkvKernel); + sageQuantQkvKernel); } else { - TLLM_THROW("Unsupported SageQuantQk quantType: %d.", static_cast(params.quantType)); + TLLM_THROW("SageQuant Q/K output must be INT8 or FP8 E4M3"); } - - // One block of the y dimension per (head, sequence) for Qk, per head for V. int const numHeadSeqs = params.numHeads * params.batchSize; uint32_t const gridX = static_cast(std::max(1, (params.smCount * 32) / numHeadSeqs)); - uint32_t const gridY = static_cast(numHeadSeqs); - uint32_t const gridZ = VStage_ > 0 ? 2U : 1U; - dim3 const launchGrid{gridX, gridY, gridZ}; - check_cuda_error(cudaLaunchKernel(kernelFunc, launchGrid, dim3{64U, 1U, 1U}, kernelArgs, 0, params.stream)); - check_cuda_error(cudaPeekAtLastError()); + constexpr uint32_t GridZ = VStage_ == 0 ? 1U : (SmoothK_ && VStage_ == 1 ? 3U : 2U); + dim3 const launchGrid{gridX, static_cast(numHeadSeqs), GridZ}; + auto status = cudaLaunchKernel(kernelFunc, launchGrid, dim3{64U, 1U, 1U}, kernelArgs, 0, params.stream); + TLLM_CHECK_WITH_INFO(status == cudaSuccess, "%s", cudaGetErrorString(status)); + status = cudaPeekAtLastError(); + TLLM_CHECK_WITH_INFO(status == cudaSuccess, "%s", cudaGetErrorString(status)); }; switch (params.vStage) { - case 0: launchWithVStage(Int<0>{}); return; - case 1: launchWithVStage(Int<1>{}); return; - case 2: launchWithVStage(Int<2>{}); return; - default: TLLM_THROW("Unsupported SageQuantV stage: %d.", params.vStage); + case 0: launchKernel(cute::false_type{}, Int<0>{}); return; + case 1: + if (params.kSmooth) + { + launchKernel(cute::true_type{}, Int<1>{}); + } + else + { + launchKernel(cute::false_type{}, Int<1>{}); + } + return; + case 2: + if (params.kSmooth) + { + launchKernel(cute::true_type{}, Int<2>{}); + } + else + { + launchKernel(cute::false_type{}, Int<2>{}); + } + return; + default: TLLM_THROW("Unsupported SageQuantV stage %d", params.vStage); } }; - // Dispatch - if (params.headDim == 64) - { - switch (params.tokenBlockSize) - { - case 1: invokeKernel(Int<64>{}, Int<1>{}); return; - case 4: invokeKernel(Int<64>{}, Int<4>{}); return; - case 16: invokeKernel(Int<64>{}, Int<16>{}); return; - default: break; - } - } - if (params.headDim == 128) - { - switch (params.tokenBlockSize) - { - case 1: invokeKernel(Int<128>{}, Int<1>{}); return; - case 4: invokeKernel(Int<128>{}, Int<4>{}); return; - case 16: invokeKernel(Int<128>{}, Int<16>{}); return; - default: break; - } - } - if (params.headDim == 256) - { - switch (params.tokenBlockSize) - { - case 1: invokeKernel(Int<256>{}, Int<1>{}); return; - case 4: invokeKernel(Int<256>{}, Int<4>{}); return; - case 16: invokeKernel(Int<256>{}, Int<16>{}); return; - default: break; - } +#define TLLM_SAGE_DISPATCH_HEAD_DIM(HEAD_DIM) \ + if (params.headDim == HEAD_DIM) \ + { \ + switch (params.partition) \ + { \ + case SageScalePartition::HopperQ: invokeKernel(Int{}, HopperQPartition{}); return; \ + case SageScalePartition::HopperK: invokeKernel(Int{}, HopperKPartition{}); return; \ + case SageScalePartition::Contiguous: \ + switch (params.tokenBlockSize) \ + { \ + case 1: invokeKernel(Int{}, ContiguousPartition<1>{}); return; \ + case 4: invokeKernel(Int{}, ContiguousPartition<4>{}); return; \ + case 16: invokeKernel(Int{}, ContiguousPartition<16>{}); return; \ + default: break; \ + } \ + break; \ + } \ } + TLLM_SAGE_DISPATCH_HEAD_DIM(64) + TLLM_SAGE_DISPATCH_HEAD_DIM(128) + TLLM_SAGE_DISPATCH_HEAD_DIM(256) +#undef TLLM_SAGE_DISPATCH_HEAD_DIM TLLM_THROW( - "Unsupported SageQuantQk dispatch config: headDim=%d tokenBlockSize=%d", params.headDim, params.tokenBlockSize); + "Unsupported SageQuant dispatch config (head_dim must be 64, 128, or 256; contiguous token_block_size must " + "be 1, 4, or 16): headDim=%d partition=%d tokenBlockSize=%d", + params.headDim, static_cast(params.partition), params.tokenBlockSize); } void invokeSageQuant(SageQuantParams const& params) { - if (params.inputType == kernels::DATA_TYPE_FP16) + if (params.inputType == DATA_TYPE_FP16) { invokeSageQuantQkvImpl(params); return; } -#ifdef ENABLE_BF16 - if (params.inputType == kernels::DATA_TYPE_BF16) + if (params.inputType == DATA_TYPE_BF16) { invokeSageQuantQkvImpl(params); return; } -#endif - TLLM_THROW("Unsupported SageQuantQk inputType: %d", static_cast(params.inputType)); + TLLM_THROW("SageQuant input must be FP16 or BF16"); +} + +namespace +{ + +// Evaluate `fn` with the PartitionTraits the (partition, tokenBlockSize) pair selects, or return +// `fallback` when there is no such instantiation. +template +int withPartitionTraits(SageScalePartition partition, int tokenBlockSize, int fallback, Fn&& fn) +{ + switch (partition) + { + case SageScalePartition::HopperQ: return fn(PartitionTraits{}); + case SageScalePartition::HopperK: return fn(PartitionTraits{}); + case SageScalePartition::Contiguous: + switch (tokenBlockSize) + { + case 1: return fn(PartitionTraits>{}); + case 4: return fn(PartitionTraits>{}); + case 16: return fn(PartitionTraits>{}); + default: break; + } + break; + } + return fallback; +} + +} // namespace + +int getSageScaleHeadStride(SageScalePartition partition, int tokenBlockSize, int sumSeqLens, int batchSize) +{ + if (tokenBlockSize <= 0 || sumSeqLens <= 0 || batchSize <= 0) + { + return 0; + } + return withPartitionTraits(partition, tokenBlockSize, 0, + [&](auto traits) { return decltype(traits)::scaleHeadStride(sumSeqLens, batchSize); }); } } // namespace tensorrt_llm::common diff --git a/cpp/tensorrt_llm/common/sageQuant.h b/cpp/tensorrt_llm/common/sageQuant.h index 80e74d840e70..8be6835de483 100644 --- a/cpp/tensorrt_llm/common/sageQuant.h +++ b/cpp/tensorrt_llm/common/sageQuant.h @@ -28,6 +28,17 @@ namespace tensorrt_llm::common { +// How Q or K tokens are grouped into scales. See sagePartition.h. +enum class SageScalePartition +{ + // tokenBlockSize consecutive tokens share a scale, one group per tile. + Contiguous = 0, + // 2 rows 8 apart share a scale; 32 groups per 64-row tile. + HopperQ, + // 16 strided keys share a scale; 16 groups per 256-key tile. + HopperK, +}; + struct SageQuantParams { // Required arguments for SageQuantQk (Q or K): @@ -35,7 +46,9 @@ struct SageQuantParams int batchSize{}; int numHeads{}; int headDim{}; + // Only read when partition == Contiguous; the Hopper partitions fix their own group size. int tokenBlockSize{}; + SageScalePartition partition{SageScalePartition::Contiguous}; bool kSmooth{false}; int const* ptrCuSeqLensQk{nullptr}; void const* ptrQk{nullptr}; @@ -43,6 +56,9 @@ struct SageQuantParams kernels::Data_type inputType{kernels::DATA_TYPE_FP16}; kernels::Data_type quantType{kernels::DATA_TYPE_E4M3}; float* ptrQkScale{nullptr}; + // Optional source and scratch mean used to perform K-smoothing. + // (See below) collected at vStage==0, applied at vStage==1. + void const* ptrKForMean{nullptr}; float* ptrKMean{nullptr}; // Optional arguments for SageQuantV: // vStage: 0: disabled, 1: collect scales, 2: quantize @@ -52,11 +68,32 @@ struct SageQuantParams void const* ptrV{nullptr}; void* ptrVQuant{nullptr}; float* ptrVScale{nullptr}; - // Hardware into. Required. + // Hardware info. Required. int smCount{}; cudaStream_t stream{}; }; void invokeSageQuant(SageQuantParams const& params); +// The scale grouping a consumer kernel expects. SM90 scales the tokens a single thread owns; the +// other architectures scale contiguous token blocks. +inline SageScalePartition getSageQPartition(bool perThread) +{ + return perThread ? SageScalePartition::HopperQ : SageScalePartition::Contiguous; +} + +inline SageScalePartition getSageKPartition(bool perThread) +{ + return perThread ? SageScalePartition::HopperK : SageScalePartition::Contiguous; +} + +// Scale-buffer geometry for a partition, so that callers can size the scale buffers and fill in +// the max_nblock the consumer kernel reads. Returns 0 when the tensor is not sage-quantized +// (tokenBlockSize 0), when the shape is empty, or when the (partition, tokenBlockSize) pair has no +// kernel instantiation, so it is safe to call from noexcept sizing paths. + +// Scales per head, i.e. the head stride of the scale buffer. Grows with the batch size as well as +// the token count, since every sequence reserves a spare tile. +int getSageScaleHeadStride(SageScalePartition partition, int tokenBlockSize, int sumSeqLens, int batchSize); + } // namespace tensorrt_llm::common diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_q_k_v_128_sage_2_16_1_output_bf16_tma_ws_sm90.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_q_k_v_128_sage_2_16_1_output_bf16_tma_ws_sm90.cubin.tar.zst new file mode 100644 index 000000000000..15391398dd07 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_q_k_v_128_sage_2_16_1_output_bf16_tma_ws_sm90.cubin.tar.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d2158b8cd7e9447fcc3c03b3be1cbae79e9cba71b836ffaf44971349c3b4981 +size 29820 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_qkv_128_sage_64_64_256_output_bf16_tma_ws_sm90.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_qkv_128_sage_64_64_256_output_bf16_tma_ws_sm90.cubin.tar.zst deleted file mode 100644 index a59880793ba6..000000000000 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_64_256_S_qkv_128_sage_64_64_256_output_bf16_tma_ws_sm90.cubin.tar.zst +++ /dev/null @@ -1,3 +0,0 @@ -version https://git-lfs.github.com/spec/v1 -oid sha256:32df67beec207fb10e70b3263ca63691c7663c4bdcbd2cdf465944423b01f202 -size 30671 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_bf16_sm89.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_bf16_sm89.cubin.tar.zst deleted file mode 100644 index dc1a01323257..000000000000 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_bf16_sm89.cubin.tar.zst +++ /dev/null @@ -1,3 +0,0 @@ -version https://git-lfs.github.com/spec/v1 -oid sha256:9b40ba309cac166a460d23a52c81b7d447b29fd713a93d9c61bd5dafa8b6da01 -size 9903 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_fp16_sm89.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_fp16_sm89.cubin.tar.zst deleted file mode 100644 index da5e93cacffc..000000000000 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_128_sage_64_32_32_output_fp16_sm89.cubin.tar.zst +++ /dev/null @@ -1,3 +0,0 @@ -version https://git-lfs.github.com/spec/v1 -oid sha256:ed7fd86676ec365f4b940f0704788d365002b7a58be03267203ce3437ca6c911 -size 9902 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_bf16_sm89.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_bf16_sm89.cubin.tar.zst deleted file mode 100644 index c911a792381e..000000000000 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_bf16_sm89.cubin.tar.zst +++ /dev/null @@ -1,3 +0,0 @@ -version https://git-lfs.github.com/spec/v1 -oid sha256:8087c7edc83f9f160bff87aac8b305ffde485d1ce3c8fa699c47ed79f6bda62a -size 8986 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_fp16_sm89.cubin.tar.zst b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_fp16_sm89.cubin.tar.zst deleted file mode 100644 index 36775c513196..000000000000 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/cubin/fmha_v2_flash_attention_e4m3_fp32_64_32_S_qkv_80_sage_64_32_32_output_fp16_sm89.cubin.tar.zst +++ /dev/null @@ -1,3 +0,0 @@ -version https://git-lfs.github.com/spec/v1 -oid sha256:d475071e4730387dc3ca7713a24cbdfac98c4b00c0f5a1f9d5f558ac4b93921f -size 8987 diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.cpp b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.cpp index 72b4e6128331..3cafcbc162ce 100644 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.cpp +++ b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.cpp @@ -87,8 +87,8 @@ FusedMHARunnerV2::FusedMHARunnerV2(MHARunnerFixedParams fixedParams) TLLM_CHECK_WITH_INFO((mSM == kSM_80 || mSM == kSM_86 || mSM == kSM_89 || mSM == kSM_90 || tensorrt_llm::common::isSM100Family(mSM) || mSM == kSM_120 || mSM == kSM_121), "Unsupported architecture"); - TLLM_CHECK_WITH_INFO((mFixedParams.dataType == DATA_TYPE_FP16 || mFixedParams.dataType == DATA_TYPE_BF16 - || mFixedParams.dataType == DATA_TYPE_E4M3), + TLLM_CHECK_WITH_INFO( + (mFixedParams.dataType == DATA_TYPE_FP16 || mFixedParams.dataType == DATA_TYPE_BF16 || isFp8Selected()), "Unsupported data type"); if (tensorrt_llm::common::isSM100Family(mSM)) { @@ -277,7 +277,7 @@ void FusedMHARunnerV2::setupKernelParams(MHARunnerParams runnerParams) // 2 scales prepared for scaleBmm1 in the device memory: float scale, float (scale with log2e). int64_t scaleBmm1PtrOffset = (mLaunchParams.useBase2ExpTrick ? kIdxScaleSoftmaxLog2Ptr : kIdxScaleSoftmaxPtr); // Only fp8 kernels need to load scales from the device memory. - if (mFixedParams.dataType == DATA_TYPE_E4M3) + if (isFp8Selected()) { mKernelParams.scale_bmm1_d = reinterpret_cast(runnerParams.scaleBmm1Ptr + scaleBmm1PtrOffset); mKernelParams.scale_bmm2_d = reinterpret_cast(runnerParams.scaleBmm2Ptr); @@ -321,7 +321,7 @@ void FusedMHARunnerV2::setupLaunchParams(MHARunnerParams runnerParams) mLaunchParams.enableAttnLogitSoftcapping = mFixedParams.attnLogitSoftcappingScale != 0.f; // BF16 FMHA only accumulates on FP32. // E4M3 FMHA only supports fp32 accumulation currently. - mLaunchParams.force_fp32_acc = mFixedParams.dataType == DATA_TYPE_BF16 || mFixedParams.dataType == DATA_TYPE_E4M3 + mLaunchParams.force_fp32_acc = mFixedParams.dataType == DATA_TYPE_BF16 || isFp8Selected() || mFixedParams.forceFp32Acc || runnerParams.forceFp32Acc; // The attention mask type. mLaunchParams.attention_mask_type = mFixedParams.attentionMaskType; @@ -382,7 +382,7 @@ void FusedMHARunnerV2::setupLaunchParams(MHARunnerParams runnerParams) // Only warp-specialized FMHA kernels support FP8 on Hopper. // Separate Q + KV input layout: enable warp-specialization kernels when s > 512, otherwise use ampere-style flash // attention kernels. - if (isSm90 && (mFixedParams.dataType == DATA_TYPE_E4M3 || (separateQKvInput && runnerParams.kvSeqLen > 512))) + if (isSm90 && (isFp8Selected() || (separateQKvInput && runnerParams.kvSeqLen > 512))) { mLaunchParams.flash_attention = true; mLaunchParams.force_unroll = true; diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.h b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.h index ab2c82a54451..cebee69d79fb 100644 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.h +++ b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fmhaRunner.h @@ -81,6 +81,14 @@ class FusedMHARunnerV2 // Get the kernel sequence that support the max sequence length (only used by non-flash-attention kernels). int getSFromMaxSeqLen(int const max_seq_len) const; + // Is the input selected as fp8? SageAttention is selected as DATA_TYPE_KV_INT8_E4M3, which + // records that it runs int8 Q/K with an e4m3 PV. Everywhere in this runner it behaves like the + // plain fp8 kernels: one byte per element, fp32 accumulation, and device-side scales. + bool isFp8Selected() const + { + return mFixedParams.dataType == DATA_TYPE_E4M3 || mFixedParams.dataType == DATA_TYPE_KV_INT8_E4M3; + } + private: // The attention fixed params (mostly related to the attention structure). MHARunnerFixedParams mFixedParams; diff --git a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h index a5b0c80a8428..9fbe72c1d73b 100644 --- a/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h +++ b/cpp/tensorrt_llm/kernels/contextFusedMultiHeadAttention/fused_multihead_attention_common.h @@ -137,11 +137,11 @@ struct MHARunnerFixedParams int tpSize = 1; // The tensor parallel rank (alibi). int tpRank = 0; - // q tensor quant block size in sage attention + // q tensor quant block size in sage attention, in tokens per scale. 0 disables sage attention. int sageBlockSizeQ = 0; - // k tensor quant block size in sage attention + // k tensor quant block size in sage attention, in tokens per scale. int sageBlockSizeK = 0; - // v tensor quant block size in sage attention + // v tensor quant block size in sage attention, in channels per scale. int sageBlockSizeV = 0; // Use sparse MLA ? bool useSparseMLA = false; @@ -341,7 +341,7 @@ struct MHARunnerParams float* qScalePtr; float* kScalePtr; float* vScalePtr; - // q, k, v block size in sageattention + // number of scales per head for q, k, v in sageattention int qMaxNBlock; int kMaxNBlock; int vMaxNBlock; @@ -491,14 +491,29 @@ struct Fused_multihead_attention_params_v2 // is input/output padded bool is_s_padded = false; - // SageAttention parameters + // SageAttention parameters. This is the reference description of the scale layout; the + // quantizer that produces these buffers lives in common/sageQuant.cu. + // + // Q and K are quantized to INT8 along the sequence axis and their scales are amax/127; V is + // quantized to e4m3 along the channel axis, one scale per channel, with the amax taken over + // every token of every sequence. + // + // How Q and K tokens are grouped into scales depends on the GPU. SM100 groups a run of + // sage_block_size consecutive tokens. SM90 groups the tokens a single thread owns -- 2 query + // rows and 16 keys, both strided -- and stores the scales swizzled so that a thread's scales + // are contiguous, so a q/k scale buffer is not a plain per-token-block array. + // + // The q/k buffers have no batch dimension. A sequence's scales start at a base derived from + // its cu_seqlens entry, so the buffer can be sized from the batch size and the total token + // count alone, without knowing the longest sequence. Each sequence reserves a spare tile, + // because a partial final tile is still indexed in full. struct SageAttention { struct Scales { - // ceil(max_seqlen / block_size) + // number of scales per head, i.e. the stride between heads. Unused for v. int max_nblock; - // The scale of each block, layout: (B, H, max_nblock) + // the scales, layout: (H, max_nblock) for q and k, (H_kv, D) for v float* scales; } q, k, v; } sage; diff --git a/cpp/tensorrt_llm/kernels/fmhaDispatcher.cpp b/cpp/tensorrt_llm/kernels/fmhaDispatcher.cpp index bfa2f4e278ea..7efd5efab2d9 100644 --- a/cpp/tensorrt_llm/kernels/fmhaDispatcher.cpp +++ b/cpp/tensorrt_llm/kernels/fmhaDispatcher.cpp @@ -68,6 +68,14 @@ FmhaDispatcher::FmhaDispatcher(MHARunnerFixedParams fixedParams) } else { + // SageAttention sets different QK/V datatypes. Unlike trtllmGen, the fmha_v2 kernelMetaInfo + // carries one input datatype. Collapse pair to the composite type where it's registered. + if (mFixedParams.dataType == DATA_TYPE_INT8 && mFixedParams.dataTypeKv == DATA_TYPE_KV_INT8_E4M3) + { + mFixedParams.dataType = DATA_TYPE_KV_INT8_E4M3; + fixedParams.dataType = DATA_TYPE_KV_INT8_E4M3; + } + TLLM_CHECK_WITH_INFO(mFixedParams.dataType == mFixedParams.dataTypeKv, "KV cache data type %s is not the same as input data type %s.", data_type_to_string(mFixedParams.dataTypeKv).c_str(), data_type_to_string(mFixedParams.dataType).c_str()); @@ -153,6 +161,26 @@ bool FmhaDispatcher::isSupported() } else { + if (mFixedParams.sageBlockSizeQ > 0 || mFixedParams.sageBlockSizeK > 0 || mFixedParams.sageBlockSizeV > 0) + { + // This backend implements SageAttention on SM90 only, with one kernel, for block + // sizes (2, 16, 1) and separate Q/K/V. + int const sm = tensorrt_llm::common::getSMVersion(); + if (sm != kSM_90) + { + TLLM_LOG_WARNING("fmha_v2 implements SageAttention on SM90 only, got sm_%d.", sm); + return false; + } + if (mFixedParams.sageBlockSizeQ != 2 || mFixedParams.sageBlockSizeK != 16 + || mFixedParams.sageBlockSizeV != 1) + { + TLLM_LOG_WARNING( + "SageAttention on SM90 supports exactly one block-size combination, (q, k, v) = (2, 16, 1), " + "got (%d, %d, %d).", + mFixedParams.sageBlockSizeQ, mFixedParams.sageBlockSizeK, mFixedParams.sageBlockSizeV); + return false; + } + } foundKernels = mFMHARunner->isFmhaSupported(); } if (!foundKernels) diff --git a/cpp/tensorrt_llm/kernels/multiHeadAttentionCommon.h b/cpp/tensorrt_llm/kernels/multiHeadAttentionCommon.h index bfd98db66c19..cec09f6bb768 100644 --- a/cpp/tensorrt_llm/kernels/multiHeadAttentionCommon.h +++ b/cpp/tensorrt_llm/kernels/multiHeadAttentionCommon.h @@ -111,6 +111,8 @@ static inline size_t get_size_in_bytes(size_t n, Data_type dtype) case DATA_TYPE_E2M1: TLLM_CHECK_WITH_INFO(n % 2 == 0, "Not supported."); return n / 2; case DATA_TYPE_E4M3: return n; case DATA_TYPE_E5M2: return n; + // SageAttention is int8 Q/K with an e4m3 PV; both halves are one byte. + case DATA_TYPE_KV_INT8_E4M3: return n; default: TLLM_CHECK_WITH_INFO(false, "FMHA Data Type is not supported."); return 0; } } diff --git a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp index 518c80671132..a41158dd4661 100644 --- a/cpp/tensorrt_llm/nanobind/thop/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/thop/bindings.cpp @@ -255,16 +255,17 @@ void initBindings(nb::module_& m) nb::arg("quant_q_buffer") = std::nullopt, nb::arg("flash_mla_tile_scheduler_metadata") = std::nullopt, nb::arg("flash_mla_num_splits") = std::nullopt, nb::arg("sage_attn_num_elts_per_blk_q") = 0, nb::arg("sage_attn_num_elts_per_blk_k") = 0, nb::arg("sage_attn_num_elts_per_blk_v") = 0, - nb::arg("sage_attn_qk_int8") = false, nb::arg("num_contexts") = 0, nb::arg("num_ctx_tokens") = 0, - nb::arg("trtllm_gen_jit_warmup") = false, nb::arg("aux_kv_cache_pool_ptr") = std::nullopt, - nb::arg("is_cross") = false, nb::arg("cross_kv") = std::nullopt, - nb::arg("relative_attention_bias") = std::nullopt, nb::arg("relative_attention_max_distance") = 0, - nb::arg("spec_decoding_target_max_draft_tokens") = std::nullopt, nb::arg("quant_scale_qkv") = std::nullopt, - nb::arg("dsv4_inv_rope_cos_sin_cache") = std::nullopt, nb::arg("enable_dsv4_epilogue_fusion") = false, - nb::arg("force_prepare_spec_dec_tree_mask") = false, nb::arg("max_num_sequences") = std::nullopt, - nb::arg("kv_norm_weight") = std::nullopt, nb::arg("kv_norm_eps") = 1e-6, - nb::arg("skip_correction_threshold") = 0.0, nb::arg("uses_spcompress") = std::nullopt, - "Multi-head attention operation", nb::call_guard()); + nb::arg("sage_attn_qk_int8") = false, nb::arg("sage_attn_smooth_k") = false, nb::arg("num_contexts") = 0, + nb::arg("num_ctx_tokens") = 0, nb::arg("trtllm_gen_jit_warmup") = false, + nb::arg("aux_kv_cache_pool_ptr") = std::nullopt, nb::arg("is_cross") = false, + nb::arg("cross_kv") = std::nullopt, nb::arg("relative_attention_bias") = std::nullopt, + nb::arg("relative_attention_max_distance") = 0, nb::arg("spec_decoding_target_max_draft_tokens") = std::nullopt, + nb::arg("quant_scale_qkv") = std::nullopt, nb::arg("dsv4_inv_rope_cos_sin_cache") = std::nullopt, + nb::arg("enable_dsv4_epilogue_fusion") = false, nb::arg("force_prepare_spec_dec_tree_mask") = false, + nb::arg("max_num_sequences") = std::nullopt, nb::arg("kv_norm_weight") = std::nullopt, + nb::arg("kv_norm_eps") = 1e-6, nb::arg("skip_correction_threshold") = 0.0, + nb::arg("uses_spcompress") = std::nullopt, "Multi-head attention operation", + nb::call_guard()); m.def( "get_helix_workspace_size_per_rank", diff --git a/cpp/tensorrt_llm/thop/attentionOp.cpp b/cpp/tensorrt_llm/thop/attentionOp.cpp index ffd09ef927a8..4d1dd59cf305 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.cpp +++ b/cpp/tensorrt_llm/thop/attentionOp.cpp @@ -1217,14 +1217,14 @@ void attention(torch::Tensor q, std::optional k, std::optional mla_bmm2_scale, std::optional quant_q_buffer, std::optional flash_mla_tile_scheduler_metadata, std::optional flash_mla_num_splits, int64_t sage_attn_num_elts_per_blk_q, int64_t sage_attn_num_elts_per_blk_k, int64_t sage_attn_num_elts_per_blk_v, - bool sage_attn_qk_int8, int64_t num_contexts, int64_t num_ctx_tokens, bool trtllm_gen_jit_warmup, - std::optional aux_kv_cache_pool_ptr, bool const is_cross, std::optional cross_kv, - std::optional relative_attention_bias, int64_t relative_attention_max_distance, - std::optional spec_decoding_target_max_draft_tokens, std::optional quant_scale_qkv, - std::optional dsv4_inv_rope_cos_sin_cache, bool enable_dsv4_epilogue_fusion, - bool const force_prepare_spec_dec_tree_mask, std::optional const max_num_sequences, - std::optional kv_norm_weight, double kv_norm_eps, double skip_correction_threshold, - std::optional uses_spcompress) + bool sage_attn_qk_int8, bool sage_attn_smooth_k, int64_t num_contexts, int64_t num_ctx_tokens, + bool trtllm_gen_jit_warmup, std::optional aux_kv_cache_pool_ptr, bool const is_cross, + std::optional cross_kv, std::optional relative_attention_bias, + int64_t relative_attention_max_distance, std::optional spec_decoding_target_max_draft_tokens, + std::optional quant_scale_qkv, std::optional dsv4_inv_rope_cos_sin_cache, + bool enable_dsv4_epilogue_fusion, bool const force_prepare_spec_dec_tree_mask, + std::optional const max_num_sequences, std::optional kv_norm_weight, double kv_norm_eps, + double skip_correction_threshold, std::optional uses_spcompress) { TLLM_LOG_TRACE("Attention op starts at layer %d", local_layer_idx); // Use these tensors to infer if the attention is using KV cache @@ -1358,6 +1358,7 @@ void attention(torch::Tensor q, std::optional k, std::optionalmSageAttnNumEltsPerBlkK = static_cast(sage_attn_num_elts_per_blk_k); op->mSageAttnNumEltsPerBlkV = static_cast(sage_attn_num_elts_per_blk_v); op->mSageAttnQkInt8 = sage_attn_qk_int8; + op->mSageAttnSmoothK = sage_attn_smooth_k; op->mFP8AttenOutput = is_fp8_out; op->mPagedContextFMHA = use_paged_context_fmha; op->mCrossAttention = is_cross; diff --git a/cpp/tensorrt_llm/thop/attentionOp.h b/cpp/tensorrt_llm/thop/attentionOp.h index 0897b0eff367..1c12001c226a 100644 --- a/cpp/tensorrt_llm/thop/attentionOp.h +++ b/cpp/tensorrt_llm/thop/attentionOp.h @@ -88,9 +88,9 @@ void attention(torch::Tensor q, std::optional k, std::optional flash_mla_tile_scheduler_metadata = std::nullopt, std::optional flash_mla_num_splits = std::nullopt, int64_t sage_attn_num_elts_per_blk_q = 0, int64_t sage_attn_num_elts_per_blk_k = 0, int64_t sage_attn_num_elts_per_blk_v = 0, bool sage_attn_qk_int8 = false, - int64_t num_contexts = 0, int64_t num_ctx_tokens = 0, bool trtllm_gen_jit_warmup = false, - std::optional aux_kv_cache_pool_ptr = std::nullopt, bool const is_cross = false, - std::optional cross_kv = std::nullopt, + bool sage_attn_smooth_k = false, int64_t num_contexts = 0, int64_t num_ctx_tokens = 0, + bool trtllm_gen_jit_warmup = false, std::optional aux_kv_cache_pool_ptr = std::nullopt, + bool const is_cross = false, std::optional cross_kv = std::nullopt, std::optional relative_attention_bias = std::nullopt, int64_t relative_attention_max_distance = 0, std::optional spec_decoding_target_max_draft_tokens = std::nullopt, std::optional quant_scale_qkv = std::nullopt, diff --git a/docs/source/features/visualgen-quantized-attention.md b/docs/source/features/visualgen-quantized-attention.md index 4f3141f70434..bc2da9b9c713 100644 --- a/docs/source/features/visualgen-quantized-attention.md +++ b/docs/source/features/visualgen-quantized-attention.md @@ -9,9 +9,9 @@ This feature is in **beta** stage. APIs, supported models, and optimization opti - [Choosing and Tuning a Recipe](#choosing-and-tuning-a-recipe) - [Configuration Surface](#configuration-surface) - [QK16PV8 Attention Kernels in the CUTEDSL Backend](#qk16pv8-attention-kernels-in-the-cutedsl-backend) -- [SageAttention (TRTLLM)](#sageattention-in-the-trtllm-backend) +- [SageAttention in the `TRTLLM` backend](#sageattention-in-the-trtllm-backend) - [FP8 and MXFP8 in the cuDNN Backend](#fp8-and-mxfp8-in-the-cudnn-backend) -- [MXFP8 / NVFP4 (CUTEDSL / FlashInfer)](#mxfp8-and-nvfp4-in-the-cutedsl-and-flashinfer-backends) +- [MXFP8 and NVFP4 in the `CUTEDSL` and `FlashInfer` backends](#mxfp8-and-nvfp4-in-the-cutedsl-and-flashinfer-backends) - [Interaction With Other Features](#interaction-with-other-features) ## Overview @@ -26,8 +26,9 @@ A recipe is the tuple `(qk_dtype, v_dtype, (q_block_size, k_block_size, v_block_ | Backend | `qk_dtype` | `v_dtype` | `(q, k, v)` block sizes | Common name | |---|---|---|---|---| -| `TRTLLM` | `int8` | `fp8` | `(1, 1, 1)`, `(1, 4, 1)`, `(1, 16, 1)` | SageAttention (INT8 QK) | -| `TRTLLM` | `fp8` | `fp8` | `(1, 1, 1)`, `(1, 4, 1)` | SageAttention (FP8 QK) | +| `TRTLLM` | `int8` | `fp8` | `(2, 16, 1)` | SageAttention (INT8 QK) for Hopper | +| `TRTLLM` | `int8` | `fp8` | `(1, 1, 1)`, `(1, 4, 1)`, `(1, 16, 1)` | SageAttention (INT8 QK) for Blackwell | +| `TRTLLM` | `fp8` | `fp8` | `(1, 1, 1)`, `(1, 4, 1)` | SageAttention (FP8 QK) for Blackwell | | `CUDNN` | `fp8` | `fp8` | `(0, 0, 0)` | cuDNN FP8 | | `CUDNN` | `mxfp8` | `mxfp8` | `(0, 0, 0)` | cuDNN MXFP8 | | `CUTEDSL` | `bf16` | `fp8` | `(0, 0, 0)` | QK16PV8 | @@ -46,7 +47,7 @@ Video quality is generally more sensitive to BMM1 accuracy than BMM2 accuracy, s - QK16PV8 keeps Q/K in BF16 and only quantizes V, making it the most conservative quantized-attention recipe. - On B200/GB200, SageAttention with INT8 Q/K typically matches QK16PV8 quality while delivering higher end-to-end throughput. - On B300/GB300, start with MXFP8 when optimizing the quality-throughput balance. SageAttention with FP8 Q/K remains an alternative when the `TRTLLM` backend is preferred for the surrounding workload. -- For SageAttention with INT8 Q/K, the default `(1, 16, 1)` block-size recipe works well for most cases. Use `(1, 4, 1)` when video quality is not satisfactory. +- For SageAttention with INT8 Q/K, the default `(2, 16, 1)` or `(1, 16, 1)` block-size recipes works well for most cases. Use `(1, 4, 1)` when video quality is not satisfactory. - The `CUDNN` backend requires Q/K and V to use the same format. Use `fp8` for per-tensor scaling or `mxfp8` for block scaling. - For `CUTEDSL` MXFP8 or NVFP4 recipes, `v_block_size: 1` uses a separate V scale per head and channel, while `v_block_size: 0` uses one tensor-wide V scale. Try the per-channel variant when the tensor-wide scale loses quality. @@ -109,8 +110,9 @@ attention_config: **Requirements and behavior.** -- SageAttention is supported on B200/GB200 and B300/GB300 GPUs. -- On B200/GB200, use the recommended `qk_dtype: "int8"` recipe. +- SageAttention is supported on H100/H200/H800/H20, B200/GB200, and B300/GB300 GPUs. +- On H100/H200/H800/H20, only one recipe with fixed `qk_dtype: "int8"` is supported. +- On B200/GB200, it is recommended to use the `qk_dtype: "int8"` recipes. - On B300/GB300, use `qk_dtype: "fp8"` and evaluate output quality because it can be less accurate than the INT8 Q/K recipe on B200/GB200. **Configuration.** @@ -145,6 +147,8 @@ attention_config: v_block_size: 1 ``` +For Hopper SageAttention, set `q_block_size` to `2`. + ## FP8 and MXFP8 in the cuDNN Backend **What it does.** The `CUDNN` backend uses cuDNN fused SDPA. Its quantized recipes use the same format for BMM1 and BMM2: diff --git a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py index a3b9379fce52..2466feef8471 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py @@ -196,6 +196,7 @@ def forward( sage_attn_num_elts_per_blk_k=forward_args.sage_attn_num_elts_per_blk_k, sage_attn_num_elts_per_blk_v=forward_args.sage_attn_num_elts_per_blk_v, sage_attn_qk_int8=forward_args.sage_attn_qk_int8, + sage_attn_smooth_k=forward_args.sage_attn_smooth_k, is_fused_qkv=forward_args.is_fused_qkv, update_kv_cache=forward_args.update_kv_cache, cross_kv=forward_args.cross_kv, diff --git a/tensorrt_llm/_torch/attention/backends/interface.py b/tensorrt_llm/_torch/attention/backends/interface.py index f638265d81ce..50f08a55569e 100644 --- a/tensorrt_llm/_torch/attention/backends/interface.py +++ b/tensorrt_llm/_torch/attention/backends/interface.py @@ -958,6 +958,7 @@ class AttentionForwardArgs: sage_attn_num_elts_per_blk_k: int = 0 sage_attn_num_elts_per_blk_v: int = 0 sage_attn_qk_int8: bool = False + sage_attn_smooth_k: bool = False # Packed QKV for non-MLA attention. MLA always passes a separate query. is_fused_qkv: bool = False diff --git a/tensorrt_llm/_torch/visual_gen/attention_backend/trtllm.py b/tensorrt_llm/_torch/visual_gen/attention_backend/trtllm.py index 53d29a2adf35..6e5371487c84 100644 --- a/tensorrt_llm/_torch/visual_gen/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/visual_gen/attention_backend/trtllm.py @@ -398,6 +398,7 @@ def forward( "sage_attn_num_elts_per_blk_k": quant_cfg.k_block_size, "sage_attn_num_elts_per_blk_v": quant_cfg.v_block_size, "sage_attn_qk_int8": quant_cfg.qk_dtype == "int8", + "sage_attn_smooth_k": quant_cfg.smooth_k, } else: if k is None and v is None: diff --git a/tensorrt_llm/visual_gen/args.py b/tensorrt_llm/visual_gen/args.py index b08d322f95a0..dc9c0e30184f 100644 --- a/tensorrt_llm/visual_gen/args.py +++ b/tensorrt_llm/visual_gen/args.py @@ -95,6 +95,14 @@ class QuantAttentionConfig(StrictBaseModel): "V quantization block size on the hidden dimension; 0 uses one tensor-wide V scale." ), ) + smooth_k: bool = Field( + False, + status="prototype", + description=( + "Subtract the per-channel K mean before quantizing K, which narrows the range the " + "quantized type has to cover. SageAttention only." + ), + ) # Discriminated union of sparse attention configs. @@ -133,15 +141,24 @@ class AttentionConfig(StrictBaseModel): @model_validator(mode="after") def _validate_quant_attention_config(self) -> "AttentionConfig": - # Recipe tuple: (qk_dtype, v_dtype, (q_block, k_block, v_block)). + # SAGE supports different recipes for different architectures. SAGE_RECIPES = { - ("int8", "fp8", (1, 1, 1)), - ("int8", "fp8", (1, 4, 1)), - ("int8", "fp8", (1, 16, 1)), - ("fp8", "fp8", (1, 1, 1)), - ("fp8", "fp8", (1, 4, 1)), + 90: { + ("int8", "fp8", (2, 16, 1)), + }, + 100: { + ("int8", "fp8", (1, 1, 1)), + ("int8", "fp8", (1, 4, 1)), + ("int8", "fp8", (1, 16, 1)), + ("fp8", "fp8", (1, 1, 1)), + ("fp8", "fp8", (1, 4, 1)), + }, + 103: { + ("fp8", "fp8", (1, 1, 1)), + ("fp8", "fp8", (1, 4, 1)), + }, } - # cuDNN fused SDPA quantizes both GEMMs with the same element format. + # Other recipes verify the hardware at corresponding backend implementations. CUDNN_RECIPES = { ("fp8", "fp8", (0, 0, 0)), ("mxfp8", "mxfp8", (0, 0, 0)), @@ -169,20 +186,18 @@ def _validate_quant_attention_config(self) -> "AttentionConfig": (q_config.q_block_size, q_config.k_block_size, q_config.v_block_size), ) if self.backend == "TRTLLM": - if recipe in SAGE_RECIPES: - # int8 Q/K SAGE has a compiled cubin only on SM100. - if q_config.qk_dtype == "int8" and get_sm_version() != 100: - raise ValueError( - f"int8 Q/K SAGE quantized attention (backend='TRTLLM', " - f"qk_dtype='int8', v_dtype='{q_config.v_dtype}') only supports sm_100." - ) - else: + recipes = SAGE_RECIPES.get(get_sm_version(), set()) + if recipe not in recipes: raise ValueError( f"Unsupported quant_attention_config={self.quant_attention_config!r} " - f"for backend='TRTLLM'. Supported SAGE recipes " - f"(qk_dtype, v_dtype, (q_block, k_block, v_block)): " - f"{sorted(SAGE_RECIPES)}." + f"for backend='TRTLLM'. Supported SAGE recipes on this device: " + f"{sorted(recipes)}." ) + elif q_config.smooth_k: + raise ValueError( + f"smooth_k is a SageAttention option and requires backend='TRTLLM', got " + f"backend='{self.backend}'." + ) elif self.backend == "CUTEDSL": if recipe not in CUTEDSL_RECIPES: raise ValueError( diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 31659d64e83c..27626bc1bc63 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -74,6 +74,7 @@ l0_h100: - unittest/_torch/visual_gen/test_attention_flashinfer.py - unittest/_torch/thop/parallel_hw_agnostic - unittest/_torch/visual_gen/kernels/parallel_hw_agnostic + - unittest/_torch/visual_gen/test_attention_trtllm_sage.py - unittest/_torch/thop/serial # Only key models in H100: llama/gemma/gpt-oss - unittest/_torch/modeling -k "modeling_llama" diff --git a/tests/unittest/_torch/visual_gen/sparse_attention/sol/test_sol_attention.py b/tests/unittest/_torch/visual_gen/sparse_attention/sol/test_sol_attention.py index d0ef18211375..35a8034bc8c5 100644 --- a/tests/unittest/_torch/visual_gen/sparse_attention/sol/test_sol_attention.py +++ b/tests/unittest/_torch/visual_gen/sparse_attention/sol/test_sol_attention.py @@ -26,7 +26,7 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import Mock, patch import pytest import torch @@ -904,13 +904,16 @@ def test_sol_and_attention_quantization_are_mutually_exclusive( ) -> None: """A recipe the backend accepts on its own is still rejected next to SOL.""" - AttentionConfig(backend=backend, quant_attention_config=QuantAttentionConfig(**quant_recipe)) - with pytest.raises(ValidationError, match="SOL and quant_attention_config"): + with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=100): AttentionConfig( - backend=backend, - quant_attention_config=QuantAttentionConfig(**quant_recipe), - sparse_attention_config=SolAttentionConfig(), + backend=backend, quant_attention_config=QuantAttentionConfig(**quant_recipe) ) + with pytest.raises(ValidationError, match="SOL and quant_attention_config"): + AttentionConfig( + backend=backend, + quant_attention_config=QuantAttentionConfig(**quant_recipe), + sparse_attention_config=SolAttentionConfig(), + ) @_CPU_ONLY diff --git a/tests/unittest/_torch/visual_gen/test_attention_integration.py b/tests/unittest/_torch/visual_gen/test_attention_integration.py index 4c7c87b94993..64fcb0d14f14 100644 --- a/tests/unittest/_torch/visual_gen/test_attention_integration.py +++ b/tests/unittest/_torch/visual_gen/test_attention_integration.py @@ -14,10 +14,6 @@ import torch.nn.functional as F from tensorrt_llm._torch.modules.rms_norm import RMSNorm - -# ============================================================================ -# Flash Attention 4 availability -# ============================================================================ from tensorrt_llm._torch.visual_gen.attention_backend.cudnn import CuDNNAttention from tensorrt_llm._torch.visual_gen.attention_backend.cute_dsl import _cute_dsl_import_error from tensorrt_llm._torch.visual_gen.attention_backend.flash_attn4 import _flash_attn_fwd as _fa4_fwd @@ -26,7 +22,6 @@ RingAttention, UlyssesAttention, ) -from tensorrt_llm._torch.visual_gen.attention_backend.trtllm import TrtllmAttention from tensorrt_llm._torch.visual_gen.attention_backend.vanilla import VanillaAttention from tensorrt_llm._torch.visual_gen.config import ( DiffusionModelConfig, @@ -36,6 +31,7 @@ # Import new integrated versions from tensorrt_llm._torch.visual_gen.modules.attention import Attention, QKVMode, apply_rotary_emb +from tensorrt_llm._utils import get_sm_version from tensorrt_llm.visual_gen.args import ( AttentionConfig, QuantAttentionConfig, @@ -185,9 +181,7 @@ def _require_attention_backend( except ImportError as e: pytest.skip(f"cuDNN detected hardware/library incompatibility: {e}") if attn_backend == "CUTEDSL": - compute_capability = torch.cuda.get_device_capability() - gpu_arch = f"sm_{compute_capability[0]}{compute_capability[1]}a" - if gpu_arch not in ("sm_100a", "sm_103a"): + if get_sm_version() not in (100, 103): pytest.skip("CUTEDSL attention test requires a supported Blackwell-class GPU") @@ -313,43 +307,6 @@ def test_attn2d_with_sequence_parallel_disabled_allowed(self): assert isinstance(attn.attn, VanillaAttention) -def _build_sage_routed_attention(qkv_mode: QKVMode): - """Build a TRTLLM-SAGE-configured Attention for backend-routing checks.""" - quant_cfg = QuantAttentionConfig( - qk_dtype="int8", q_block_size=1, k_block_size=16, v_block_size=1 - ) - config = create_model_config( - hidden_size=512, - num_heads=4, - head_dim=128, - attn_backend="TRTLLM", - quant_attention_config=quant_cfg, - skip_create_weights_in_init=True, - ) - attn = Attention( - hidden_size=512, - num_attention_heads=4, - head_dim=128, - qkv_mode=qkv_mode, - config=config, - ) - return attn, quant_cfg - - -class TestSageAttentionBackendRouting: - def test_self_attention_uses_trtllm_sage_backend(self): - attn, quant_cfg = _build_sage_routed_attention(QKVMode.FUSE_QKV) - assert attn.attn_backend == "TRTLLM" - assert isinstance(attn.attn, TrtllmAttention) - assert attn.attn.quant_attention_config == quant_cfg - assert not attn.attn.support_fused_qkv() - - def test_cross_attention_with_sage_config_falls_back_to_vanilla(self): - attn, _ = _build_sage_routed_attention(QKVMode.SEPARATE_QKV) - assert attn.attn_backend == "VANILLA" - assert isinstance(attn.attn, VanillaAttention) - - # ============================================================================ # Test functions # ============================================================================ @@ -467,10 +424,10 @@ def test_sage_attention_self_attention(qk_dtype: str, batch_size: int, seq_len: 3. Outputs are finite (no NaN/Inf) 4. Approximate agreement with naive (cosine similarity > 0.99) """ - compute_capability = torch.cuda.get_device_capability() - gpu_arch = f"sm_{compute_capability[0]}{compute_capability[1]}a" - if qk_dtype == "int8" and gpu_arch not in ["sm_100a"]: - pytest.skip("Int8 kernels are only available for SM100 devices.") + # This test configures the contiguous-block recipe (q_block_size=1), which requires SM100. The + # SM90 recipe is covered by test_attention_trtllm_sage.py::test_attention_trtllm_sage_hopper. + if get_sm_version() != 100: + pytest.skip("The contiguous-block SageAttention recipe is only available on SM100 devices.") print("\n" + "=" * 60) print(f"Testing SageAttention (qk_dtype={qk_dtype}, B={batch_size}, S={seq_len})") print("=" * 60) diff --git a/tests/unittest/_torch/visual_gen/test_attention_trtllm_sage.py b/tests/unittest/_torch/visual_gen/test_attention_trtllm_sage.py index 91c337c1960f..10853ed9ffe0 100644 --- a/tests/unittest/_torch/visual_gen/test_attention_trtllm_sage.py +++ b/tests/unittest/_torch/visual_gen/test_attention_trtllm_sage.py @@ -9,16 +9,10 @@ from tensorrt_llm._torch.attention.backends import interface as attention_backend_interface from tensorrt_llm._torch.attention.backends import utils as attention_backend_utils +from tensorrt_llm._utils import get_sm_version from tensorrt_llm.llmapi.llm_args import SkipSoftmaxAttentionConfig -def _cuda_cc(): - if torch.cuda.is_available(): - return torch.cuda.get_device_capability() - else: - return -1, -1 - - def _repeat_kv(hidden_states: torch.Tensor, gqa_groups: int) -> torch.Tensor: bsz, n_kv_heads, seqlen, head_dim = hidden_states.shape if gqa_groups == 1: @@ -60,7 +54,9 @@ def _test_attention_trtllm_sage( amp_mul_v: Optional[float] = None, skip_softmax: bool = False, sage_attn_qk_int8: bool = False, + sage_attn_num_elts_per_blk_q: int = 1, sage_attn_num_elts_per_blk_k: Optional[int] = None, + sage_attn_smooth_k: bool = False, out_dtype: torch.dtype = torch.bfloat16, ) -> Tuple[torch.Tensor, torch.Tensor, float, float, float]: torch.manual_seed(1234) @@ -121,15 +117,17 @@ def _test_attention_trtllm_sage( "attention_mask": mask_type, } - # SageAttention separate-QKV requires these block sizes. + # Which block sizes are legal depends on the GPU: SM100 scales contiguous token blocks, SM90 + # only accepts (q, k, v) = (2, 16, 1). if sage_attn_num_elts_per_blk_k is None: sage_attn_num_elts_per_blk_k = 16 if sage_attn_qk_int8 else 1 attn_kwargs.update( { - "sage_attn_num_elts_per_blk_q": 1, + "sage_attn_num_elts_per_blk_q": sage_attn_num_elts_per_blk_q, "sage_attn_num_elts_per_blk_k": sage_attn_num_elts_per_blk_k, "sage_attn_num_elts_per_blk_v": 1, "sage_attn_qk_int8": sage_attn_qk_int8, + "sage_attn_smooth_k": sage_attn_smooth_k, } ) @@ -164,21 +162,20 @@ def _test_attention_trtllm_sage( @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for TRTLLM attention.") -@pytest.mark.skipif( - _cuda_cc()[0] != 10, reason="TRTLLM SageAttention test requires CUDA major version 10." -) @pytest.mark.parametrize("batch_size", [1, 2, 4]) @pytest.mark.parametrize("gqa_groups", [1, 2, 4]) @pytest.mark.parametrize("num_heads", [4, 12, 16]) -@pytest.mark.parametrize("seq_len", [128, 256, 1024, 8192]) +@pytest.mark.parametrize("seq_len", [128, 256, 1000, 8192]) @pytest.mark.parametrize("skip_softmax", [False, True]) +@pytest.mark.parametrize("sage_attn_smooth_k", [False, True]) @pytest.mark.parametrize( - "out_dtype,sage_attn_qk_int8,sage_attn_num_elts_per_blk_k,atol,rtol", + "sage_attn_qk_int8,sage_attn_num_elts_per_blk_q,sage_attn_num_elts_per_blk_k,atol,rtol," + "supported_sms", [ - (torch.bfloat16, True, 4, 1e-1, 4e-2), - (torch.bfloat16, True, 16, 5e-1, 5e-1), - (torch.bfloat16, True, 4, 3e-1, 2e-1), - (torch.bfloat16, False, 1, 3e-1, 2e-1), + pytest.param(True, 2, 16, 5e-1, 5e-1, (90,), id="sm90-int8-q2k16"), + pytest.param(True, 1, 4, 1e-1, 4e-2, (100,), id="sm100-int8-q1k4"), + pytest.param(True, 1, 16, 5e-1, 5e-1, (100,), id="sm100-int8-q1k16"), + pytest.param(False, 1, 1, 3e-1, 2e-1, (100, 103), id="sm10x-fp8-q1k1"), ], ) def test_attention_trtllm_sage( @@ -186,15 +183,19 @@ def test_attention_trtllm_sage( num_heads: int, gqa_groups: int, batch_size: int, - out_dtype: torch.dtype, skip_softmax: bool, sage_attn_qk_int8: bool, + sage_attn_smooth_k: bool, + sage_attn_num_elts_per_blk_q: int, sage_attn_num_elts_per_blk_k: int, atol: float, rtol: float, + supported_sms: Tuple[int, ...], ): - if sage_attn_qk_int8 and _cuda_cc()[1] == 3: - pytest.skip("SM103 does not have Int8 Tensor Cores.") + # Skip-softmax is unavailable on SM90. + sm_version = get_sm_version() + if sm_version not in supported_sms or (skip_softmax and sm_version == 90): + pytest.skip(f"Configuration is unsupported on SM{sm_version}.") out_tllm, out_native, max_abs, mean_abs, cos_sim = _test_attention_trtllm_sage( num_heads=num_heads, @@ -205,8 +206,10 @@ def test_attention_trtllm_sage( amp_mul=3.2, skip_softmax=skip_softmax, sage_attn_qk_int8=sage_attn_qk_int8, + sage_attn_smooth_k=sage_attn_smooth_k, + sage_attn_num_elts_per_blk_q=sage_attn_num_elts_per_blk_q, sage_attn_num_elts_per_blk_k=sage_attn_num_elts_per_blk_k, - out_dtype=out_dtype, + out_dtype=torch.bfloat16, ) assert out_tllm.shape == out_native.shape, "Shape mismatch" diff --git a/tests/unittest/_torch/visual_gen/test_visual_gen_args.py b/tests/unittest/_torch/visual_gen/test_visual_gen_args.py index 006cb38fd73e..e90434e29bce 100644 --- a/tests/unittest/_torch/visual_gen/test_visual_gen_args.py +++ b/tests/unittest/_torch/visual_gen/test_visual_gen_args.py @@ -95,13 +95,14 @@ def test_quant_config_rejected_on_unsupported_backend(self): ) def test_quant_config_rejected_when_unsupported(self): - with pytest.raises(ValidationError, match="Unsupported quant_attention_config"): - AttentionConfig( - backend="TRTLLM", - quant_attention_config=QuantAttentionConfig( - qk_dtype="int8", q_block_size=1, k_block_size=127, v_block_size=1 - ), - ) + with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=100): + with pytest.raises(ValidationError, match="Unsupported quant_attention_config"): + AttentionConfig( + backend="TRTLLM", + quant_attention_config=QuantAttentionConfig( + qk_dtype="int8", q_block_size=1, k_block_size=127, v_block_size=1 + ), + ) @pytest.mark.parametrize( ("backend", "quant_config"), @@ -122,14 +123,15 @@ def test_quant_config_rejected_when_unsupported(self): ], ) def test_vsa_and_quantization_are_mutually_exclusive(self, backend, quant_config): - with pytest.raises( - ValidationError, match="VSA and quant_attention_config are mutually exclusive" - ): - AttentionConfig( - backend=backend, - quant_attention_config=quant_config, - sparse_attention_config=VideoSparseAttentionConfig(vsa_sparsity=0.9), - ) + with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=100): + with pytest.raises( + ValidationError, match="VSA and quant_attention_config are mutually exclusive" + ): + AttentionConfig( + backend=backend, + quant_attention_config=quant_config, + sparse_attention_config=VideoSparseAttentionConfig(vsa_sparsity=0.9), + ) def test_skip_softmax_and_sage_quantization_can_be_combined(self): # int8 Q/K SAGE has a compiled cubin only on SM100; pin the SM so this @@ -150,19 +152,22 @@ def test_skip_softmax_and_sage_quantization_can_be_combined(self): assert attention.sparse_attention_config.algorithm == "skip_softmax" @pytest.mark.parametrize( - ("qk_dtype", "q_block_size", "k_block_size", "v_block_size"), + ("sm_ver", "qk_dtype", "q_block_size", "k_block_size", "v_block_size"), [ - ("int8", 1, 1, 1), - ("int8", 1, 4, 1), - ("int8", 1, 16, 1), - ("fp8", 1, 1, 1), - ("fp8", 1, 4, 1), + (90, "int8", 2, 16, 1), + (100, "int8", 1, 1, 1), + (100, "int8", 1, 4, 1), + (100, "int8", 1, 16, 1), + (100, "fp8", 1, 1, 1), + (100, "fp8", 1, 4, 1), + (103, "fp8", 1, 1, 1), + (103, "fp8", 1, 4, 1), ], ) - def test_supported_quant_config_sage(self, qk_dtype, q_block_size, k_block_size, v_block_size): - # int8 Q/K SAGE has a compiled cubin only on SM100; pin the SM so this - # supported-recipe check is host-independent (CI CPU stages have no GPU). - with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=100): + def test_supported_quant_config_sage( + self, sm_ver, qk_dtype, q_block_size, k_block_size, v_block_size + ): + with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=sm_ver): attention = AttentionConfig( backend="TRTLLM", quant_attention_config=QuantAttentionConfig( @@ -175,18 +180,15 @@ def test_supported_quant_config_sage(self, qk_dtype, q_block_size, k_block_size, assert attention.quant_attention_config is not None - @pytest.mark.parametrize("sm_version", [90, 107, 120]) - def test_int8_sage_rejected_on_non_sm100(self, sm_version): - # int8 Q/K SAGE only has an SM100 cubin; validation must fail fast on - # any other SM instead of silently falling back to unfused MHA. - with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=sm_version): - with pytest.raises(ValidationError, match="only supports sm_100"): - AttentionConfig( - backend="TRTLLM", - quant_attention_config=QuantAttentionConfig( - qk_dtype="int8", q_block_size=1, k_block_size=1, v_block_size=1 - ), - ) + @pytest.mark.parametrize("backend", ["CUTEDSL", "CUDNN", "FLASHINFER", "VANILLA"]) + def test_smooth_k_rejected_on_non_trtllm_backend(self, backend): + with pytest.raises(ValidationError, match="smooth_k is a SageAttention option"): + AttentionConfig( + backend=backend, + quant_attention_config=QuantAttentionConfig( + qk_dtype="bf16", v_dtype="fp8", smooth_k=True + ), + ) def test_supported_quant_config_cute(self): attention = AttentionConfig( @@ -275,11 +277,12 @@ def test_fp8_qk_dtype_rejected_on_flashinfer(self) -> None: ) def test_blockscaled_qk_dtype_rejected_on_trtllm(self): - with pytest.raises(ValidationError, match="Unsupported quant_attention_config"): - AttentionConfig( - backend="TRTLLM", - quant_attention_config=QuantAttentionConfig(qk_dtype="nvfp4"), - ) + with patch("tensorrt_llm.visual_gen.args.get_sm_version", return_value=100): + with pytest.raises(ValidationError, match="Unsupported quant_attention_config"): + AttentionConfig( + backend="TRTLLM", + quant_attention_config=QuantAttentionConfig(qk_dtype="nvfp4"), + ) def test_supported_quant_config_cudnn_fp8(self): attention = AttentionConfig(