From 77e3371f5fd6c2b56a5c35441df05c3ce682c47e Mon Sep 17 00:00:00 2001 From: Emre Albayrak Date: Thu, 17 Sep 2026 20:38:59 +0300 Subject: [PATCH 1/2] [diffusion] attention: add fp8_fa_sm120 FP8 backend for SM120 GPUs Opt-in attention backend around an FP8 (E4M3) flash-attention forward written in CuTe-DSL for SM120 (GeForce RTX 50, RTX PRO 6000 Blackwell). - kernels/ops/attention/fp8_fa_sm120: CuTe-DSL kernel, fused Triton quantization of strided BF16 Q/K/V views, and a reusable plan shared by all DiT layers. - multimodal_gen attention backend with forward and forward_varlen; falls back to cuDNN SDPA for causal, batch > 1, head_dim != 128, non-BF16 and non-SM120 calls. - AttentionBackendEnum.FP8_FA_SM120, CUDA resolver, docs rows and unit tests. MiniMax-H3 Ref2VA on RTX PRO 6000, 50 steps: 6.87 -> 5.61 s/it against cuDNN SDPA. --- .../sglang-diffusion/attention_backends.mdx | 15 + .../ops/attention/fp8_fa_sm120/__init__.py | 11 + .../ops/attention/fp8_fa_sm120/fused_prep.py | 243 +++++++++ .../ops/attention/fp8_fa_sm120/kernel.py | 475 ++++++++++++++++++ .../ops/attention/fp8_fa_sm120/plan.py | 180 +++++++ .../attention/backends/fp8_fa_sm120_attn.py | 216 ++++++++ .../multimodal_gen/runtime/platforms/cuda.py | 31 ++ .../runtime/platforms/interface.py | 1 + .../test/unit/test_fp8_fa_sm120_attn.py | 174 +++++++ 9 files changed, 1346 insertions(+) create mode 100644 python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py create mode 100644 python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py create mode 100644 python/sglang/kernels/ops/attention/fp8_fa_sm120/kernel.py create mode 100644 python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py create mode 100644 python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py diff --git a/docs/docs/sglang-diffusion/attention_backends.mdx b/docs/docs/sglang-diffusion/attention_backends.mdx index d9c6e33d697b..5ba2e5cb1f7d 100644 --- a/docs/docs/sglang-diffusion/attention_backends.mdx +++ b/docs/docs/sglang-diffusion/attention_backends.mdx @@ -139,6 +139,11 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa RAIN_FUSION_ATTN Requires attentions which can be installed with sgl_kernel_npu; available only for NPU. + + fp8_fa_sm120 + FP8_FA_SM120 + FP8 (E4M3) dense attention for SM120 GPUs (GeForce RTX 50, RTX PRO Blackwell), written in CuTe-DSL. Q/K/V are quantized per head on the fly; the output stays BF16. Non-causal, batch 1, head dim 128, BF16 inputs; other calls use torch_cudnn_sdpa. Opt-in only: outputs differ from BF16 attention by about 5% relative RMS per call. Each new sequence length compiles once (about 10 s). + @@ -790,6 +795,16 @@ Hybrid window attention constraints: ✅ NPU-only. Requires attentions from sgl_kernel_npu Configuration via --attention-backend-config. + + fp8_fa_sm120 + Yes + ❌ + ❌ + ❌ + ❌ + ❌ + CUDA SM120 only; falls back to torch_cudnn_sdpa on other devices and for unsupported calls. + diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py new file mode 100644 index 000000000000..9d1bf36c8f45 --- /dev/null +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py @@ -0,0 +1,11 @@ +# SPDX-License-Identifier: Apache-2.0 +"""SM120 FP8 attention: CuTe-DSL kernel, fused Triton prep and the reusable plan.""" + +from sglang.kernels.ops.attention.fp8_fa_sm120.plan import ( + HEAD_DIM, + FP8AttentionPlan, + plan_key, + validate_inputs, +) + +__all__ = ["HEAD_DIM", "FP8AttentionPlan", "plan_key", "validate_inputs"] diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py new file mode 100644 index 000000000000..3e4a36c36ece --- /dev/null +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py @@ -0,0 +1,243 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused FP8 preparation for the SM120 kernel: two Triton passes into the plan buffers. + +Replaces a Torch chain (FP32 copy, amax, divide, cast, +padded copy, V transpose copy) with one per-head amax pass and one pack pass that +writes the kernel's own buffers directly: + + Q/K [H, padded_S, 128] same token/column order as the input + V [H, 128, padded_S] transposed here, so the kernel stays unchanged + +Scales stay per head and the quantization recipe is the same (amax/448, div.rn, +clamp, E4M3), so the kernel's numerics do not move. Padding rows are zeroed once +by the plan and never written here. +""" + +import torch +import triton +import triton.language as tl + + +# Triton's `/` lowers to div.full.f32 (approximate); Torch divides with div.rn. +# The FP8 bytes must match the Torch prepare() path, so the exact division is kept. +@triton.jit +def _divide_rn(numerator, denominator): + return tl.inline_asm_elementwise( + "div.rn.f32 $0, $1, $2;", + "=f,f,f", + [numerator, denominator], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _load_tile( + base, + head, + tokens, + columns, + valid, + head_stride: tl.constexpr, + token_stride: tl.constexpr, +): + offsets = head * head_stride + tokens[:, None] * token_stride + columns[None, :] + return tl.load(base + offsets, mask=valid, other=0.0).to(tl.float32) + + +@triton.jit +def _quantize(values, scale): + scaled = _divide_rn(values, scale) + scaled = tl.maximum(tl.minimum(scaled, 448.0), -448.0) + return scaled.to(tl.float8e4nv) + + +# Pass 1: one program per (head, token block); each reduces its Q/K/V tile to +# three scalars and folds them into maximums[operand, head]. +@triton.jit +def _head_amax_qkv( + q, + k, + v, + maximums, + sequence, + head_stride: tl.constexpr, + token_stride: tl.constexpr, + heads: tl.constexpr, + block_tokens: tl.constexpr, +): + head = tl.program_id(0) + tokens = tl.program_id(1) * block_tokens + tl.arange(0, block_tokens) + columns = tl.arange(0, 128) + valid = tokens[:, None] < sequence + + q_tile = _load_tile(q, head, tokens, columns, valid, head_stride, token_stride) + tl.atomic_max(maximums + 0 * heads + head, tl.max(tl.abs(q_tile))) + + k_tile = _load_tile(k, head, tokens, columns, valid, head_stride, token_stride) + tl.atomic_max(maximums + 1 * heads + head, tl.max(tl.abs(k_tile))) + + v_tile = _load_tile(v, head, tokens, columns, valid, head_stride, token_stride) + tl.atomic_max(maximums + 2 * heads + head, tl.max(tl.abs(v_tile))) + + +# Torch evaluates `tensor / 448.0` as `tensor * (1/448)` with the reciprocal in +# FP32 (CPU-scalar divisor fast path). The same product keeps the scales bit-equal. +INVERSE_FP8_MAX = (torch.ones((), dtype=torch.float32) / 448.0).item() + + +# scales[operand, head] = amax * (1/448), or 1 for an all-zero head. +@triton.jit +def _head_scales( + maximums, + scales, + inverse_fp8_max, + count: tl.constexpr, + block: tl.constexpr, +): + offsets = tl.arange(0, block) + valid = offsets < count + maximum = tl.load(maximums + offsets, mask=valid, other=0.0) + scale = tl.where(maximum > 0.0, maximum * inverse_fp8_max, 1.0) + tl.store(scales + offsets, scale, mask=valid) + + +# Pass 2: same grid as pass 1. Q/K tiles go out in [token, column] order; +# the V tile is transposed in registers/shared and goes out in [column, token]. +@triton.jit +def _pack_qkv( + q, + k, + v, + q_fp8, + k_fp8, + v_fp8, + scales, + sequence, + head_stride: tl.constexpr, + token_stride: tl.constexpr, + padded_queries: tl.constexpr, + padded_keys: tl.constexpr, + heads: tl.constexpr, + block_tokens: tl.constexpr, + permute_keys: tl.constexpr, +): + head = tl.program_id(0) + tokens = tl.program_id(1) * block_tokens + tl.arange(0, block_tokens) + columns = tl.arange(0, 128) + valid = tokens[:, None] < sequence + + q_scale = tl.load(scales + 0 * heads + head) + k_scale = tl.load(scales + 1 * heads + head) + v_scale = tl.load(scales + 2 * heads + head) + + # Q is padded to the 128-row query tile, K/V to the 32-key tile. + hsd_offsets = head * padded_queries * 128 + tokens[:, None] * 128 + columns[None, :] + k_offsets = head * padded_keys * 128 + tokens[:, None] * 128 + columns[None, :] + + q_tile = _load_tile(q, head, tokens, columns, valid, head_stride, token_stride) + tl.store(q_fp8 + hsd_offsets, _quantize(q_tile, q_scale), mask=valid) + + k_tile = _load_tile(k, head, tokens, columns, valid, head_stride, token_stride) + tl.store(k_fp8 + k_offsets, _quantize(k_tile, k_scale), mask=valid) + + # The kernel packs P straight from the score fragment; V^T's stored key + # order must follow it: stored position p holds key + # p//16*16 + 8*((p%4)//2) + 2*((p%16)//4) + p%2 (see key_order_for_positions). + # The permutation is applied on the LOAD side as a row gather: every gathered + # row is still 256 contiguous bytes, and the transposed store stays contiguous. + # Permuting the store instead scatters bytes inside 16 B windows and doubled + # the prepare time. Padding keys are never written. + if permute_keys: + source_tokens = ( + tokens // 16 * 16 + + 8 * ((tokens % 4) // 2) + + 2 * ((tokens % 16) // 4) + + tokens % 2 + ) + else: + source_tokens = tokens + valid_source = source_tokens[:, None] < sequence + + hds_offsets = ( + head * 128 * padded_keys + columns[:, None] * padded_keys + tokens[None, :] + ) + valid_transposed = source_tokens[None, :] < sequence + + v_tile = _load_tile( + v, head, source_tokens, columns, valid_source, head_stride, token_stride + ) + tl.store( + v_fp8 + hds_offsets, tl.trans(_quantize(v_tile, v_scale)), mask=valid_transposed + ) + + +def _attach_fused_state(plan, block_tokens): + q, k, v = plan.inputs + for tensor in (k, v): + if tensor.stride() != q.stride(): + raise ValueError("Q/K/V must share strides for fused preparation") + if q.stride(2) != 1: + raise ValueError("Head dimension must be contiguous for fused preparation") + if block_tokens & (block_tokens - 1) or block_tokens < 16: + raise ValueError( + "block_tokens must be a power of two >= 16 (key permutation groups are 16 wide)" + ) + + heads = q.shape[0] + plan.block_tokens = block_tokens + plan.head_stride = q.stride(0) + plan.token_stride = q.stride(1) + plan.maximums = torch.zeros((3, heads), device=q.device, dtype=torch.float32) + + +def fused_prepare(plan): + """Three launches: per-head amax, scale finalize, direct pack into the plan buffers.""" + q, k, v = plan.inputs + heads = q.shape[0] + padded_queries = plan.q_fp8.shape[1] + padded_keys = plan.k_fp8.shape[1] + grid = (heads, triton.cdiv(plan.sequence, plan.block_tokens)) + + with torch.cuda.device(q.device): + plan.maximums.zero_() + _head_amax_qkv[grid]( + q, + k, + v, + plan.maximums, + plan.sequence, + head_stride=plan.head_stride, + token_stride=plan.token_stride, + heads=heads, + block_tokens=plan.block_tokens, + num_warps=8, + ) + _head_scales[(1,)]( + plan.maximums, + plan.scales, + INVERSE_FP8_MAX, + count=3 * heads, + block=triton.next_power_of_2(3 * heads), + num_warps=1, + ) + _pack_qkv[grid]( + q, + k, + v, + plan.q_fp8, + plan.k_fp8, + plan.v_fp8, + plan.scales, + plan.sequence, + head_stride=plan.head_stride, + token_stride=plan.token_stride, + padded_queries=padded_queries, + padded_keys=padded_keys, + heads=heads, + block_tokens=plan.block_tokens, + permute_keys=plan.key_order is not None, + num_warps=8, + ) + plan._prepared = True diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/kernel.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/kernel.py new file mode 100644 index 000000000000..b318f723baee --- /dev/null +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/kernel.py @@ -0,0 +1,475 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FP8 flash-attention forward for SM120 in CuTe-DSL: E4M3 QK and PV, FP32 softmax. + + Q, K [H, padded_S, 128] E4M3, queries padded to 128 rows, keys to 32 + V^T [H, 128, padded_S] E4M3, keys stored in key_order_for_positions() order + scales [3, H] FP32 per-head descale of Q, K, V + O [H, S, 128] BF16 + LSE [H, S] FP32 + +One CTA covers 128 queries x 32 keys on 128 threads: 4 warps, 32 query rows per warp. +Every thread keeps the online-softmax state (running max m, running sum l) of its four +query rows in registers; the max is reduced across the quad with two butterfly shuffles +per key tile, the sum once after the key loop. + +Per key tile: K and V^T arrive in one of two shared-memory stages and the next tile is +requested right after the barrier, so the copy overlaps the math. QK runs as four E4M3 +mma.sync blocks into FP32 score registers. The probabilities are scaled by 256, cast to +E4M3 and packed in registers from the score fragment into the A fragment of PV; they +never go through shared memory. The A fragment holds the keys of every 16-key group in +a fixed permuted order, which is why V^T is stored in that order. PV accumulates into +FP32 O registers that are rescaled to the new max on every tile and normalized by l at +the end. + +Non-causal only. No dropout, no autograd. Inputs must be finite. +""" + +import math + +import cutlass +import cutlass.cute as cute +import cutlass.utils as utils + + +def _make_gmem_tiled_copy(atom_copy, dtype, copy_bits, minor_size, num_threads): + copy_elems = copy_bits // dtype.width + shape_dim_1 = minor_size // copy_elems + thread_layout = cute.make_layout( + (num_threads // shape_dim_1, shape_dim_1), stride=(shape_dim_1, 1) + ) + value_layout = cute.make_layout((1, copy_elems)) + return cute.make_tiled_copy_tv(atom_copy, thread_layout, value_layout) + + +def _make_smem_layout_fp8(dtype, copy_bits, smem_tiler): + major_size = smem_tiler[1] + row_bytes = major_size * dtype.width // 8 + chunk_bytes = copy_bits // 8 + swizzle_bits = min(int(math.log2(row_bytes // chunk_bytes)), 3) + base_bits = int(math.log2(chunk_bytes)) + shift_bits = int(math.log2(128 // chunk_bytes)) + swizzle = cute.make_swizzle(swizzle_bits, base_bits, shift_bits) + atom = cute.make_layout((8, major_size), stride=(major_size, 1)) + layout = cute.tile_to_shape(atom, smem_tiler, (0, 1, 2)) + return layout, swizzle + + +@cute.kernel +def _fp8_attention_sm120( + softmax_scale: cutlass.Float32, + mQ: cute.Tensor, + mK: cute.Tensor, + mV: cute.Tensor, + mO: cute.Tensor, + mLSE: cute.Tensor, + mScales: cute.Tensor, + sQ_layout: cute.Layout, + sK_layout: cute.Layout, + sV_layout: cute.Layout, + sQ_swizzle: cute.Swizzle, + sK_swizzle: cute.Swizzle, + sV_swizzle: cute.Swizzle, + tiled_copy_Q: cute.TiledCopy, + tiled_copy_K: cute.TiledCopy, + tiled_copy_V: cute.TiledCopy, + tiled_mma: cute.TiledMma, + cta_tiler: cutlass.Constexpr = (128, 32, 128), +): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, bidz = cute.arch.block_idx() + + qk_scale = mScales[0, bidz] * mScales[1, bidz] * softmax_scale + pv_scale = mScales[2, bidz] / 256.0 + sequence = mLSE.shape[0] + + m = cute.make_rmem_tensor((4,), cutlass.Float32) + m.fill(-cutlass.Float32.inf) + + l = cute.make_rmem_tensor((4,), cutlass.Float32) + l.fill(0.0) + + gQ = cute.local_tile( + mQ[None, None, bidz], cta_tiler, (bidx, None, 0), proj=(1, None, 1) + ) + gK = cute.local_tile( + mK[None, None, bidz], cta_tiler, (None, None, 0), proj=(None, 1, 1) + ) + gV = cute.local_tile(mV[None, None, bidz], (128, 32), (0, None)) + gO = cute.local_tile(mO[None, None, bidz], (128, 128), (bidx, 0)) + + @cute.struct + class SharedStorageQKV: + q: cute.struct.Align[ + cute.struct.MemRange[mQ.element_type, cute.cosize(sQ_layout)], 16 + ] + k: cute.struct.Align[ + cute.struct.MemRange[mK.element_type, cute.cosize(sK_layout)], 16 + ] + v: cute.struct.Align[ + cute.struct.MemRange[mV.element_type, cute.cosize(sV_layout)], 16 + ] + + smem = utils.SmemAllocator() + storage = smem.allocate(SharedStorageQKV.size_in_bytes(), byte_alignment=16) + sQ = SharedStorageQKV(storage).q.get_tensor(sQ_layout, swizzle=sQ_swizzle) + sK = SharedStorageQKV(storage).k.get_tensor(sK_layout, swizzle=sK_swizzle) + sV = SharedStorageQKV(storage).v.get_tensor(sV_layout, swizzle=sV_swizzle) + + thr_copy_Q = tiled_copy_Q.get_slice(tidx) + thr_copy_K = tiled_copy_K.get_slice(tidx) + thr_copy_V = tiled_copy_V.get_slice(tidx) + + tQgQ = thr_copy_Q.partition_S(gQ) + tKgK = thr_copy_K.partition_S(gK) + tVgV = thr_copy_V.partition_S(gV) + + tQsQ = thr_copy_Q.partition_D(sQ) + tKsK = thr_copy_K.partition_D(sK) + tVsV = thr_copy_V.partition_D(sV) + + k_tile_count = cute.size(tKgK, mode=[3]) + + thr_mma = tiled_mma.get_slice(tidx) + tCgO = thr_mma.partition_C(gO) + + tCsQ = thr_mma.partition_A(sQ) + tCsK = thr_mma.partition_B(sK) + tCsV = thr_mma.partition_B(sV) + + tCrQ = tiled_mma.make_fragment_A(tCsQ[None, None, None, 0]) + tCrK = tiled_mma.make_fragment_B(tCsK[None, None, None, 0]) + tCrV = tiled_mma.make_fragment_B(tCsV[None, None, None, 0]) + + acc_shape = thr_mma.partition_shape_C((128, 32)) + tCrC = cute.make_rmem_tensor(acc_shape, cutlass.Float32) + + acc_shape_O = thr_mma.partition_shape_C((128, 128)) + tCrO = cute.make_rmem_tensor(acc_shape_O, cutlass.Float32) + tCrO.fill(0.0) + + tCrP = cute.make_rmem_tensor( + cute.make_layout(((4, 2, 2), 2, 1), stride=((1, 4, 8), 16, 0)), + mQ.element_type, + ) + num_k_block_PV = cute.size(tCrP, mode=[2]) + + pack_shape = ((2, 2), 2, (2, 2)) + tCrC_as_pack = cute.make_tensor( + tCrC.iterator, + cute.make_layout(pack_shape, stride=((1, 2), 4, (8, 16))), + ) + tCrP_as_pack = cute.make_tensor( + tCrP.iterator, + cute.make_layout(pack_shape, stride=((1, 4), 16, (2, 8))), + ) + + atom_copy_s2r_Q = cute.make_copy_atom( + cute.nvgpu.warp.LdMatrix8x16x8bOp(transpose=False, num_matrices=4), + mQ.element_type, + ) + atom_copy_s2r_K = cute.make_copy_atom( + cute.nvgpu.warp.LdMatrix8x16x8bOp(transpose=False, num_matrices=4), + mK.element_type, + ) + atom_copy_s2r_V = cute.make_copy_atom( + cute.nvgpu.warp.LdMatrix8x16x8bOp(transpose=False, num_matrices=4), + mV.element_type, + ) + + tiled_copy_s2r_Q = cute.make_tiled_copy_A(atom_copy_s2r_Q, tiled_mma) + tiled_copy_s2r_K = cute.make_tiled_copy_B(atom_copy_s2r_K, tiled_mma) + tiled_copy_s2r_V = cute.make_tiled_copy_B(atom_copy_s2r_V, tiled_mma) + + ldmatrix_Q = tiled_copy_s2r_Q.get_slice(tidx) + ldmatrix_K = tiled_copy_s2r_K.get_slice(tidx) + ldmatrix_V = tiled_copy_s2r_V.get_slice(tidx) + + tCsQ_copy_view = ldmatrix_Q.partition_S(sQ) + tCrQ_copy_view = ldmatrix_Q.retile(tCrQ) + tCsK_copy_view = ldmatrix_K.partition_S(sK) + tCrK_copy_view = ldmatrix_K.retile(tCrK) + tCsV_copy_view = ldmatrix_V.partition_S(sV) + tCrV_copy_view = ldmatrix_V.retile(tCrV) + + num_k_block = cute.size(tCrQ, mode=[2]) + + tCsQ_p = tCsQ_copy_view[None, None, None, 0] + + cute.copy(tiled_copy_Q, tQgQ[None, None, None], tQsQ[None, None, None, 0]) + cute.copy(tiled_copy_K, tKgK[None, None, None, 0], tKsK[None, None, None, 0]) + cute.copy(tiled_copy_V, tVgV[None, None, None, 0], tVsV[None, None, None, 0]) + cute.arch.cp_async_commit_group() + + for k_tile in range(k_tile_count): + stage = k_tile % 2 + next_stage = 1 - stage + + cute.arch.cp_async_wait_group(0) + cute.arch.sync_threads() + + if k_tile + 1 < k_tile_count: + cute.copy( + tiled_copy_K, + tKgK[None, None, None, k_tile + 1], + tKsK[None, None, None, next_stage], + ) + cute.copy( + tiled_copy_V, + tVgV[None, None, None, k_tile + 1], + tVsV[None, None, None, next_stage], + ) + cute.arch.cp_async_commit_group() + + tCsK_p = tCsK_copy_view[None, None, None, stage] + tCsV_p = tCsV_copy_view[None, None, None, stage] + + tCrC.fill(0.0) + + for k_block in cutlass.range(num_k_block, unroll_full=True): + cute.copy( + tiled_copy_s2r_Q, + tCsQ_p[None, None, k_block], + tCrQ_copy_view[None, None, k_block], + ) + cute.copy( + tiled_copy_s2r_K, + tCsK_p[None, None, k_block], + tCrK_copy_view[None, None, k_block], + ) + cute.gemm( + tiled_mma, + tCrC, + tCrQ[None, None, k_block], + tCrK[None, None, k_block], + tCrC, + ) + + for m_tile in cutlass.range_constexpr(2): + for row in cutlass.range_constexpr(2): + state = m_tile * 2 + row + tCrC_row = tCrC[(None, row), m_tile, None] + row_scores = tCrC_row.load() * qk_scale + + if cutlass.const_expr(mK.shape[0] != sequence): + tCrC_row.store(row_scores) + for column_group in cutlass.range_constexpr(4): + for column_pair in cutlass.range_constexpr(2): + key_row = ( + k_tile * 32 + + (tidx % 4) * 2 + + column_group * 8 + + column_pair + ) + if key_row >= sequence: + tCrC_row[ + column_pair, column_group + ] = -cutlass.Float32.inf + row_scores = tCrC_row.load() + + local_max = row_scores.reduce( + cute.ReductionOp.MAX, + -cutlass.Float32.inf, + 0, + ) + + neighbor_max = cute.arch.shuffle_sync_bfly(local_max, offset=1) + pair_max = cute.arch.fmax(local_max, neighbor_max) + + neighbor_max = cute.arch.shuffle_sync_bfly(pair_max, offset=2) + tile_max = cute.arch.fmax(pair_max, neighbor_max) + + m_new = cute.arch.fmax(m[state], tile_max) + + alpha = cute.math.exp2( + (m[state] - m_new) * math.log2(math.e), + fastmath=True, + ) + p = cute.math.exp2( + (row_scores - m_new) * math.log2(math.e), + fastmath=True, + ) + + p_partial_sums = cute.make_rmem_tensor((4,), cutlass.Float32) + for column_group in cutlass.range_constexpr(4): + p_partial_sums[column_group] = ( + p[0, column_group] + p[1, column_group] + ) + + for level in cutlass.range_constexpr(2): + for sum_index in cutlass.range_constexpr(2 >> level): + p_partial_sums[sum_index] = ( + p_partial_sums[2 * sum_index] + + p_partial_sums[2 * sum_index + 1] + ) + tile_sum = p_partial_sums[0] + + l[state] = alpha * l[state] + tile_sum + m[state] = m_new + + tCrC_row.store(p) + + tCrO_row = tCrO[(None, row), m_tile, None] + tCrO_row.store(tCrO_row.load() * alpha) + + tCrP_as_pack.store((tCrC_as_pack.load() * 256.0).to(mQ.element_type)) + + for k_block in cutlass.range(num_k_block_PV, unroll_full=True): + cute.copy( + tiled_copy_s2r_V, + tCsV_p[None, None, k_block], + tCrV_copy_view[None, None, k_block], + ) + cute.gemm( + tiled_mma, + tCrO, + tCrP[None, None, k_block], + tCrV[None, None, k_block], + tCrO, + ) + + warp_id = tidx // 32 + lane_id = tidx % 32 + quad_id = lane_id // 4 + + for m_tile in cutlass.range_constexpr(2): + for row in cutlass.range_constexpr(2): + state = m_tile * 2 + row + row_sum = l[state] + + neighbor_sum = cute.arch.shuffle_sync_bfly(row_sum, offset=1) + row_sum += neighbor_sum + + neighbor_sum = cute.arch.shuffle_sync_bfly(row_sum, offset=2) + row_sum += neighbor_sum + + l[state] = row_sum + lse = m[state] + cute.math.log(row_sum) + + query_row = bidx * 128 + m_tile * 64 + warp_id * 16 + quad_id + row * 8 + + valid_query = True + if cutlass.const_expr(sequence % 128 != 0): + valid_query = query_row < sequence + + if valid_query and lane_id % 4 == 0: + mLSE[query_row, bidz] = lse + + tCrO_row = tCrO[(None, row), m_tile, None] + inverse_row_sum = pv_scale / row_sum + tCrO_row.store(tCrO_row.load() * inverse_row_sum) + + if valid_query: + tCrO_row_out = cute.make_fragment_like(tCrO_row, mO.element_type) + tCrO_row_out.store(tCrO_row.load().to(mO.element_type)) + cute.autovec_copy(tCrO_row_out, tCgO[(None, row), m_tile, None]) + + +@cute.jit +def fp8_attention_host( + mQ: cute.Tensor, + mK: cute.Tensor, + mV: cute.Tensor, + mO: cute.Tensor, + mLSE: cute.Tensor, + mScales: cute.Tensor, + softmax_scale: cutlass.Float32, + stream, +): + mQ = cute.make_tensor(mQ.iterator, cute.select(mQ.layout, mode=[1, 2, 0])) + mK = cute.make_tensor(mK.iterator, cute.select(mK.layout, mode=[1, 2, 0])) + mV = cute.make_tensor(mV.iterator, cute.select(mV.layout, mode=[1, 2, 0])) + mO = cute.make_tensor(mO.iterator, cute.select(mO.layout, mode=[1, 2, 0])) + mLSE = cute.make_tensor(mLSE.iterator, cute.select(mLSE.layout, mode=[1, 0])) + + mma_op = cute.nvgpu.warp.MmaFP8Op( + mQ.element_type, + cutlass.Float32, + (16, 8, 32), + ) + tiled_mma = cute.make_tiled_mma( + mma_op, + (4, 1, 1), + permutation_mnk=(128, 16, 32), + ) + + copy_bits = 128 + num_threads = 128 + + sQ_layout, sQ_swizzle = _make_smem_layout_fp8( + mQ.element_type, + copy_bits, + (128, 128, 1), + ) + sK_layout, sK_swizzle = _make_smem_layout_fp8( + mK.element_type, + copy_bits, + (32, 128, 2), + ) + sV_layout, sV_swizzle = _make_smem_layout_fp8( + mV.element_type, + copy_bits, + (128, 32, 2), + ) + + atom_copy_g2s = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(cache_mode=cute.nvgpu.LoadCacheMode.GLOBAL), + mQ.element_type, + num_bits_per_copy=copy_bits, + ) + tiled_copy_Q = _make_gmem_tiled_copy( + atom_copy_g2s, + mQ.element_type, + copy_bits, + 128, + num_threads, + ) + tiled_copy_K = _make_gmem_tiled_copy( + atom_copy_g2s, + mK.element_type, + copy_bits, + 128, + num_threads, + ) + tiled_copy_V = _make_gmem_tiled_copy( + atom_copy_g2s, + mV.element_type, + copy_bits, + 32, + num_threads, + ) + + _fp8_attention_sm120( + softmax_scale, + mQ, + mK, + mV, + mO, + mLSE, + mScales, + sQ_layout, + sK_layout, + sV_layout, + sQ_swizzle, + sK_swizzle, + sV_swizzle, + tiled_copy_Q, + tiled_copy_K, + tiled_copy_V, + tiled_mma, + ).launch( + grid=(cute.ceil_div(mQ.shape[0], 128), 1, mQ.shape[2]), + block=(num_threads, 1, 1), + stream=stream, + ) + + +def key_order_for_positions(positions): + """Key held at each stored V^T position. + + Inside every 16-key group, position 4t + i holds key 8*(i//2) + 2t + i%2: + the k order of the PV A fragment packed straight from the score fragment. + """ + group = positions // 16 * 16 + quad = (positions % 16) // 4 + element = positions % 4 + return group + 8 * (element // 2) + 2 * quad + element % 2 diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py new file mode 100644 index 000000000000..dad93fcbcb20 --- /dev/null +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py @@ -0,0 +1,180 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reusable FP8 attention plan on SGLang's [S, H, 128] strided BF16 views. + +One plan owns one compiled kernel and its E4M3, scale, output and LSE buffers for a +fixed (sequence, heads, input strides, softmax scale, device). bind_inputs() points +the plan at new Q/K/V storage with the same shape and strides, so every DiT layer +reuses one compilation. prepare() quantizes with the fused Triton passes, +launch_prepared() runs the kernel. The BF16 output buffer [S, H, 128] belongs to +the plan and is overwritten by the next launch. +""" + +import math + +import cuda.bindings.driver as cuda +import cutlass +import cutlass.cute as cute +import torch +from cutlass.cute.runtime import from_dlpack + +from sglang.kernels.ops.attention.fp8_fa_sm120.fused_prep import ( + _attach_fused_state, + fused_prepare, +) +from sglang.kernels.ops.attention.fp8_fa_sm120.kernel import ( + fp8_attention_host, + key_order_for_positions, +) + +HEAD_DIM = 128 +# Q and K/V are padded separately; the kernel compiles its key mask only when S % 32 != 0. +QUERY_TILE = 128 +KEY_TILE = 32 + + +def validate_inputs(q, k, v, softmax_scale): + """Strided [S, H, 128] BF16 views on one CUDA device; returns the scale as float.""" + if ( + q.ndim != 3 + or q.shape[-1] != HEAD_DIM + or q.shape != k.shape + or q.shape != v.shape + ): + raise ValueError("Q/K/V must have equal shape [sequence, heads, 128]") + if min(q.shape[:2]) < 1: + raise ValueError("Sequence and heads must be positive") + for tensor in (q, k, v): + if ( + not tensor.is_cuda + or tensor.device != q.device + or tensor.dtype != torch.bfloat16 + ): + raise ValueError("Q/K/V must share a CUDA device and use BF16") + if tensor.stride(-1) != 1 or tensor.stride(-2) != HEAD_DIM: + raise ValueError( + "Feature and head dimensions must be contiguous within each token" + ) + if tensor.data_ptr() % 16 or tensor.stride(0) % 8: + raise ValueError("Token rows must be aligned to 16 bytes") + if tensor.stride() != q.stride(): + raise ValueError("Q/K/V must share strides") + scale = HEAD_DIM**-0.5 if softmax_scale is None else float(softmax_scale) + if not math.isfinite(scale): + raise ValueError("Softmax scale must be finite") + return scale + + +def plan_key(q, softmax_scale): + """Cache key: everything the compiled kernel and the prep constants depend on.""" + return ( + q.shape[0], + q.shape[1], + tuple(q.stride()), + q.device.index, + float(softmax_scale), + ) + + +class FP8AttentionPlan: + """Compiled kernel plus buffers for fixed strided [S, H, 128] BF16 inputs.""" + + def __init__(self, q, k, v, softmax_scale=None, block_tokens=64): + scale = validate_inputs(q, k, v, softmax_scale) + self.sequence = q.shape[0] + self.heads = q.shape[1] + self.softmax_scale = scale + self.kernel_scale = cutlass.Float32(scale) + self.padded_queries = ( + (self.sequence + QUERY_TILE - 1) // QUERY_TILE * QUERY_TILE + ) + self.padded_keys = (self.sequence + KEY_TILE - 1) // KEY_TILE * KEY_TILE + self.device = q.device + self.input_strides = tuple(q.stride()) + # The kernel side works in [H, S, 128]; these are views, not copies. + self.inputs = (q.permute(1, 0, 2), k.permute(1, 0, 2), v.permute(1, 0, 2)) + + with torch.cuda.device(q.device): + self.q_fp8 = torch.zeros( + (self.heads, self.padded_queries, HEAD_DIM), + device=q.device, + dtype=torch.float8_e4m3fn, + ) + self.k_fp8 = torch.zeros( + (self.heads, self.padded_keys, HEAD_DIM), + device=q.device, + dtype=torch.float8_e4m3fn, + ) + self.v_fp8 = torch.zeros( + (self.heads, HEAD_DIM, self.padded_keys), + device=q.device, + dtype=torch.float8_e4m3fn, + ) + self.scales = torch.empty( + (3, self.heads), device=q.device, dtype=torch.float32 + ) + # V^T is stored in the key order the kernel's register P pack expects. + positions = torch.arange(self.padded_keys, device=q.device) + self.key_order = key_order_for_positions(positions) + self.output = torch.empty(q.shape, device=q.device, dtype=torch.bfloat16) + self.lse = torch.empty( + (self.heads, self.sequence), device=q.device, dtype=torch.float32 + ) + + # Byte DLPack views work across Torch versions without FP8 DLPack support. + mQ = from_dlpack(self.q_fp8.view(torch.uint8), assumed_align=16) + mK = from_dlpack(self.k_fp8.view(torch.uint8), assumed_align=16) + mV = from_dlpack(self.v_fp8.view(torch.uint8), assumed_align=16) + mQ.element_type = cutlass.Float8E4M3FN + mK.element_type = cutlass.Float8E4M3FN + mV.element_type = cutlass.Float8E4M3FN + mO = from_dlpack(self.output.permute(1, 0, 2), assumed_align=16) + mLSE = from_dlpack(self.lse, assumed_align=16) + mScales = from_dlpack(self.scales, assumed_align=16) + self.tensor_views = (mQ, mK, mV, mO, mLSE, mScales) + + stream = cuda.CUstream(torch.cuda.current_stream(q.device).cuda_stream) + self.compiled = cute.compile( + fp8_attention_host, + *self.tensor_views, + self.kernel_scale, + stream, + ) + + _attach_fused_state(self, block_tokens) + self._prepared = False + + def bind_inputs(self, q, k, v): + """Point the plan at new Q/K/V storage with the same shape and strides.""" + expected_shape = (self.sequence, self.heads, HEAD_DIM) + for tensor in (q, k, v): + if tuple(tensor.shape) != expected_shape: + raise ValueError( + f"Plan expects shape {expected_shape}, got {tuple(tensor.shape)}" + ) + if tuple(tensor.stride()) != self.input_strides: + raise ValueError( + f"Plan expects strides {self.input_strides}, got {tuple(tensor.stride())}" + ) + if tensor.device != self.device or tensor.dtype != torch.bfloat16: + raise ValueError("Plan inputs must stay BF16 on the plan's device") + if tensor.data_ptr() % 16: + raise ValueError("Token rows must be aligned to 16 bytes") + self.inputs = (q.permute(1, 0, 2), k.permute(1, 0, 2), v.permute(1, 0, 2)) + self._prepared = False + + def prepare(self): + """Quantize the bound inputs: per-head amax, scales, packed Q/K and permuted V^T.""" + fused_prepare(self) + + def launch_prepared(self): + """Run attention on the prepared buffers; returns (output [S, H, 128], lse [H, S]).""" + if not self._prepared: + raise RuntimeError("Call prepare() before launch_prepared()") + with torch.cuda.device(self.device): + stream = cuda.CUstream(torch.cuda.current_stream(self.device).cuda_stream) + self.compiled(*self.tensor_views, self.kernel_scale, stream) + return self.output, self.lse + + def __call__(self): + self.prepare() + return self.launch_prepared() diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py new file mode 100644 index 000000000000..f0b93c426770 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py @@ -0,0 +1,216 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FP8 attention backend for SM120 GPUs (GeForce RTX 50, RTX PRO Blackwell). + +Opt in with ``--attention-backend fp8_fa_sm120``. Q, K and V are quantized per head to +E4M3 on the fly (two fused Triton passes), attention runs in FP8 with FP32 +accumulation, and the output is BF16. Expect about 5% relative RMS against cuDNN BF16 +per call on normally distributed inputs. + +Scope: dense, non-causal, batch 1, head_dim 128, BF16 inputs on an SM120 device. +Everything else goes to cuDNN SDPA. Each distinct sequence length compiles once +(about 10 s); later calls with the same shape and strides reuse the plan. +""" + +import torch + +from sglang.kernels.ops.attention.fp8_fa_sm120 import ( + HEAD_DIM, + FP8AttentionPlan, + plan_key, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( + AttentionBackend, + AttentionImpl, + AttentionMetadata, + trailing_padding_used_len, +) +from sglang.multimodal_gen.runtime.layers.attention.backends.sdpa import CudnnSDPAImpl +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# Shared by every impl instance: each DiT layer owns an impl, and per-layer plans would +# hold one copy of the E4M3 and output buffers per layer. Layers run in sequence on one stream. +_PLAN_CACHE: dict[tuple, FP8AttentionPlan] = {} + + +class FP8FlashAttentionSM120Backend(AttentionBackend): + accept_output_buffer: bool = False + + @staticmethod + def get_supported_head_sizes() -> list[int]: + return [HEAD_DIM] + + @staticmethod + def get_enum() -> AttentionBackendEnum: + return AttentionBackendEnum.FP8_FA_SM120 + + @staticmethod + def get_impl_cls() -> type["FP8FlashAttentionSM120Impl"]: + return FP8FlashAttentionSM120Impl + + +class FP8FlashAttentionSM120Impl(AttentionImpl): + def __init__( + self, + num_heads: int, + head_size: int, + causal: bool, + softmax_scale: float, + num_kv_heads: int | None = None, + prefix: str = "", + **extra_impl_args, + ) -> None: + self.num_heads = num_heads + self.head_size = head_size + self.causal = causal + self.softmax_scale = softmax_scale + self.fallback = CudnnSDPAImpl( + num_heads=num_heads, + head_size=head_size, + causal=causal, + softmax_scale=softmax_scale, + num_kv_heads=num_kv_heads, + prefix=prefix, + **extra_impl_args, + ) + self.plans = _PLAN_CACHE + self._reported_fallbacks: set[str] = set() + + # --- dispatch ------------------------------------------------------------- + + def _fallback_reason(self, query, key, value) -> str | None: + """None when the kernel applies to these [S, H, D] tensors, else why not.""" + if self.causal: + return "causal attention" + if query.shape != key.shape or key.shape != value.shape: + return "Q/K/V shapes differ" + if query.ndim != 3 or query.shape[-1] != HEAD_DIM: + return f"head_dim {query.shape[-1]} (kernel is head_dim {HEAD_DIM})" + if query.dtype != torch.bfloat16: + return f"dtype {query.dtype} (kernel takes BF16)" + if not query.is_cuda: + return "non-CUDA tensors" + if torch.cuda.get_device_capability(query.device) != (12, 0): + return "device is not SM120" + return None + + def _report_fallback(self, reason: str) -> None: + if reason in self._reported_fallbacks: + return + self._reported_fallbacks.add(reason) + logger.warning( + "fp8_fa_sm120 attention: %s; using cuDNN SDPA for these calls.", reason + ) + + # --- kernel path ---------------------------------------------------------- + + def _get_plan(self, query, key, value) -> FP8AttentionPlan: + key_ = plan_key(query, self.softmax_scale) + plan = self.plans.get(key_) + if plan is None: + logger.info( + "fp8_fa_sm120 attention: compiling for S=%d H=%d (once per shape)", + query.shape[0], + query.shape[1], + ) + plan = FP8AttentionPlan(query, key, value, self.softmax_scale) + self.plans[key_] = plan + else: + plan.bind_inputs(query, key, value) + return plan + + def _run(self, query, key, value): + """[S, H, 128] strided BF16 in, the plan's contiguous [S, H, 128] BF16 out.""" + plan = self._get_plan(query, key, value) + plan.prepare() + output, _ = plan.launch_prepared() + return output + + # --- AttentionImpl -------------------------------------------------------- + + def forward( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attn_metadata: AttentionMetadata, + ) -> torch.Tensor: + # [B, S, H, D]; the kernel handles one dense sequence at a time. + if query.shape[0] != 1: + self._report_fallback(f"batch {query.shape[0]}") + return self.fallback.forward(query, key, value, attn_metadata) + reason = self._fallback_reason(query[0], key[0], value[0]) + if reason is not None: + self._report_fallback(reason) + return self.fallback.forward(query, key, value, attn_metadata) + try: + output = self._run(query[0], key[0], value[0]) + except ValueError as error: + self._report_fallback(str(error)) + return self.fallback.forward(query, key, value, attn_metadata) + return output.unsqueeze(0) + + def forward_varlen( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + *, + cu_seqlens: torch.Tensor, + max_seqlen: int, + cu_seqlens_host: tuple[int, ...] | None = None, + ) -> torch.Tensor: + bounds = ( + cu_seqlens_host + if cu_seqlens_host is not None + else tuple(int(item) for item in cu_seqlens.tolist()) + ) + reason = self._fallback_reason(query, key, value) + if reason is not None: + self._report_fallback(reason) + return self.fallback.forward_varlen( + query, + key, + value, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + cu_seqlens_host=bounds, + ) + + # MiniMax-H3 packs one live document as (0, used, total); the tail is + # 64-aligned padding that downstream masks, so it stays zero. + used = trailing_padding_used_len( + total_tokens=query.shape[0], + max_seqlen=max_seqlen, + bounds=bounds, + ) + try: + if used is not None: + live_output = self._run(query[:used], key[:used], value[:used]) + if used == query.shape[0]: + return live_output + output = torch.zeros_like(query) + output[:used].copy_(live_output) + return output + + output = torch.empty_like(query) + for start, stop in zip(bounds[:-1], bounds[1:]): + if start == stop: + continue + segment_output = self._run( + query[start:stop], key[start:stop], value[start:stop] + ) + output[start:stop].copy_(segment_output) + return output + except ValueError as error: + self._report_fallback(str(error)) + return self.fallback.forward_varlen( + query, + key, + value, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + cu_seqlens_host=bounds, + ) diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index f1878f8f0fb0..19388b1b9f85 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -516,6 +516,36 @@ def resolve(cls, platform) -> AttentionBackendEnum: return AttentionBackendEnum.FA +class _FP8FlashAttentionSM120BackendResolver(_CudaAttentionBackendResolver): + backend = AttentionBackendEnum.FP8_FA_SM120 + + # CuTe-DSL mma.sync FP8 kernel written for SM120 (GeForce RTX 50 / RTX PRO + # Blackwell). Dense non-causal head_dim 128 only; the backend itself falls + # back to cuDNN SDPA for other calls. + @classmethod + def resolve(cls, platform) -> str | AttentionBackendEnum: + if platform.get_device_capability() != (12, 0): + logger.warning( + "fp8_fa_sm120 attention needs an SM120 device; falling back to cuDNN SDPA." + ) + return AttentionBackendEnum.TORCH_CUDNN_SDPA + try: + import cutlass.cute # noqa: F401 + import triton # noqa: F401 + + from sglang.multimodal_gen.runtime.layers.attention.backends.fp8_fa_sm120_attn import ( # noqa: F401 + FP8FlashAttentionSM120Backend, + ) + + return "sglang.multimodal_gen.runtime.layers.attention.backends.fp8_fa_sm120_attn.FP8FlashAttentionSM120Backend" + except ImportError as error: + logger.warning( + "fp8_fa_sm120 attention backend failed to import (%s); falling back to cuDNN SDPA.", + error, + ) + return AttentionBackendEnum.TORCH_CUDNN_SDPA + + _CUDA_ATTENTION_BACKEND_RESOLVERS = { resolver.backend: resolver for resolver in ( @@ -539,6 +569,7 @@ def resolve(cls, platform) -> AttentionBackendEnum: _SubBlockSparseAttentionBackendResolver, _FlashAttention2BackendResolver, _FlashAttentionBackendResolver, + _FP8FlashAttentionSM120BackendResolver, ) } diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 129e92c5d210..3372784748dc 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -49,6 +49,7 @@ class AttentionBackendEnum(enum.Enum): SOL_ATTN = enum.auto() SUBBLOCK_SPARSE_ATTN = enum.auto() CUBE_SPARSE_ATTN = enum.auto() + FP8_FA_SM120 = enum.auto() NO_ATTENTION = enum.auto() def __str__(self): diff --git a/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py b/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py new file mode 100644 index 000000000000..42dcc95788c3 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py @@ -0,0 +1,174 @@ +# SPDX-License-Identifier: Apache-2.0 +"""fp8_fa_sm120 attention backend against cuDNN SDPA on MiniMax-H3 style inputs. + +Needs an SM120 GPU; skipped elsewhere. Shapes are kept small (8 heads) so the +CuTe-DSL compile and the runs stay in seconds. +""" + +import pytest +import torch + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (12, 0), + reason="fp8_fa_sm120 attention needs an SM120 GPU", +) + +HEADS = 8 +HEAD_DIM = 128 +SOFTMAX_SCALE = HEAD_DIM**-0.5 + + +def _make_impls(causal=False): + from sglang.multimodal_gen.runtime.layers.attention.backends.fp8_fa_sm120_attn import ( + FP8FlashAttentionSM120Impl, + ) + from sglang.multimodal_gen.runtime.layers.attention.backends.sdpa import ( + CudnnSDPAImpl, + ) + + ours = FP8FlashAttentionSM120Impl( + num_heads=HEADS, head_size=HEAD_DIM, causal=causal, softmax_scale=SOFTMAX_SCALE + ) + reference = CudnnSDPAImpl( + num_heads=HEADS, head_size=HEAD_DIM, causal=causal, softmax_scale=SOFTMAX_SCALE + ) + return ours, reference + + +def _fused_qkv_views(sequence, seed=0): + """Q/K/V as views into one [S, 3*H*D] projection output, like the DiT hands them over.""" + generator = torch.Generator(device="cuda").manual_seed(seed) + qkv = torch.randn( + (sequence, 3 * HEADS * HEAD_DIM), + device="cuda", + dtype=torch.bfloat16, + generator=generator, + ) + q = qkv[:, 0 : HEADS * HEAD_DIM].view(sequence, HEADS, HEAD_DIM) + k = qkv[:, HEADS * HEAD_DIM : 2 * HEADS * HEAD_DIM].view(sequence, HEADS, HEAD_DIM) + v = qkv[:, 2 * HEADS * HEAD_DIM :].view(sequence, HEADS, HEAD_DIM) + return q, k, v + + +def _error_metrics(candidate, reference): + candidate = candidate.float() + reference = reference.float() + relative_rms = (candidate - reference).pow(2).mean().sqrt() / reference.pow( + 2 + ).mean().sqrt() + cosine = torch.nn.functional.cosine_similarity( + candidate.flatten(), reference.flatten(), dim=0 + ) + return float(relative_rms), float(cosine) + + +@pytest.mark.parametrize("sequence", [4096, 5980]) +def test_forward_matches_cudnn(sequence): + ours, reference = _make_impls() + q, k, v = _fused_qkv_views(sequence) + + output = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + expected = reference.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + + assert output.shape == expected.shape + assert output.dtype == torch.bfloat16 + assert torch.isfinite(output.float()).all() + relative_rms, cosine = _error_metrics(output, expected) + # Per-head E4M3 quantization of Q/K/V: ~5% relative RMS on normal inputs. + assert relative_rms < 0.08, relative_rms + assert cosine > 0.995, cosine + + +def test_plan_is_reused_across_calls(): + ours, reference = _make_impls() + q, k, v = _fused_qkv_views(4096, seed=1) + first = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None).clone() + plans_after_first = len(ours.plans) + + q2, k2, v2 = _fused_qkv_views(4096, seed=2) + second = ours.forward(q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None) + assert len(ours.plans) == plans_after_first + + # A second impl (another DiT layer) shares the same plan. + other, _ = _make_impls() + other.forward(q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None) + assert len(other.plans) == plans_after_first + + expected = reference.forward( + q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None + ) + relative_rms, _ = _error_metrics(second, expected) + assert relative_rms < 0.08, relative_rms + assert not torch.equal(first, second) + + +def test_varlen_trailing_padding_keeps_tail_zero(): + ours, _ = _make_impls() + used = 5980 + total = 6016 # 64-aligned tail padding, as MiniMax-H3 packs it + q, k, v = _fused_qkv_views(total) + cu_seqlens = torch.tensor([0, used, total], device="cuda", dtype=torch.int32) + + output = ours.forward_varlen( + q, + k, + v, + cu_seqlens=cu_seqlens, + max_seqlen=used, + cu_seqlens_host=(0, used, total), + ) + live = ours.forward( + q[:used].unsqueeze(0), k[:used].unsqueeze(0), v[:used].unsqueeze(0), None + )[0] + + assert output.shape == q.shape + assert torch.equal(output[:used], live) + assert torch.count_nonzero(output[used:]) == 0 + + +def test_varlen_multiple_segments(): + ours, reference = _make_impls() + bounds = (0, 1024, 3072, 4096) + q, k, v = _fused_qkv_views(bounds[-1]) + cu_seqlens = torch.tensor(bounds, device="cuda", dtype=torch.int32) + + output = ours.forward_varlen( + q, k, v, cu_seqlens=cu_seqlens, max_seqlen=2048, cu_seqlens_host=bounds + ) + expected = reference.forward_varlen( + q, k, v, cu_seqlens=cu_seqlens, max_seqlen=2048, cu_seqlens_host=bounds + ) + + for start, stop in zip(bounds[:-1], bounds[1:]): + relative_rms, _ = _error_metrics(output[start:stop], expected[start:stop]) + assert relative_rms < 0.08, (start, stop, relative_rms) + + +def test_causal_falls_back_to_cudnn(): + ours, reference = _make_impls(causal=True) + q, k, v = _fused_qkv_views(1024) + plans_before = len(ours.plans) + + output = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + expected = reference.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + + assert torch.equal(output, expected) + assert len(ours.plans) == plans_before + + +def test_backend_resolves_by_name(): + from sglang.multimodal_gen.runtime.platforms import ( + AttentionBackendEnum, + current_platform, + ) + + class_path = current_platform.get_attn_backend_cls_str( + AttentionBackendEnum.FP8_FA_SM120, HEAD_DIM, torch.bfloat16 + ) + assert class_path.endswith("fp8_fa_sm120_attn.FP8FlashAttentionSM120Backend") + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) From eae2853c44f6953687e4725415d5b6443366d63a Mon Sep 17 00:00:00 2001 From: Emre Albayrak Date: Sun, 20 Sep 2026 12:39:12 +0300 Subject: [PATCH 2/2] [diffusion] fp8_fa_sm120: allocate buffers per call, cache compiled kernels only _PLAN_CACHE kept one plan per (S, H, strides, scale) forever. Each plan owned about 1.0 GiB of FP8, output and LSE buffers at S=30272 H=56 and a reference to the last BF16 Q/K/V views (1.2 GiB), so a worker serving several sequence lengths ran out of memory. - plan.py: fp8_attention() replaces FP8AttentionPlan. The compiled kernel is cached per (S, H, device); the workspace comes from the caching allocator on every call and the output belongs to the caller. No input references are kept. Allocation plus DLPack views cost 47 us on an RTX 5080 against 2.9 ms (S=4096) and 127 ms (S=30272) of kernel time. - fused_prep.py: the pack pass writes every position of the padded buffers (stores masked on the buffer position, grid up to padded_queries), so the workspace can come from torch.empty. The V^T key permutation left unwritten positions inside the last 16-key group; they were only safe because the plan zeroed its buffers once. - Tests: memory returns to baseline across distinct sequence lengths; no NaN byte survives the prep in a 0xFF-filled workspace. Output and LSE are bit-identical to the previous commit for S in (1052, 4096, 5980, 30272), H in (8, 56). --- .../ops/attention/fp8_fa_sm120/__init__.py | 9 +- .../ops/attention/fp8_fa_sm120/fused_prep.py | 117 +++++---- .../ops/attention/fp8_fa_sm120/plan.py | 222 ++++++++---------- .../attention/backends/fp8_fa_sm120_attn.py | 35 +-- .../test/unit/test_fp8_fa_sm120_attn.py | 64 ++++- 5 files changed, 215 insertions(+), 232 deletions(-) diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py index 9d1bf36c8f45..45df385e6f52 100644 --- a/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/__init__.py @@ -1,11 +1,12 @@ # SPDX-License-Identifier: Apache-2.0 -"""SM120 FP8 attention: CuTe-DSL kernel, fused Triton prep and the reusable plan.""" +"""SM120 FP8 attention: CuTe-DSL kernel, fused Triton prep and the per-call entry point.""" from sglang.kernels.ops.attention.fp8_fa_sm120.plan import ( HEAD_DIM, - FP8AttentionPlan, - plan_key, + Workspace, + fp8_attention, + kernel_key, validate_inputs, ) -__all__ = ["HEAD_DIM", "FP8AttentionPlan", "plan_key", "validate_inputs"] +__all__ = ["HEAD_DIM", "Workspace", "fp8_attention", "kernel_key", "validate_inputs"] diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py index 3e4a36c36ece..19d96269f912 100644 --- a/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/fused_prep.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Fused FP8 preparation for the SM120 kernel: two Triton passes into the plan buffers. +"""Fused FP8 preparation for the SM120 kernel: two Triton passes into the workspace. Replaces a Torch chain (FP32 copy, amax, divide, cast, padded copy, V transpose copy) with one per-head amax pass and one pack pass that @@ -9,14 +9,18 @@ V [H, 128, padded_S] transposed here, so the kernel stays unchanged Scales stay per head and the quantization recipe is the same (amax/448, div.rn, -clamp, E4M3), so the kernel's numerics do not move. Padding rows are zeroed once -by the plan and never written here. +clamp, E4M3), so the kernel's numerics do not move. The pack pass writes every +position of the padded buffers, padding included, because the workspace comes +from torch.empty. """ import torch import triton import triton.language as tl +# Power of two >= 16: the V^T key permutation works in groups of 16 tokens. +BLOCK_TOKENS = 64 + # Triton's `/` lowers to div.full.f32 (approximate); Torch divides with div.rn. # The FP8 bytes must match the Torch prepare() path, so the exact division is kept. @@ -103,8 +107,10 @@ def _head_scales( tl.store(scales + offsets, scale, mask=valid) -# Pass 2: same grid as pass 1. Q/K tiles go out in [token, column] order; -# the V tile is transposed in registers/shared and goes out in [column, token]. +# Pass 2: one program per (head, token block) up to padded_queries. Q/K tiles go +# out in [token, column] order; the V tile is transposed in registers/shared and +# goes out in [column, token]. Loads are masked on the input token, stores on the +# buffer position, so padding positions receive quantized zeros. @triton.jit def _pack_qkv( q, @@ -121,7 +127,6 @@ def _pack_qkv( padded_keys: tl.constexpr, heads: tl.constexpr, block_tokens: tl.constexpr, - permute_keys: tl.constexpr, ): head = tl.program_id(0) tokens = tl.program_id(1) * block_tokens + tl.arange(0, block_tokens) @@ -135,12 +140,14 @@ def _pack_qkv( # Q is padded to the 128-row query tile, K/V to the 32-key tile. hsd_offsets = head * padded_queries * 128 + tokens[:, None] * 128 + columns[None, :] k_offsets = head * padded_keys * 128 + tokens[:, None] * 128 + columns[None, :] + in_queries = tokens[:, None] < padded_queries + in_keys = tokens[:, None] < padded_keys q_tile = _load_tile(q, head, tokens, columns, valid, head_stride, token_stride) - tl.store(q_fp8 + hsd_offsets, _quantize(q_tile, q_scale), mask=valid) + tl.store(q_fp8 + hsd_offsets, _quantize(q_tile, q_scale), mask=in_queries) k_tile = _load_tile(k, head, tokens, columns, valid, head_stride, token_stride) - tl.store(k_fp8 + k_offsets, _quantize(k_tile, k_scale), mask=valid) + tl.store(k_fp8 + k_offsets, _quantize(k_tile, k_scale), mask=in_keys) # The kernel packs P straight from the score fragment; V^T's stored key # order must follow it: stored position p holds key @@ -148,96 +155,78 @@ def _pack_qkv( # The permutation is applied on the LOAD side as a row gather: every gathered # row is still 256 contiguous bytes, and the transposed store stays contiguous. # Permuting the store instead scatters bytes inside 16 B windows and doubled - # the prepare time. Padding keys are never written. - if permute_keys: - source_tokens = ( - tokens // 16 * 16 - + 8 * ((tokens % 4) // 2) - + 2 * ((tokens % 16) // 4) - + tokens % 2 - ) - else: - source_tokens = tokens + # the prepare time. + source_tokens = ( + tokens // 16 * 16 + + 8 * ((tokens % 4) // 2) + + 2 * ((tokens % 16) // 4) + + tokens % 2 + ) valid_source = source_tokens[:, None] < sequence hds_offsets = ( head * 128 * padded_keys + columns[:, None] * padded_keys + tokens[None, :] ) - valid_transposed = source_tokens[None, :] < sequence + in_keys_transposed = tokens[None, :] < padded_keys v_tile = _load_tile( v, head, source_tokens, columns, valid_source, head_stride, token_stride ) tl.store( - v_fp8 + hds_offsets, tl.trans(_quantize(v_tile, v_scale)), mask=valid_transposed + v_fp8 + hds_offsets, + tl.trans(_quantize(v_tile, v_scale)), + mask=in_keys_transposed, ) -def _attach_fused_state(plan, block_tokens): - q, k, v = plan.inputs - for tensor in (k, v): - if tensor.stride() != q.stride(): - raise ValueError("Q/K/V must share strides for fused preparation") - if q.stride(2) != 1: - raise ValueError("Head dimension must be contiguous for fused preparation") - if block_tokens & (block_tokens - 1) or block_tokens < 16: - raise ValueError( - "block_tokens must be a power of two >= 16 (key permutation groups are 16 wide)" - ) - - heads = q.shape[0] - plan.block_tokens = block_tokens - plan.head_stride = q.stride(0) - plan.token_stride = q.stride(1) - plan.maximums = torch.zeros((3, heads), device=q.device, dtype=torch.float32) - +def fused_prepare(q, k, v, workspace): + """Three launches: per-head amax, scale finalize, direct pack into the workspace. -def fused_prepare(plan): - """Three launches: per-head amax, scale finalize, direct pack into the plan buffers.""" - q, k, v = plan.inputs + q, k, v are [H, S, 128] views with one shared stride set; the workspace buffers are + sized for this S and H. + """ heads = q.shape[0] - padded_queries = plan.q_fp8.shape[1] - padded_keys = plan.k_fp8.shape[1] - grid = (heads, triton.cdiv(plan.sequence, plan.block_tokens)) + sequence = q.shape[1] + padded_queries = workspace.q_fp8.shape[1] + padded_keys = workspace.k_fp8.shape[1] + head_stride = q.stride(0) + token_stride = q.stride(1) with torch.cuda.device(q.device): - plan.maximums.zero_() - _head_amax_qkv[grid]( + _head_amax_qkv[(heads, triton.cdiv(sequence, BLOCK_TOKENS))]( q, k, v, - plan.maximums, - plan.sequence, - head_stride=plan.head_stride, - token_stride=plan.token_stride, + workspace.maximums, + sequence, + head_stride=head_stride, + token_stride=token_stride, heads=heads, - block_tokens=plan.block_tokens, + block_tokens=BLOCK_TOKENS, num_warps=8, ) _head_scales[(1,)]( - plan.maximums, - plan.scales, + workspace.maximums, + workspace.scales, INVERSE_FP8_MAX, count=3 * heads, block=triton.next_power_of_2(3 * heads), num_warps=1, ) - _pack_qkv[grid]( + _pack_qkv[(heads, triton.cdiv(padded_queries, BLOCK_TOKENS))]( q, k, v, - plan.q_fp8, - plan.k_fp8, - plan.v_fp8, - plan.scales, - plan.sequence, - head_stride=plan.head_stride, - token_stride=plan.token_stride, + workspace.q_fp8, + workspace.k_fp8, + workspace.v_fp8, + workspace.scales, + sequence, + head_stride=head_stride, + token_stride=token_stride, padded_queries=padded_queries, padded_keys=padded_keys, heads=heads, - block_tokens=plan.block_tokens, - permute_keys=plan.key_order is not None, + block_tokens=BLOCK_TOKENS, num_warps=8, ) - plan._prepared = True diff --git a/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py b/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py index dad93fcbcb20..b263da3b5c9e 100644 --- a/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py +++ b/python/sglang/kernels/ops/attention/fp8_fa_sm120/plan.py @@ -1,36 +1,37 @@ # SPDX-License-Identifier: Apache-2.0 -"""Reusable FP8 attention plan on SGLang's [S, H, 128] strided BF16 views. - -One plan owns one compiled kernel and its E4M3, scale, output and LSE buffers for a -fixed (sequence, heads, input strides, softmax scale, device). bind_inputs() points -the plan at new Q/K/V storage with the same shape and strides, so every DiT layer -reuses one compilation. prepare() quantizes with the fused Triton passes, -launch_prepared() runs the kernel. The BF16 output buffer [S, H, 128] belongs to -the plan and is overwritten by the next launch. +"""FP8 attention on SGLang's [S, H, 128] strided BF16 views. + +fp8_attention() quantizes Q/K/V per head with the fused Triton passes, runs the CuTe-DSL +kernel and returns (output [S, H, 128] BF16, lse [H, S] FP32). The compiled kernel is +cached per (sequence, heads, device); the E4M3, scale, output and LSE buffers come from +the caching allocator on every call and belong to the caller, nothing is retained +between calls. """ +import logging import math import cuda.bindings.driver as cuda import cutlass import cutlass.cute as cute +import msgspec import torch from cutlass.cute.runtime import from_dlpack -from sglang.kernels.ops.attention.fp8_fa_sm120.fused_prep import ( - _attach_fused_state, - fused_prepare, -) -from sglang.kernels.ops.attention.fp8_fa_sm120.kernel import ( - fp8_attention_host, - key_order_for_positions, -) +from sglang.kernels.ops.attention.fp8_fa_sm120.fused_prep import fused_prepare +from sglang.kernels.ops.attention.fp8_fa_sm120.kernel import fp8_attention_host + +logger = logging.getLogger(__name__) HEAD_DIM = 128 # Q and K/V are padded separately; the kernel compiles its key mask only when S % 32 != 0. QUERY_TILE = 128 KEY_TILE = 32 +# Compiled kernels only, KBs of host state each and about 10 s to build; the buffers +# they run on are allocated per call. +_KERNEL_CACHE: dict[tuple, object] = {} + def validate_inputs(q, k, v, softmax_scale): """Strided [S, H, 128] BF16 views on one CUDA device; returns the scale as float.""" @@ -64,117 +65,86 @@ def validate_inputs(q, k, v, softmax_scale): return scale -def plan_key(q, softmax_scale): - """Cache key: everything the compiled kernel and the prep constants depend on.""" - return ( - q.shape[0], - q.shape[1], - tuple(q.stride()), - q.device.index, - float(softmax_scale), +def kernel_key(q): + """Cache key: the kernel bakes the buffer shapes, which follow from (S, H).""" + return (q.shape[0], q.shape[1], q.device.index) + + +class Workspace(msgspec.Struct, frozen=True, kw_only=True): + """Kernel-side buffers for one call; the kernel works in [H, S, 128].""" + + q_fp8: torch.Tensor + k_fp8: torch.Tensor + v_fp8: torch.Tensor + scales: torch.Tensor + maximums: torch.Tensor + output: torch.Tensor + lse: torch.Tensor + + +def _allocate_workspace(q): + sequence = q.shape[0] + heads = q.shape[1] + padded_queries = (sequence + QUERY_TILE - 1) // QUERY_TILE * QUERY_TILE + padded_keys = (sequence + KEY_TILE - 1) // KEY_TILE * KEY_TILE + device = q.device + return Workspace( + q_fp8=torch.empty( + (heads, padded_queries, HEAD_DIM), device=device, dtype=torch.float8_e4m3fn + ), + k_fp8=torch.empty( + (heads, padded_keys, HEAD_DIM), device=device, dtype=torch.float8_e4m3fn + ), + v_fp8=torch.empty( + (heads, HEAD_DIM, padded_keys), device=device, dtype=torch.float8_e4m3fn + ), + scales=torch.empty((3, heads), device=device, dtype=torch.float32), + maximums=torch.zeros((3, heads), device=device, dtype=torch.float32), + output=torch.empty(q.shape, device=device, dtype=torch.bfloat16), + lse=torch.empty((heads, sequence), device=device, dtype=torch.float32), ) -class FP8AttentionPlan: - """Compiled kernel plus buffers for fixed strided [S, H, 128] BF16 inputs.""" - - def __init__(self, q, k, v, softmax_scale=None, block_tokens=64): - scale = validate_inputs(q, k, v, softmax_scale) - self.sequence = q.shape[0] - self.heads = q.shape[1] - self.softmax_scale = scale - self.kernel_scale = cutlass.Float32(scale) - self.padded_queries = ( - (self.sequence + QUERY_TILE - 1) // QUERY_TILE * QUERY_TILE - ) - self.padded_keys = (self.sequence + KEY_TILE - 1) // KEY_TILE * KEY_TILE - self.device = q.device - self.input_strides = tuple(q.stride()) - # The kernel side works in [H, S, 128]; these are views, not copies. - self.inputs = (q.permute(1, 0, 2), k.permute(1, 0, 2), v.permute(1, 0, 2)) - - with torch.cuda.device(q.device): - self.q_fp8 = torch.zeros( - (self.heads, self.padded_queries, HEAD_DIM), - device=q.device, - dtype=torch.float8_e4m3fn, - ) - self.k_fp8 = torch.zeros( - (self.heads, self.padded_keys, HEAD_DIM), - device=q.device, - dtype=torch.float8_e4m3fn, - ) - self.v_fp8 = torch.zeros( - (self.heads, HEAD_DIM, self.padded_keys), - device=q.device, - dtype=torch.float8_e4m3fn, +def _kernel_views(workspace): + # Byte DLPack views work across Torch versions without FP8 DLPack support. + mQ = from_dlpack(workspace.q_fp8.view(torch.uint8), assumed_align=16) + mK = from_dlpack(workspace.k_fp8.view(torch.uint8), assumed_align=16) + mV = from_dlpack(workspace.v_fp8.view(torch.uint8), assumed_align=16) + mQ.element_type = cutlass.Float8E4M3FN + mK.element_type = cutlass.Float8E4M3FN + mV.element_type = cutlass.Float8E4M3FN + mO = from_dlpack(workspace.output.permute(1, 0, 2), assumed_align=16) + mLSE = from_dlpack(workspace.lse, assumed_align=16) + mScales = from_dlpack(workspace.scales, assumed_align=16) + return (mQ, mK, mV, mO, mLSE, mScales) + + +def fp8_attention(q, k, v, softmax_scale=None): + """Attention over strided [S, H, 128] BF16 views; returns (output [S, H, 128], lse [H, S]).""" + scale = validate_inputs(q, k, v, softmax_scale) + kernel_scale = cutlass.Float32(scale) + key = kernel_key(q) + + with torch.cuda.device(q.device): + workspace = _allocate_workspace(q) + views = _kernel_views(workspace) + stream = cuda.CUstream(torch.cuda.current_stream(q.device).cuda_stream) + + compiled = _KERNEL_CACHE.get(key) + if compiled is None: + logger.info( + "fp8_fa_sm120 attention: compiling for S=%d H=%d (once per shape)", + q.shape[0], + q.shape[1], ) - self.scales = torch.empty( - (3, self.heads), device=q.device, dtype=torch.float32 - ) - # V^T is stored in the key order the kernel's register P pack expects. - positions = torch.arange(self.padded_keys, device=q.device) - self.key_order = key_order_for_positions(positions) - self.output = torch.empty(q.shape, device=q.device, dtype=torch.bfloat16) - self.lse = torch.empty( - (self.heads, self.sequence), device=q.device, dtype=torch.float32 - ) - - # Byte DLPack views work across Torch versions without FP8 DLPack support. - mQ = from_dlpack(self.q_fp8.view(torch.uint8), assumed_align=16) - mK = from_dlpack(self.k_fp8.view(torch.uint8), assumed_align=16) - mV = from_dlpack(self.v_fp8.view(torch.uint8), assumed_align=16) - mQ.element_type = cutlass.Float8E4M3FN - mK.element_type = cutlass.Float8E4M3FN - mV.element_type = cutlass.Float8E4M3FN - mO = from_dlpack(self.output.permute(1, 0, 2), assumed_align=16) - mLSE = from_dlpack(self.lse, assumed_align=16) - mScales = from_dlpack(self.scales, assumed_align=16) - self.tensor_views = (mQ, mK, mV, mO, mLSE, mScales) - - stream = cuda.CUstream(torch.cuda.current_stream(q.device).cuda_stream) - self.compiled = cute.compile( - fp8_attention_host, - *self.tensor_views, - self.kernel_scale, - stream, - ) - - _attach_fused_state(self, block_tokens) - self._prepared = False - - def bind_inputs(self, q, k, v): - """Point the plan at new Q/K/V storage with the same shape and strides.""" - expected_shape = (self.sequence, self.heads, HEAD_DIM) - for tensor in (q, k, v): - if tuple(tensor.shape) != expected_shape: - raise ValueError( - f"Plan expects shape {expected_shape}, got {tuple(tensor.shape)}" - ) - if tuple(tensor.stride()) != self.input_strides: - raise ValueError( - f"Plan expects strides {self.input_strides}, got {tuple(tensor.stride())}" - ) - if tensor.device != self.device or tensor.dtype != torch.bfloat16: - raise ValueError("Plan inputs must stay BF16 on the plan's device") - if tensor.data_ptr() % 16: - raise ValueError("Token rows must be aligned to 16 bytes") - self.inputs = (q.permute(1, 0, 2), k.permute(1, 0, 2), v.permute(1, 0, 2)) - self._prepared = False - - def prepare(self): - """Quantize the bound inputs: per-head amax, scales, packed Q/K and permuted V^T.""" - fused_prepare(self) - - def launch_prepared(self): - """Run attention on the prepared buffers; returns (output [S, H, 128], lse [H, S]).""" - if not self._prepared: - raise RuntimeError("Call prepare() before launch_prepared()") - with torch.cuda.device(self.device): - stream = cuda.CUstream(torch.cuda.current_stream(self.device).cuda_stream) - self.compiled(*self.tensor_views, self.kernel_scale, stream) - return self.output, self.lse - - def __call__(self): - self.prepare() - return self.launch_prepared() + compiled = cute.compile(fp8_attention_host, *views, kernel_scale, stream) + _KERNEL_CACHE[key] = compiled + + fused_prepare( + q=q.permute(1, 0, 2), + k=k.permute(1, 0, 2), + v=v.permute(1, 0, 2), + workspace=workspace, + ) + compiled(*views, kernel_scale, stream) + return workspace.output, workspace.lse diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py index f0b93c426770..b8faa26f7503 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py @@ -8,16 +8,13 @@ Scope: dense, non-causal, batch 1, head_dim 128, BF16 inputs on an SM120 device. Everything else goes to cuDNN SDPA. Each distinct sequence length compiles once -(about 10 s); later calls with the same shape and strides reuse the plan. +(about 10 s); later calls with the same shape reuse the compiled kernel. The FP8 and +output buffers are allocated per call and belong to the caller. """ import torch -from sglang.kernels.ops.attention.fp8_fa_sm120 import ( - HEAD_DIM, - FP8AttentionPlan, - plan_key, -) +from sglang.kernels.ops.attention.fp8_fa_sm120 import HEAD_DIM, fp8_attention from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( AttentionBackend, AttentionImpl, @@ -30,10 +27,6 @@ logger = init_logger(__name__) -# Shared by every impl instance: each DiT layer owns an impl, and per-layer plans would -# hold one copy of the E4M3 and output buffers per layer. Layers run in sequence on one stream. -_PLAN_CACHE: dict[tuple, FP8AttentionPlan] = {} - class FP8FlashAttentionSM120Backend(AttentionBackend): accept_output_buffer: bool = False @@ -75,7 +68,6 @@ def __init__( prefix=prefix, **extra_impl_args, ) - self.plans = _PLAN_CACHE self._reported_fallbacks: set[str] = set() # --- dispatch ------------------------------------------------------------- @@ -106,26 +98,9 @@ def _report_fallback(self, reason: str) -> None: # --- kernel path ---------------------------------------------------------- - def _get_plan(self, query, key, value) -> FP8AttentionPlan: - key_ = plan_key(query, self.softmax_scale) - plan = self.plans.get(key_) - if plan is None: - logger.info( - "fp8_fa_sm120 attention: compiling for S=%d H=%d (once per shape)", - query.shape[0], - query.shape[1], - ) - plan = FP8AttentionPlan(query, key, value, self.softmax_scale) - self.plans[key_] = plan - else: - plan.bind_inputs(query, key, value) - return plan - def _run(self, query, key, value): - """[S, H, 128] strided BF16 in, the plan's contiguous [S, H, 128] BF16 out.""" - plan = self._get_plan(query, key, value) - plan.prepare() - output, _ = plan.launch_prepared() + """[S, H, 128] strided BF16 in, a new contiguous [S, H, 128] BF16 out.""" + output, _ = fp8_attention(query, key, value, softmax_scale=self.softmax_scale) return output # --- AttentionImpl -------------------------------------------------------- diff --git a/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py b/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py index 42dcc95788c3..c240c8f1b7a5 100644 --- a/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py +++ b/python/sglang/multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py @@ -79,20 +79,26 @@ def test_forward_matches_cudnn(sequence): assert cosine > 0.995, cosine -def test_plan_is_reused_across_calls(): +def _compiled_kernels(): + from sglang.kernels.ops.attention.fp8_fa_sm120.plan import _KERNEL_CACHE + + return len(_KERNEL_CACHE) + + +def test_kernel_is_reused_across_calls(): ours, reference = _make_impls() q, k, v = _fused_qkv_views(4096, seed=1) - first = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None).clone() - plans_after_first = len(ours.plans) + first = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + kernels_after_first = _compiled_kernels() q2, k2, v2 = _fused_qkv_views(4096, seed=2) second = ours.forward(q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None) - assert len(ours.plans) == plans_after_first + assert _compiled_kernels() == kernels_after_first - # A second impl (another DiT layer) shares the same plan. + # A second impl (another DiT layer) shares the same compiled kernel. other, _ = _make_impls() other.forward(q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None) - assert len(other.plans) == plans_after_first + assert _compiled_kernels() == kernels_after_first expected = reference.forward( q2.unsqueeze(0), k2.unsqueeze(0), v2.unsqueeze(0), None @@ -144,16 +150,58 @@ def test_varlen_multiple_segments(): assert relative_rms < 0.08, (start, stop, relative_rms) +def test_distinct_shapes_release_memory(): + """A call must not retain GPU buffers keyed by shape: each distinct sequence length + used to keep its FP8/output buffers and the last Q/K/V alive, ~2.2 GiB at H3 size.""" + ours, _ = _make_impls() + shapes = [_fused_qkv_views(1052), _fused_qkv_views(1088)] + q, k, v = shapes[0] + ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + + for q, k, v in shapes: + output = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) + assert torch.isfinite(output.float()).all() + del output + + torch.cuda.synchronize() + assert torch.cuda.memory_allocated() == baseline + + +def test_prep_defines_every_padded_byte(): + """The FP8 buffers come from torch.empty, so the prep must write every padding + position; an E4M3 NaN byte left in the V^T tail makes masked keys poison the row.""" + from sglang.kernels.ops.attention.fp8_fa_sm120.fused_prep import fused_prepare + from sglang.kernels.ops.attention.fp8_fa_sm120.plan import _allocate_workspace + + q, k, v = _fused_qkv_views(1052) + workspace = _allocate_workspace(q) + for buffer in (workspace.q_fp8, workspace.k_fp8, workspace.v_fp8): + buffer.view(torch.uint8).fill_(0xFF) + + fused_prepare( + q=q.permute(1, 0, 2), + k=k.permute(1, 0, 2), + v=v.permute(1, 0, 2), + workspace=workspace, + ) + + for buffer in (workspace.q_fp8, workspace.k_fp8, workspace.v_fp8): + raw = buffer.view(torch.uint8) + assert not ((raw == 0xFF) | (raw == 0x7F)).any() + + def test_causal_falls_back_to_cudnn(): ours, reference = _make_impls(causal=True) q, k, v = _fused_qkv_views(1024) - plans_before = len(ours.plans) + kernels_before = _compiled_kernels() output = ours.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) expected = reference.forward(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), None) assert torch.equal(output, expected) - assert len(ours.plans) == plans_before + assert _compiled_kernels() == kernels_before def test_backend_resolves_by_name():