[DSA] Integrate Q8KV8 FP8 Sparse MLA Prefill into the DSA Backend (DeepSeek-V3.2) (#30514)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
Ho-Ren (Jack) Chuang
2026-07-19 11:58:16 +08:00
committed by GitHub
co-authored by Claude Opus 4.8 Xiaoyu Zhang
parent cce5fe7696
commit b8ec544946
10 changed files with 801 additions and 95 deletions
@@ -64,7 +64,7 @@ import { DeepSeekV32Deployment } from "/src/snippets/autoregressive/deepseek-v32
### 3.2 Configuration Tips ### 3.2 Configuration Tips
- **Short-sequence MHA prefill (adaptive):** For prefill sequences shorter than 2048 tokens (default threshold), the DSA backend automatically switches to standard MHA (using FlashAttention variable-length on SM90, TRT-LLM ragged MHA on SM100). To extend this to longer sequences set env var `SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD` to a larger value (potential minor accuracy trade-off). - **Short-sequence MHA prefill (adaptive):** For prefill sequences shorter than 2048 tokens (default threshold), the DSA backend automatically switches to standard MHA (using FlashAttention variable-length on SM90, TRT-LLM ragged MHA on SM100). To extend this to longer sequences set env var `SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD` to a larger value (potential minor accuracy trade-off).
- **DSA prefill/decode attention kernels (`--dsa-prefill-backend`, `--dsa-decode-backend`):** The `dsa` backend is automatically selected for DeepSeek-V3.2. Available kernels: `flashmla_sparse`, `flashmla_kv`, `flashmla_auto`, `fa3` (Hopper only), `tilelang` (GPU/HPU/NPU), `aiter` (AMD, decode only), `trtllm` (Blackwell only). Defaults: Hopper BF16 KV → `flashmla_sparse` prefill / `fa3` decode; Hopper FP8 KV → `flashmla_kv` both; Blackwell BF16 → `flashmla_sparse` / `trtllm`; Blackwell FP8 → `trtllm` both. - **DSA prefill/decode attention kernels (`--dsa-prefill-backend`, `--dsa-decode-backend`):** The `dsa` backend is automatically selected for DeepSeek-V3.2. Available kernels: `flashmla_sparse`, `flashmla_sparse_q8` (native FP8 e4m3 sparse prefill — no fp8→bf16 dequantization round-trip; Hopper SM90 + `--kv-cache-dtype fp8_e4m3` only, prefill only), `flashmla_kv`, `flashmla_auto`, `fa3` (Hopper only), `tilelang` (GPU/HPU/NPU), `aiter` (AMD, decode only), `trtllm` (Blackwell only). Defaults: Hopper BF16 KV → `flashmla_sparse` prefill / `fa3` decode; Hopper FP8 KV → `flashmla_kv` both; Blackwell BF16 → `flashmla_sparse` / `trtllm`; Blackwell FP8 → `trtllm` both.
- **Index Cache:** Reuses indexer results across layers for efficiency at negligible accuracy cost. For **GLM-5** specifically, append `--json-model-override-args '{"index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSFFFSFSSSFSFFSFFSSS"}'` for a better speed/accuracy tradeoff. - **Index Cache:** Reuses indexer results across layers for efficiency at negligible accuracy cost. For **GLM-5** specifically, append `--json-model-override-args '{"index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSFFFSFSSSFSFFSFFSSS"}'` for a better speed/accuracy tradeoff.
- **HiSparse (experimental):** Reduces per-request GPU memory during long-context decode by offloading KV data to CPU pinned memory. Requires PD disaggregation mode (decode instance only). See [HiSparse Guide](../../../docs/advanced_features/hisparse_guide). - **HiSparse (experimental):** Reduces per-request GPU memory during long-context decode by offloading KV data to CPU pinned memory. Requires PD disaggregation mode (decode instance only). See [HiSparse Guide](../../../docs/advanced_features/hisparse_guide).
- **NVFP4 on Blackwell:** Specify `--quantization modelopt_fp4` and `--moe-runner-backend flashinfer_trtllm` (recommended) / `flashinfer_cutlass` / `flashinfer_cutedsl`. Full example: - **NVFP4 on Blackwell:** Specify `--quantization modelopt_fp4` and `--moe-runner-backend flashinfer_trtllm` (recommended) / `flashinfer_cutlass` / `flashinfer_cutedsl`. Full example:
@@ -422,6 +422,12 @@ Internally, the DSA backend dispatches to different sub-backends for prefill and
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>✅</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>✅</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Default prefill on Hopper and Blackwell (BF16)</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Default prefill on Hopper and Blackwell (BF16)</td>
</tr> </tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>flashmla_sparse_q8</strong></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>✅</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>❌</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Native FP8 (q8×kv8) sparse prefill on Hopper (SM90); requires <code>--kv-cache-dtype fp8_e4m3</code></td>
</tr>
<tr> <tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>flashmla_kv</strong></td> <td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong>flashmla_kv</strong></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>✅</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>✅</td>
@@ -1408,7 +1408,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--dsa-prefill-backend`</td> <td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--dsa-prefill-backend`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DSA backend for the prefill stage (overrides `--attention-backend` when running DeepSeek DSA-style attention).</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DSA backend for the prefill stage (overrides `--attention-backend` when running DeepSeek DSA-style attention).</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Auto (hardware-dependent)</td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Auto (hardware-dependent)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>flashmla_sparse</code>, <code>flashmla_kv</code>, <code>flashmla_auto</code>, <code>fa3</code>, <code>tilelang</code>, <code>aiter</code>, <code>trtllm</code></td> <td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>flashmla_sparse</code>, <code>flashmla_sparse_q8</code>, <code>flashmla_kv</code>, <code>flashmla_auto</code>, <code>fa3</code>, <code>tilelang</code>, <code>aiter</code>, <code>trtllm</code></td>
</tr> </tr>
<tr> <tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--dsa-decode-backend`</td> <td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--dsa-decode-backend`</td>
@@ -14,11 +14,11 @@ limitations under the License.
==============================================================================*/ ==============================================================================*/
// SM90 FP8 native sparse MLA prefill kernel. // SM90 FP8 native sparse MLA prefill kernel.
//
// Algorithm inspired by DeepSeek FlashMLA // Algorithm inspired by DeepSeek FlashMLA
// (https://github.com/deepseek-ai/FlashMLA); the kernel itself is a // (https://github.com/deepseek-ai/FlashMLA); the kernel itself is a
// clean-room re-implementation targeting the Q8KV8 sparse prefill path. // clean-room re-implementation targeting the Q8KV8 sparse prefill path.
//
// Design: Native fp8 GMMA path // Design: Native fp8 GMMA path
// QK GEMM: fp8 SS (E4M3 x E4M3 -> F32, k=32, 2x throughput vs bf16) // QK GEMM: fp8 SS (E4M3 x E4M3 -> F32, k=32, 2x throughput vs bf16)
// PV GEMM: fp8 RS/SS (E4M3 x E4M3 -> F32, V physically transposed in smem) // PV GEMM: fp8 RS/SS (E4M3 x E4M3 -> F32, V physically transposed in smem)
@@ -138,14 +138,35 @@ struct SparseMlaQ8Kv8PrefillKernel {
array_aligned<bf16, cosize_v<SmemLayoutO>> o; // B_H * D_V/2 bf16 array_aligned<bf16, cosize_v<SmemLayoutO>> o; // B_H * D_V/2 bf16
} q_o; } q_o;
array_aligned<fp8_t, cosize_v<SmemLayoutK>> k[2]; // 2x K double-buffer, fp8 array_aligned<fp8_t, cosize_v<SmemLayoutK>> k[2]; // 2x K double-buffer, fp8
array_aligned<fp8_t, cosize_v<SmemLayoutVt>> vt[2]; // 2x Vt transposed buffer, fp8 // Vt is split into four half-buffers (256x64 each) so each half can be
// handed to its consumer warpgroup as soon as it is transposed:
// vt_loc[i]: i=0 block0-left (WG0 local), i=1 block1-right (WG1 local)
// vt_rem[i]: i=0 block1-left (WG0 remote), i=1 block0-right (WG1 remote)
array_aligned<fp8_t, cosize_v<SmemLayoutHalfVt>> vt_loc[2];
array_aligned<fp8_t, cosize_v<SmemLayoutHalfVt>> vt_rem[2];
array_aligned<fp8_t, 128 * 36> s[2]; // 2x S buffer, padded to a 36B row stride to avoid bank conflicts. array_aligned<fp8_t, 128 * 36> s[2]; // 2x S buffer, padded to a 36B row stride to avoid bank conflicts.
bool is_kv_valid[2][B_TOPK]; // double-buffer the VALIDITY DATA by iteration-pair
float2 sM[32]; // parity, matching bar_is_kv_valid_ready[2]. The consumer arrives
// bar_k_free BEFORE mask_rP reads these bits (perf: releases the producer's
// next K-load early), so a single-buffered array let the producer overwrite
// pair N+1's bits while a laggard consumer was still masking pair N ->
// wrong -INF pattern (nondeterministic corruption growing with CTA
// count). Dims: [k-buf/warpgroup][pair
// parity][slot]; producer ahead-ness is bounded to one pair by the
// bar_k_free protocol, so depth 2 suffices (same argument as the barrier).
bool is_kv_valid[2][2][B_TOPK];
// double-buffer the WG0<->WG1 running-max exchange
// by iteration-pair parity (same recipe as is_kv_valid): any +-1-iteration
// overrun of the producer side lands in the OTHER slot instead of
// overwriting a value the peer is still reading. Empirical signature of
// the race: single-row, ~half-heads rescale blow-up (peer_scale computed
// from the wrong iteration's max), vanishing on same-tensor retry.
float2 sM[2][32];
float2 sL[64]; float2 sL[64];
float final_max_logits[64], final_lse[64]; float final_max_logits[64], final_lse[64];
transac_bar_t bar_q, bar_k0_ready[2], bar_k1_ready[2], bar_is_kv_valid_ready; transac_bar_t bar_q, bar_k0_ready[2], bar_k1_ready[2],
bar_is_kv_valid_ready[2]; // double-buffer is_kv_valid_ready
transac_bar_t bar_k0_free, bar_k1_free; transac_bar_t bar_k0_free, bar_k1_free;
// Consumers arrive after PV drains; the producer waits before reusing the Vt buffer. // Consumers arrive after PV drains; the producer waits before reusing the Vt buffer.
// These barriers are separate from K-free so K buffers can be released earlier. // These barriers are separate from K-free so K buffers can be released earlier.
@@ -187,7 +208,10 @@ struct SparseMlaQ8Kv8PrefillKernel {
plan.bar_k0_ready[i].init(128); plan.bar_k0_ready[i].init(128);
plan.bar_k1_ready[i].init(128); plan.bar_k1_ready[i].init(128);
} }
plan.bar_is_kv_valid_ready.init(16); CUTE_UNROLL
for (int i = 0; i < 2; ++i) {
plan.bar_is_kv_valid_ready[i].init(16); // double-buffer
}
CUTE_UNROLL CUTE_UNROLL
for (int i = 0; i < 2; ++i) { for (int i = 0; i < 2; ++i) {
// Transaction barriers for Vt buffer safety: 128 arrivals from each consumer WG. // Transaction barriers for Vt buffer safety: 128 arrivals from each consumer WG.
@@ -215,8 +239,11 @@ struct SparseMlaQ8Kv8PrefillKernel {
const fp8_t* gQ = reinterpret_cast<const fp8_t*>(params.q) + s_q_idx * (int64_t)params.stride_q_s_q + const fp8_t* gQ = reinterpret_cast<const fp8_t*>(params.q) + s_q_idx * (int64_t)params.stride_q_s_q +
q_h_idx * B_H * (int64_t)params.stride_q_h_q; q_h_idx * B_H * (int64_t)params.stride_q_h_q;
// Vectorized Q loading via cp.async.cg (16 bytes per op) // Vectorized Q loading via cp.async.cg (16 bytes per op).
constexpr int Q_GROUP_SIZE = 8; // Group size 4 (not 8): with 8-row groups two warps' cp.async stores
// could overlap the same Q smem rows (WAW hazard); 4 keeps each row
// owned by exactly one group.
constexpr int Q_GROUP_SIZE = 4;
constexpr int Q_NUM_GROUPS = 128 / Q_GROUP_SIZE; constexpr int Q_NUM_GROUPS = 128 / Q_GROUP_SIZE;
constexpr int Q_ROWS_PER_GROUP = B_H / Q_NUM_GROUPS; constexpr int Q_ROWS_PER_GROUP = B_H / Q_NUM_GROUPS;
int q_ig = idx_in_warpgroup % Q_GROUP_SIZE; int q_ig = idx_in_warpgroup % Q_GROUP_SIZE;
@@ -248,6 +275,12 @@ struct SparseMlaQ8Kv8PrefillKernel {
// -------------------------------------------------------- // --------------------------------------------------------
float rM[2] = {MAX_INIT_VAL, MAX_INIT_VAL}; float rM[2] = {MAX_INIT_VAL, MAX_INIT_VAL};
float rL[2] = {0.0f, 0.0f}; float rL[2] = {0.0f, 0.0f};
// WG1 consumes WG0's peer P (s[0]) which is in WG0's LOCAL max frame
// (max(prev,block0)), not the combined frame. peer_scale = exp2(WG0_local - combined)
// brings it into the combined frame (<=1, ==1 at low magnitude => no-op). Clamped
// away from 0 so the 1/peer_scale pre-scale cannot overflow. WG0's peer (s[1]) is
// already combined-frame, so peer_scale stays 1.0 for WG0.
float peer_scale[2] = {1.0f, 1.0f};
Tensor rO = partition_fragment_C(TiledMMA_PV_LocalP{}, Shape<Int<B_H>, Int<D_V / 2>>{}); Tensor rO = partition_fragment_C(TiledMMA_PV_LocalP{}, Shape<Int<B_H>, Int<D_V / 2>>{});
Tensor rP = partition_fragment_C(TiledMMA_QK{}, Shape<Int<B_H>, Int<B_TOPK>>{}); Tensor rP = partition_fragment_C(TiledMMA_QK{}, Shape<Int<B_H>, Int<B_TOPK>>{});
cute::fill(rO, 0.0f); cute::fill(rO, 0.0f);
@@ -259,6 +292,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
Tensor rP_fp8_local = make_tensor<fp8_t>(rP_fp8_layout_t{}); Tensor rP_fp8_local = make_tensor<fp8_t>(rP_fp8_layout_t{});
bool cur_bar_wait_phase = 0; bool cur_bar_wait_phase = 0;
bool kv_valid_phase[2] = {false, false}; // per-instance phase for bar_is_kv_valid_ready[2]
struct Warpgroup0 {}; struct Warpgroup0 {};
struct Warpgroup1 {}; struct Warpgroup1 {};
@@ -272,28 +306,39 @@ struct SparseMlaQ8Kv8PrefillKernel {
gemm_ss(clear_accum, tiled_mma_QK, sQ_tile, sK_tile, rP, idx_in_warpgroup); gemm_ss(clear_accum, tiled_mma_QK, sQ_tile, sK_tile, rP, idx_in_warpgroup);
}; };
auto mask_rP = [&](auto wg_tag) { auto mask_rP = [&](auto wg_tag, int block_idx) {
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>; constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
plan.bar_is_kv_valid_ready.wait(cur_bar_wait_phase); // bar_is_kv_valid_ready was single-buffered but arrived EARLY by the
// producer and waited LATE here (after the prefetched QK gemm). The producer runs 1
// loop-iter ahead, so it can arrive the single-bit barrier TWICE before this wait
// consumes the first flip -> the parity aliases back and the laggard consumer waits for
// a flip that never comes (producer wedged at bar_vt_free) -> deadlock
// under load. Fix: double-buffer by iteration parity (producer is at most 1 loop-iter
// ahead => always the OTHER instance) with an independent per-instance phase.
int _kv_buf = (block_idx >> 1) & 1;
plan.bar_is_kv_valid_ready[_kv_buf].wait(kv_valid_phase[_kv_buf]);
kv_valid_phase[_kv_buf] ^= 1;
CUTE_UNROLL CUTE_UNROLL
for (int row_idx = 0; row_idx < 2; ++row_idx) { for (int row_idx = 0; row_idx < 2; ++row_idx) {
CUTE_UNROLL CUTE_UNROLL
for (int i = row_idx * 2; i < size(rP); i += 4) { for (int i = row_idx * 2; i < size(rP); i += 4) {
int col = 8 * (i / 4) + (idx_in_warpgroup % 4) * 2; int col = 8 * (i / 4) + (idx_in_warpgroup % 4) * 2;
if (!plan.is_kv_valid[IS_WG1][col]) rP(i) = -INFINITY; if (!plan.is_kv_valid[IS_WG1][_kv_buf][col]) rP(i) = -INFINITY; // parity-buffered
if (!plan.is_kv_valid[IS_WG1][col + 1]) rP(i + 1) = -INFINITY; if (!plan.is_kv_valid[IS_WG1][_kv_buf][col + 1]) rP(i + 1) = -INFINITY;
} }
} }
}; };
// online_softmax: compute softmax on rP (f32), then convert to fp8 // online_softmax: compute softmax on rP (f32), then convert to fp8
auto online_softmax_and_rescale_o = [&](auto wg_tag) { // _par = iteration-pair parity ((block_idx>>1)&1) selecting
// the sM slot for this iteration's WG0<->WG1 max exchange.
auto online_softmax_and_rescale_o = [&](auto wg_tag, int _par) {
// mask_rP already waits for the validity mask. // mask_rP already waits for the validity mask.
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>; constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
const float scale = qk_combined_scale_div_log2; const float scale = qk_combined_scale_div_log2;
float r_sM[2]; float r_sM[2];
if constexpr (IS_WG1) { if constexpr (IS_WG1) {
*(float2*)r_sM = plan.sM[idx_in_warpgroup / 4]; *(float2*)r_sM = plan.sM[_par][idx_in_warpgroup / 4];
} }
float new_maxs[2]; float new_maxs[2];
CUTE_UNROLL CUTE_UNROLL
@@ -307,6 +352,12 @@ struct SparseMlaQ8Kv8PrefillKernel {
cur_max = max(cur_max, __shfl_xor_sync(0xffffffff, cur_max, 2)); cur_max = max(cur_max, __shfl_xor_sync(0xffffffff, cur_max, 2));
cur_max *= scale; cur_max *= scale;
new_maxs[row_idx] = max(IS_WG1 ? r_sM[row_idx] : rM[row_idx], cur_max); new_maxs[row_idx] = max(IS_WG1 ? r_sM[row_idx] : rM[row_idx], cur_max);
// peer P (WG0's s[0]) is in WG0's local frame r_sM; bring it to the
// combined frame new_maxs. <=1, ==1 at low magnitude. Floor at 2^-30 so the
// 1/peer_scale pre-scale stays finite (peer weight there is negligible anyway).
if constexpr (IS_WG1) {
peer_scale[row_idx] = fmaxf(exp2f(r_sM[row_idx] - new_maxs[row_idx]), exp2f(-30.0f));
}
float scale_for_o = exp2f(rM[row_idx] - new_maxs[row_idx]); float scale_for_o = exp2f(rM[row_idx] - new_maxs[row_idx]);
CUTE_UNROLL CUTE_UNROLL
for (int i = row_idx * 2; i < size(rO); i += 4) { for (int i = row_idx * 2; i < size(rO); i += 4) {
@@ -326,7 +377,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
} }
__syncwarp(); __syncwarp();
if (idx_in_warpgroup % 4 == 0) { if (idx_in_warpgroup % 4 == 0) {
plan.sM[idx_in_warpgroup / 4] = *(float2*)new_maxs; plan.sM[_par][idx_in_warpgroup / 4] = *(float2*)new_maxs;
} }
rM[0] = new_maxs[0]; rM[0] = new_maxs[0];
rM[1] = new_maxs[1]; rM[1] = new_maxs[1];
@@ -441,32 +492,10 @@ struct SparseMlaQ8Kv8PrefillKernel {
} }
}; };
auto undo_v_transpose_col_permutation = [&]() { // No output un-permutation is needed here: the PV output columns are
// Undo the column permutation from the fp8 V transpose before writing O. // produced by the Vt B-operand, not by the permute_Cregs_fp8-rearranged P
// CLayout_64x256: col bit0 = t1_bit0 (thread), col bit3 = v1 (register). // A-operand (the permute only reorders the P contraction/K dimension,
// V transpose introduces bit0<->bit3 swap. Fix by cross-thread exchange: // never the output N).
// thread with t1_bit0=0, v1=1 <-> thread with t1_bit0=1, v1=0
// Within each 4-element group (same v2=row): idx%4 in {0,1} are v1=0, {2,3} are v1=1.
int t1_bit0 = (threadIdx.x >> 2) & 1;
#pragma unroll
for (int g = 0; g < 32; g++) {
float a = rO(4 * g + 0);
float b = rO(4 * g + 1);
float c = rO(4 * g + 2);
float d = rO(4 * g + 3);
float send0 = t1_bit0 ? a : c;
float send1 = t1_bit0 ? b : d;
float recv0 = __shfl_xor_sync(0xFFFFFFFF, send0, 4);
float recv1 = __shfl_xor_sync(0xFFFFFFFF, send1, 4);
if (t1_bit0 == 0) {
rO(4 * g + 2) = recv0;
rO(4 * g + 3) = recv1;
} else {
rO(4 * g + 0) = recv0;
rO(4 * g + 1) = recv1;
}
}
};
// ============================================================ // ============================================================
// WG0 Pipeline -- native fp8 // WG0 Pipeline -- native fp8
@@ -508,7 +537,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
CUTE_NO_UNROLL CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) { for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Vt[0] left half: (256, 64) fp8 -- only half we transpose & use // Vt[0] left half: (256, 64) fp8 -- only half we transpose & use
Tensor sVt0l = make_tensor(make_smem_ptr(plan.vt[0].data()), SmemLayoutHalfVt{}); Tensor sVt0l = make_tensor(make_smem_ptr(plan.vt_loc[0].data()), SmemLayoutHalfVt{});
if (block_idx == 0) { if (block_idx == 0) {
pipelined_wait_and_qkt_gemm_l(); pipelined_wait_and_qkt_gemm_l();
@@ -518,11 +547,16 @@ struct SparseMlaQ8Kv8PrefillKernel {
plan.bar_k0_free.arrive(); plan.bar_k0_free.arrive();
} }
mask_rP(Warpgroup0{}); mask_rP(Warpgroup0{}, block_idx); // pass block_idx for kv_buf parity
online_softmax_and_rescale_o(Warpgroup0{}); online_softmax_and_rescale_o(Warpgroup0{}, (block_idx >> 1) & 1);
save_rP_fp8_to_sS(plan.s[0].data()); save_rP_fp8_to_sS(plan.s[0].data());
NamedBarrier::arrive(256, NamedBarriers::wg0_bunch_0_ready); // was arrive-only: a +1-iteration overrun by WG0 could
// alias the named-barrier count against WG1's pending wait. Full
// rendezvous kills the aliasing; WG1's path here (mask_rP -> its
// bar_is_kv_valid wait) does not depend on anything WG0 does after
// this point, so no deadlock surface is added.
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg0_bunch_0_ready);
// Wait for Vt[0] left half only (producer + WG0 arrivals). // Wait for Vt[0] left half only (producer + WG0 arrivals).
// V[0]-RIGHT may still be transposing; WG0 doesn't need it. // V[0]-RIGHT may still be transposing; WG0 doesn't need it.
@@ -535,7 +569,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
// Overlap PV-local GMMA drain with barrier waits, sM read, and peer P load. // Overlap PV-local GMMA drain with barrier waits, sM read, and peer P load.
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg1_bunch_0_ready); NamedBarrier::arrive_and_wait(256, NamedBarriers::wg1_bunch_0_ready);
float new_rM[2], scale_factors_arr[2]; float new_rM[2], scale_factors_arr[2];
*(float2*)new_rM = plan.sM[idx_in_warpgroup / 4]; *(float2*)new_rM = plan.sM[(block_idx >> 1) & 1][idx_in_warpgroup / 4];
CUTE_UNROLL CUTE_UNROLL
for (int i = 0; i < 2; ++i) { for (int i = 0; i < 2; ++i) {
scale_factors_arr[i] = exp2f(rM[i] - new_rM[i]); scale_factors_arr[i] = exp2f(rM[i] - new_rM[i]);
@@ -556,7 +590,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
// Rescale rO: must be after wait<0> since rO is PV-local accumulator // Rescale rO: must be after wait<0> since rO is PV-local accumulator
rescale_rO(scale_factors_arr); rescale_rO(scale_factors_arr);
Tensor sVt1l = make_tensor(make_smem_ptr(plan.vt[1].data()), SmemLayoutHalfVt{}); Tensor sVt1l = make_tensor(make_smem_ptr(plan.vt_rem[0].data()), SmemLayoutHalfVt{});
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt1l, rO, idx_in_warpgroup); gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt1l, rO, idx_in_warpgroup);
warpgroup_commit_batch(); warpgroup_commit_batch();
@@ -579,8 +613,6 @@ struct SparseMlaQ8Kv8PrefillKernel {
} }
} }
undo_v_transpose_col_permutation();
reduce_L(); reduce_L();
store_O(); store_O();
@@ -614,7 +646,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
CUTE_NO_UNROLL CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) { for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Vt[1] right half: (256, 64) fp8 -- only half we transpose & use // Vt[1] right half: (256, 64) fp8 -- only half we transpose & use
Tensor sVt1r = make_tensor(make_smem_ptr(plan.vt[1].data() + 256 * B_TOPK), SmemLayoutHalfVt{}); Tensor sVt1r = make_tensor(make_smem_ptr(plan.vt_loc[1].data()), SmemLayoutHalfVt{});
if (block_idx == 0) { if (block_idx == 0) {
pipelined_wait_and_qkt_gemm_r_wg1(); pipelined_wait_and_qkt_gemm_r_wg1();
@@ -624,13 +656,20 @@ struct SparseMlaQ8Kv8PrefillKernel {
plan.bar_k1_free.arrive(); plan.bar_k1_free.arrive();
} }
mask_rP(Warpgroup1{}); mask_rP(Warpgroup1{}, block_idx); // pass block_idx for kv_buf parity
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg0_bunch_0_ready); NamedBarrier::arrive_and_wait(256, NamedBarriers::wg0_bunch_0_ready);
online_softmax_and_rescale_o(Warpgroup1{}); online_softmax_and_rescale_o(Warpgroup1{}, (block_idx >> 1) & 1);
save_rP_fp8_to_sS(plan.s[1].data()); save_rP_fp8_to_sS(plan.s[1].data());
NamedBarrier::arrive(256, NamedBarriers::wg1_bunch_0_ready); // was arrive-only — the mirror of the wg0_bunch case: if WG1 runs a
// full iteration ahead, its second arrive aliases the named-barrier
// count against WG0's pending wait at this rendezvous, so WG0 releases
// early and reads the combined max (sM) BEFORE WG1 wrote it -> wrong
// rescale factors on one CTA's 64-head half. No deadlock surface:
// WG0 reaches its wait via wg0_bunch (which WG1 arrives earlier) plus
// producer vt0_left; neither depends on WG1's progress past here.
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg1_bunch_0_ready);
// Wait for Vt[1] transpose (prod+WG1 barrier) // Wait for Vt[1] transpose (prod+WG1 barrier)
NamedBarrier::arrive_and_wait(256, vt1_for_wg1); NamedBarrier::arrive_and_wait(256, vt1_for_wg1);
@@ -643,6 +682,18 @@ struct SparseMlaQ8Kv8PrefillKernel {
warpgroup_fence_operand(rO); warpgroup_fence_operand(rO);
warpgroup_fence_operand(rP_fp8_local); warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[1].arrive(); plan.bar_vt_free[1].arrive();
// pre-scale rO (= local block1, combined frame) by 1/peer_scale so the
// peer PV (WG0's block0 P in WG0's local frame) lands correct after the post-scale:
// rO/peer_scale; rO += peer_PV; rO *= peer_scale => local + peer_scale*peer_PV.
CUTE_UNROLL
for (int row = 0; row < 2; ++row) {
float inv = 1.0f / peer_scale[row];
CUTE_UNROLL
for (int i = row * 2; i < size(rO); i += 4) {
rO(i) *= inv;
rO(i + 1) *= inv;
}
}
load_sS_to_rP(plan.s[0].data()); load_sS_to_rP(plan.s[0].data());
NamedBarrier::arrive_and_wait(256, s_consumed_ready); NamedBarrier::arrive_and_wait(256, s_consumed_ready);
@@ -650,7 +701,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
// V[0]-LEFT was signaled earlier; WG1 doesn't need it. // V[0]-LEFT was signaled earlier; WG1 doesn't need it.
NamedBarrier::arrive_and_wait(256, vt0_right_ready); NamedBarrier::arrive_and_wait(256, vt0_right_ready);
Tensor sVt0r = make_tensor(make_smem_ptr(plan.vt[0].data() + 256 * B_TOPK), SmemLayoutHalfVt{}); Tensor sVt0r = make_tensor(make_smem_ptr(plan.vt_rem[1].data()), SmemLayoutHalfVt{});
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt0r, rO, idx_in_warpgroup); gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt0r, rO, idx_in_warpgroup);
warpgroup_commit_batch(); warpgroup_commit_batch();
@@ -658,8 +709,17 @@ struct SparseMlaQ8Kv8PrefillKernel {
cur_bar_wait_phase ^= 1; cur_bar_wait_phase ^= 1;
// Overlap: start next-iteration QK-right while PV drains // Overlap: start next-iteration QK-right while PV drains
pipelined_wait_and_qkt_gemm_r_wg1(); pipelined_wait_and_qkt_gemm_r_wg1();
warpgroup_wait<1>(); warpgroup_wait<1>(); // drains the peer PV (committed before the QK-right batch)
warpgroup_fence_operand(rO); warpgroup_fence_operand(rO);
// post-scale: rO = local + peer_scale * peer_PV (combined frame).
CUTE_UNROLL
for (int row = 0; row < 2; ++row) {
CUTE_UNROLL
for (int i = row * 2; i < size(rO); i += 4) {
rO(i) *= peer_scale[row];
rO(i + 1) *= peer_scale[row];
}
}
warpgroup_fence_operand(rP_fp8_local); warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[0].arrive(); plan.bar_vt_free[0].arrive();
pipelined_wait_and_qkt_gemm_l_wg1(); pipelined_wait_and_qkt_gemm_l_wg1();
@@ -667,14 +727,21 @@ struct SparseMlaQ8Kv8PrefillKernel {
warpgroup_fence_operand(rP); warpgroup_fence_operand(rP);
plan.bar_k1_free.arrive(); plan.bar_k1_free.arrive();
} else { } else {
warpgroup_wait<0>(); warpgroup_wait<0>(); // drains the peer PV
warpgroup_fence_operand(rO); warpgroup_fence_operand(rO);
// post-scale (final iteration).
CUTE_UNROLL
for (int row = 0; row < 2; ++row) {
CUTE_UNROLL
for (int i = row * 2; i < size(rO); i += 4) {
rO(i) *= peer_scale[row];
rO(i + 1) *= peer_scale[row];
}
}
plan.bar_vt_free[0].arrive(); plan.bar_vt_free[0].arrive();
} }
} }
undo_v_transpose_col_permutation();
reduce_L(); reduce_L();
store_O(); store_O();
@@ -720,6 +787,9 @@ struct SparseMlaQ8Kv8PrefillKernel {
int64_t token_indices[2][NUM_ROWS_PER_GROUP]; int64_t token_indices[2][NUM_ROWS_PER_GROUP];
bool is_token_valid[2][NUM_ROWS_PER_GROUP]; bool is_token_valid[2][NUM_ROWS_PER_GROUP];
// Hoisted invariant: base of the `topk` trailing zero pad
// rows in kv. -1 sentinels map to (pad_base + slot) = distinct zero rows.
const int pad_base = params.s_kv - params.topk;
auto load_token_indices = [&](int block_idx) { auto load_token_indices = [&](int block_idx) {
CUTE_UNROLL CUTE_UNROLL
for (int buf_idx = 0; buf_idx < 2; ++buf_idx) { for (int buf_idx = 0; buf_idx < 2; ++buf_idx) {
@@ -727,9 +797,31 @@ struct SparseMlaQ8Kv8PrefillKernel {
for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row) { for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row) {
int offs = (block_idx + buf_idx) * B_TOPK + local_row * NUM_GROUPS + group_idx; int offs = (block_idx + buf_idx) * B_TOPK + local_row * NUM_GROUPS + group_idx;
int t = __ldg(gIndices + offs); int t = __ldg(gIndices + offs);
bool is_cur_token_valid = t >= 0 && t < params.s_kv; bool is_cur_token_valid;
if constexpr (HAVE_TOPK_LENGTH) { if constexpr (!HAVE_TOPK_LENGTH) {
is_cur_token_valid &= offs < topk_length; // Map -1 sentinels to DISTINCT zero pad rows (slot `offs` ->
// pad_base+offs, distinct within each query) -> avoids the
// duplicate-index kernel slowdown while keeping uniform full-topk
// loads (data-independent, no DP hang). Replaces the per-layer
// torch.where in the integration (eliminates an elementwise launch).
// After clamping, t in [0, s_kv) by construction (offs < topk), so
// the LOAD is always safe (real or zero-pad row).
// but the pad slots must be MASKED in the
// softmax (is_valid=false -> -INF in mask_rP), NOT scored: a zero
// KV row contributes exp(0 - max) to the denominator, which for
// few-valid rows (ctx << topk, e.g. the first tokens of a prompt)
// crushes the output by up to ~2048x. chunk<=16384 never exposed
// this because rows with ctx<=topk take the DENSE prefill path
// (per-rank chunk 2048 = topk); chunk32768 packs them into the
// sparse call, exposing the bug. Control
// flow is UNCHANGED (uniform full-topk loads, data-independent;
// only the validity bit differs -> no-hang properties preserved).
const bool t_is_pad = (t < 0);
t = t_is_pad ? (pad_base + offs) : t;
is_cur_token_valid = !t_is_pad;
} else {
is_cur_token_valid = (t >= 0 && t < params.s_kv) && (offs < topk_length);
} }
token_indices[buf_idx][local_row] = (int64_t)t * (int64_t)params.stride_kv_s_kv; token_indices[buf_idx][local_row] = (int64_t)t * (int64_t)params.stride_kv_s_kv;
is_token_valid[buf_idx][local_row] = is_cur_token_valid; is_token_valid[buf_idx][local_row] = is_cur_token_valid;
@@ -766,23 +858,28 @@ struct SparseMlaQ8Kv8PrefillKernel {
SmemTransposeV smem_transpose_v; SmemTransposeV smem_transpose_v;
using SmemLayoutTransposeV_t = typename SmemTransposeV::SmemLayoutTransposeV; using SmemLayoutTransposeV_t = typename SmemTransposeV::SmemLayoutTransposeV;
using SmemLayoutTransposeVt_t = typename SmemTransposeV::SmemLayoutTransposeVt; using SmemLayoutTransposeVt_t = typename SmemTransposeV::SmemLayoutTransposeVt;
// Half (256x64) Vt dst layout for the split vt_loc/vt_rem half-buffers.
// Source stays the FULL [64,512] K layout (tile index j handles the swizzle); only the DST
// is a half-buffer, tiles remapped tile_start..end -> 0..3. transpose_pair is tile-agnostic
// (operates on 64x64 tiles) so smem_transpose_v is reused.
using SmemTransposeV_half_t = SmemTransposeFp8_64x64<B_TOPK, D_V / 2>;
using SmemLayoutTransposeVt_half_t = typename SmemTransposeV_half_t::SmemLayoutTransposeVt;
// Use the FA3-style STSM thread layout for the fp8 V transpose. // Use the FA3-style STSM thread layout for the fp8 V transpose.
// but same composition-based framework as before. auto transpose_v_to_half = [&](int smem_k_buf, fp8_t* vt_dst, int tile_start, int tile_end) {
auto transpose_v_half = [&](int smem_k_buf, int vt_buf, int tile_start, int tile_end) {
Tensor sV_src = as_position_independent_swizzle_tensor( Tensor sV_src = as_position_independent_swizzle_tensor(
make_tensor(make_smem_ptr(plan.k[smem_k_buf].data()), SmemLayoutTransposeV_t{})); make_tensor(make_smem_ptr(plan.k[smem_k_buf].data()), SmemLayoutTransposeV_t{}));
Tensor sVt_dst = as_position_independent_swizzle_tensor( Tensor sVt_dst =
make_tensor(make_smem_ptr(plan.vt[vt_buf].data()), SmemLayoutTransposeVt_t{})); as_position_independent_swizzle_tensor(make_tensor(make_smem_ptr(vt_dst), SmemLayoutTransposeVt_half_t{}));
static_assert((D_V / 64 / 2) % 2 == 0, "half tile count must be even for pair transpose"); static_assert((D_V / 64 / 2) % 2 == 0, "half tile count must be even for pair transpose");
CUTE_UNROLL CUTE_UNROLL
for (int j = tile_start; j < tile_end; j += 2) { for (int j = tile_start; j < tile_end; j += 2) {
smem_transpose_v.transpose_pair( smem_transpose_v.transpose_pair(
flatten(sV_src(_, 0, j)), flatten(sV_src(_, 0, j)),
flatten(sVt_dst(_, 0, j)), flatten(sVt_dst(_, 0, j - tile_start)),
flatten(sV_src(_, 0, j + 1)), flatten(sV_src(_, 0, j + 1)),
flatten(sVt_dst(_, 0, j + 1))); flatten(sVt_dst(_, 0, j + 1 - tile_start)));
} }
asm volatile("" ::: "memory"); asm volatile("" ::: "memory");
}; };
@@ -797,7 +894,6 @@ struct SparseMlaQ8Kv8PrefillKernel {
CUTE_NO_UNROLL CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) { for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Indices are already loaded by the prologue or the previous iteration's prefetch. // Indices are already loaded by the prologue or the previous iteration's prefetch.
plan.bar_k0_free.wait(cur_bar_wait_phase_prod); plan.bar_k0_free.wait(cur_bar_wait_phase_prod);
plan.bar_k1_free.wait(cur_bar_wait_phase_prod); plan.bar_k1_free.wait(cur_bar_wait_phase_prod);
@@ -810,8 +906,9 @@ struct SparseMlaQ8Kv8PrefillKernel {
for (int buf_idx = 0; buf_idx < 2; ++buf_idx) for (int buf_idx = 0; buf_idx < 2; ++buf_idx)
CUTE_UNROLL CUTE_UNROLL
for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row) for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row)
plan.is_kv_valid[buf_idx][local_row * NUM_GROUPS + group_idx] = is_token_valid[buf_idx][local_row]; plan.is_kv_valid[buf_idx][(block_idx >> 1) & 1][local_row * NUM_GROUPS + group_idx] =
plan.bar_is_kv_valid_ready.arrive(); is_token_valid[buf_idx][local_row]; // parity-buffered
plan.bar_is_kv_valid_ready[(block_idx >> 1) & 1].arrive(); // index by iter parity
} }
copy_tiles(0, 0, 0, 4); copy_tiles(0, 0, 0, 4);
@@ -830,7 +927,9 @@ struct SparseMlaQ8Kv8PrefillKernel {
commit_to_mbar(plan.bar_k1_ready[0]); commit_to_mbar(plan.bar_k1_ready[0]);
asm volatile("cp.async.commit_group;\n" ::); asm volatile("cp.async.commit_group;\n" ::);
// Wait for K[0]-left (group-0) // Wait for K[0]-left. NOTE: keep cp.async.wait_group at 1, not 0 -
// wait_group 0 re-times the producer into a latent cross-iteration
// race; the consumer must overlap the in-flight K[1] load here.
asm volatile("cp.async.wait_group 1;\n" ::); asm volatile("cp.async.wait_group 1;\n" ::);
// fence.proxy.async: make cp.async data visible through generic proxy // fence.proxy.async: make cp.async data visible through generic proxy
// (required for LDSM reads in V transpose; cp.async uses async proxy) // (required for LDSM reads in V transpose; cp.async uses async proxy)
@@ -847,16 +946,19 @@ struct SparseMlaQ8Kv8PrefillKernel {
load_token_indices(block_idx + 2); load_token_indices(block_idx + 2);
} }
transpose_v_half(0, 0, 0, 4); transpose_v_to_half(0, plan.vt_loc[0].data(), 0, 4); // block0-left -> WG0 local
NamedBarrier::arrive(256, vt0_left_ready); // The arrive is placed below the wait_group-0 + fence point: the
// fence_view_async_shared there also orders these STSM (generic proxy)
// writes for the consumer's WGMMA (async proxy) reads. Without the
// fence the STSM->WGMMA proxy crossing is unordered and corrupts values.
// Transpose V[1] left before V[0] right to match the consumer handoff order. // Transpose V[1] left before V[0] right to match the consumer handoff order.
// WG0 is on the critical path (feeds WG1 via sM/wg0_bunch). // WG0 is on the critical path (feeds WG1 via sM/wg0_bunch).
// WG0 waits for vt1_for_wg0 (V[1]-LEFT) for PV-remote. // WG0 waits for vt1_for_wg0 (V[1]-LEFT) for PV-remote.
// Moving V[1]-LEFT earlier (2nd instead of 4th) reduces WG0 // Moving V[1]-LEFT earlier (2nd instead of 4th) reduces WG0
// critical-path stall by ~768 cycles per iteration. // critical-path stall.
//
// v52 CRASH FIX: K[1]-left tiles 0-3 are in cp.async group-1, // NOTE: K[1]-left tiles 0-3 are in cp.async group-1,
// NOT group-0. wait_group 1 only waits for group-0. Under high // NOT group-0. wait_group 1 only waits for group-0. Under high
// CTA counts (512+), memory bandwidth saturation delays group-1 // CTA counts (512+), memory bandwidth saturation delays group-1
// completion past the V[0]-LEFT transpose timing margin, causing // completion past the V[0]-LEFT transpose timing margin, causing
@@ -870,20 +972,25 @@ struct SparseMlaQ8Kv8PrefillKernel {
// generic proxy for LDSM reads in V transpose // generic proxy for LDSM reads in V transpose
fence_view_async_shared(); fence_view_async_shared();
asm volatile("bar.sync 7, 128;\n" ::: "memory"); asm volatile("bar.sync 7, 128;\n" ::: "memory");
NamedBarrier::arrive(256, vt0_left_ready); // covered by the fence above
// V[1]-LEFT: tiles 0-3 from K[1] -- WG0 needs this for PV-remote // V[1]-LEFT: tiles 0-3 from K[1] -- WG0 needs this for PV-remote
if (block_idx > 0) { if (block_idx > 0) {
plan.bar_vt_free[1].wait(cur_bar_wait_phase_prod); plan.bar_vt_free[1].wait(cur_bar_wait_phase_prod);
} }
transpose_v_half(1, 1, 0, 4); transpose_v_to_half(1, plan.vt_rem[0].data(), 0, 4); // block1-left -> WG0 remote
// The fence is cheap here: cp.async queue already drained by
// wait_group 0 above; orders V[1]L STSM writes before the arrive.
fence_view_async_shared();
NamedBarrier::arrive(256, vt1_for_wg0); NamedBarrier::arrive(256, vt1_for_wg0);
// V[0]-RIGHT: tiles 4-7 from K[0] // V[0]-RIGHT: tiles 4-7 from K[0]
transpose_v_half(0, 0, 4, 8); transpose_v_to_half(0, plan.vt_rem[1].data(), 4, 8); // block0-right -> WG1 remote
NamedBarrier::arrive(256, vt0_right_ready);
// V[1]-RIGHT: tiles 4-7 from K[1] // V[1]-RIGHT: tiles 4-7 from K[1]
transpose_v_half(1, 1, 4, 8); transpose_v_to_half(1, plan.vt_loc[1].data(), 4, 8); // block1-right -> WG1 local
// One fence covers both right-half transposes, then arrive both.
fence_view_async_shared();
NamedBarrier::arrive(256, vt0_right_ready);
NamedBarrier::arrive(256, vt1_for_wg1); NamedBarrier::arrive(256, vt1_for_wg1);
asm volatile("bar.sync 7, 128;\n" ::: "memory"); asm volatile("bar.sync 7, 128;\n" ::: "memory");
@@ -285,5 +285,137 @@ def _dequantize_k_cache_paged_kernel(
tl.store(dst_ptr, data, mask=mask) tl.store(dst_ptr, data, mask=mask)
def gather_dequant_requant_fp8_paged(
quant_k_cache: torch.Tensor,
page_table_1_flattened: torch.Tensor,
group_size: int = 128,
extra_rows: int = 0,
) -> torch.Tensor:
"""Gather paged fp8 KV tokens and re-pack into flat [576] fp8 layout.
The paged KV cache stores 656 bytes per token:
[512 nope_fp8 | 16 scales_f32 (4 groups) | 128 rope_bf16_bytes]
This kernel gathers the requested tokens, de-quantises nope with the
per-group scales, and re-quantises to per-tensor fp8 (scale=1.0).
Rope is cast bf16->fp8. The whole operation is fused into a single
Triton kernel to avoid allocating an intermediate bf16 buffer.
Args:
quant_k_cache: [total_num_tokens, 1, 656] fp8_e4m3fn
page_table_1_flattened: [num_tokens] int32
group_size: per-group dequant tile size (default 128)
extra_rows: number of zero-filled landing-pad rows to append at
the end of the output (used by the SM90 sparse MLA Q8KV8
kernel which over-reads past end-of-buffer for masked
indices)
Returns:
output: [num_tokens + extra_rows, 1, 576] fp8_e4m3fn
"""
dim_quant = quant_k_cache.shape[-1]
assert dim_quant == 656
quant_k_cache = quant_k_cache.view((-1, dim_quant))
num_tokens = page_table_1_flattened.shape[0]
assert quant_k_cache.dtype == torch.float8_e4m3fn
dim_nope = 512
dim_rope = 64
num_tiles = dim_nope // group_size # 4
out_dim = dim_nope + dim_rope # 576
assert num_tiles * group_size == dim_nope
total_rows = num_tokens + extra_rows
# Allocate a fresh zero-filled buffer. The extra landing-pad rows at
# the tail must read as zeros (the kernel may over-read past
# num_tokens for masked indices). A future optimization could cache
# this buffer but baseline allocates fresh.
output = torch.zeros(
(total_rows, 1, out_dim),
dtype=torch.float8_e4m3fn,
device=quant_k_cache.device,
)
num_blocks_per_token = triton.cdiv(dim_nope + dim_rope, group_size) # 5
assert num_blocks_per_token == 5
input_nope_q = quant_k_cache[:, :dim_nope]
input_nope_s = quant_k_cache[:, dim_nope : dim_nope + num_tiles * 4].view(
torch.float32
)
input_rope = quant_k_cache[:, dim_nope + num_tiles * 4 :].view(torch.bfloat16)
_gather_dequant_requant_fp8_paged_kernel[(num_tokens, num_blocks_per_token)](
output,
input_nope_q,
input_nope_s,
input_rope,
page_table_1_flattened,
output.stride(0),
input_nope_q.stride(0),
input_nope_s.stride(0),
input_rope.stride(0),
NUM_NOPE_BLOCKS=num_tiles,
GROUP_SIZE=group_size,
DIM_NOPE=dim_nope,
DIM_ROPE=dim_rope,
)
return output
@triton.jit
def _gather_dequant_requant_fp8_paged_kernel(
output_ptr,
input_nope_q_ptr,
input_nope_s_ptr,
input_rope_ptr,
page_table_1_ptr,
output_stride_0: int,
input_nope_q_stride_0: int,
input_nope_s_stride_0: int,
input_rope_stride_0: int,
NUM_NOPE_BLOCKS: tl.constexpr,
GROUP_SIZE: tl.constexpr,
DIM_NOPE: tl.constexpr,
DIM_ROPE: tl.constexpr,
):
"""Fused gather + dequant(per-group) + requant(per-tensor) -> fp8."""
token_id = tl.program_id(0)
token_id_paged = tl.load(page_table_1_ptr + token_id).to(tl.int32)
raw_block_id = tl.program_id(1)
if raw_block_id < NUM_NOPE_BLOCKS:
# nope: read fp8, mul group scale -> f32, cast to fp8_e4m3fn
effective_block_id = raw_block_id
offs_q = effective_block_id * GROUP_SIZE + tl.arange(0, GROUP_SIZE)
mask = offs_q < DIM_NOPE
ptr_q = input_nope_q_ptr + token_id_paged * input_nope_q_stride_0 + offs_q
ptr_s = (
input_nope_s_ptr
+ token_id_paged * input_nope_s_stride_0
+ effective_block_id
)
y_q = tl.load(ptr_q, mask=mask, other=0.0).to(tl.float32)
y_s = tl.load(ptr_s)
# dequant -> f32 -> requant to fp8
y = (y_q * y_s).to(tl.float8e4nv)
dst_ptr = output_ptr + token_id * output_stride_0 + offs_q
tl.store(dst_ptr, y, mask=mask)
else:
# rope: read bf16, cast to fp8
effective_block_id = raw_block_id - NUM_NOPE_BLOCKS
offs = effective_block_id * GROUP_SIZE + tl.arange(0, GROUP_SIZE)
mask = offs < DIM_ROPE
src_ptr = input_rope_ptr + token_id_paged * input_rope_stride_0 + offs
dst_ptr = output_ptr + token_id * output_stride_0 + DIM_NOPE + offs
data = tl.load(src_ptr, mask=mask).to(tl.float8e4nv)
tl.store(dst_ptr, data, mask=mask)
if __name__ == "__main__": if __name__ == "__main__":
raise Exception("UT is in quant_k_cache.py") raise Exception("UT is in quant_k_cache.py")
@@ -21,6 +21,12 @@ from sglang.kernels.ops.kvcache.cache_ops import (
from sglang.kernels.ops.kvcache.cache_ops import ( from sglang.kernels.ops.kvcache.cache_ops import (
concat_and_cast_mha_k_triton as concat_and_cast_mha_k_triton, concat_and_cast_mha_k_triton as concat_and_cast_mha_k_triton,
) )
from sglang.kernels.ops.kvcache.cache_ops import (
concat_and_cast_q_fp8_pad as concat_and_cast_q_fp8_pad,
)
from sglang.kernels.ops.kvcache.cache_ops import (
concat_and_cast_q_fp8_pad_kernel as concat_and_cast_q_fp8_pad_kernel,
)
from sglang.kernels.ops.kvcache.cache_ops import ( from sglang.kernels.ops.kvcache.cache_ops import (
launch_reshape_and_cache_flash as launch_reshape_and_cache_flash, launch_reshape_and_cache_flash as launch_reshape_and_cache_flash,
) )
@@ -264,3 +264,70 @@ def launch_reshape_and_cache_flash(
HAS_SWA=(swa_slot_mapping is not None), HAS_SWA=(swa_slot_mapping is not None),
USE_SCALE=(k_scale is not None), USE_SCALE=(k_scale is not None),
) )
@triton.jit
def concat_and_cast_q_fp8_pad_kernel(
qpad_ptr, # [num_tokens, pad_heads, NOPE+ROPE] fp8 (dst; only [:, :H, :] written)
q_nope_ptr, # [num_tokens, H, NOPE] bf16
q_rope_ptr, # [num_tokens, H, ROPE] bf16
qpad_s0,
qpad_s1,
nope_s0,
nope_s1,
rope_s0,
rope_s1,
H: tl.constexpr,
NOPE: tl.constexpr,
ROPE: tl.constexpr,
):
# One program per token: write the H active heads of the padded fp8 q buffer,
# fusing the bf16->fp8 cast (on store) with the nope/rope concat. Bit-exact vs the
# two strided copy_() it replaces; ~3.7x faster because copy_ into the
# 64-head-padded buffer is strided (~4.5x off memory-bound). Strides are passed in,
# so q_nope/q_rope may be views of a [T, H, NOPE+ROPE] q (head-stride != last-dim).
pid = tl.program_id(0)
hr = tl.arange(0, H)
qpad_head = qpad_ptr + pid * qpad_s0 + hr[:, None] * qpad_s1
no = tl.arange(0, NOPE)
src_n = tl.load(q_nope_ptr + pid * nope_s0 + hr[:, None] * nope_s1 + no[None, :])
tl.store(qpad_head + no[None, :], src_n)
ro = tl.arange(0, ROPE)
src_r = tl.load(q_rope_ptr + pid * rope_s0 + hr[:, None] * rope_s1 + ro[None, :])
tl.store(qpad_head + NOPE + ro[None, :], src_r)
def concat_and_cast_q_fp8_pad(q_fp8_pad, q_nope, q_rope, num_heads):
"""fused bf16->fp8 concat-cast of q_nope/q_rope into the active
[:, :num_heads, :] slice of the padded fp8 q buffer. Bit-exact replacement for the
two strided converting copy_() in the Q8KV8 prefill q-prep, ~3.7x faster. Requires
num_heads / nope_dim / rope_dim to be powers of two (always true for DeepSeek: 128
heads / any TP, 512 nope, 64 rope)."""
num_tokens = q_nope.shape[0]
nope_dim = q_nope.shape[-1]
rope_dim = q_rope.shape[-1]
concat_and_cast_q_fp8_pad_kernel[(num_tokens,)](
q_fp8_pad,
q_nope,
q_rope,
q_fp8_pad.stride(0),
q_fp8_pad.stride(1),
q_nope.stride(0),
q_nope.stride(1),
q_rope.stride(0),
q_rope.stride(1),
H=num_heads,
NOPE=nope_dim,
ROPE=rope_dim,
)
# ---------------------------------------------------------------------------
# Decode Context Parallel (DCP) helpers.
#
# Not part of upstream main (PR #26000 centralized the other Triton utility
# kernels into triton_ops/*). These three live here because they are DCP-only:
# - create_triton_kv_indices_for_dcp_triton: per-rank local KV indices
# - get_dcp_lens: per-rank visible KV length
# - cp_lse_ag_out_rs: merge DCP partial attention via natural-log LSE
# ---------------------------------------------------------------------------
@@ -18,7 +18,10 @@ from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.kernels.ops.attention.dsa.dequant_k_cache import dequantize_k_cache_paged from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
dequantize_k_cache_paged,
gather_dequant_requant_fp8_paged,
)
from sglang.kernels.ops.attention.dsa.quant_k_cache import quantize_k_cache from sglang.kernels.ops.attention.dsa.quant_k_cache import quantize_k_cache
from sglang.kernels.ops.attention.dsa.transform_index import ( from sglang.kernels.ops.attention.dsa.transform_index import (
transform_index_page_table_decode, transform_index_page_table_decode,
@@ -29,6 +32,7 @@ from sglang.kernels.ops.attention.utils import (
mla_quantize_and_rope_for_fp8, mla_quantize_and_rope_for_fp8,
seqlens_expand_triton, seqlens_expand_triton,
) )
from sglang.kernels.ops.kvcache.cache_ops import concat_and_cast_q_fp8_pad
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import (
@@ -323,7 +327,7 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
_DSA_IMPL_T: TypeAlias = Literal[ _DSA_IMPL_T: TypeAlias = Literal[
"flashmla_sparse", "flashmla_kv", "fa3", "tilelang", "trtllm" "flashmla_sparse", "flashmla_sparse_q8", "flashmla_kv", "fa3", "tilelang", "trtllm"
] ]
@@ -448,6 +452,40 @@ class DeepseekSparseAttnBackend(
self.device_sm_major = self.device_capability[0] self.device_sm_major = self.device_capability[0]
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
# `flashmla_sparse_q8` = the native FP8 SM90 sparse-prefill kernel. It always
# runs FP8 (requires fp8_e4m3 KV) and is SM90-only, so validate both at
# construction: an unsupported config must fail at launch rather than
# mid-forward. `flashmla_sparse` remains the bf16 path with no such
# requirement.
if self.dsa_prefill_impl == "flashmla_sparse_q8":
if self.kv_cache_dtype != torch.float8_e4m3fn:
raise ValueError(
"--dsa-prefill-backend flashmla_sparse_q8 is native FP8 and requires "
f"--kv-cache-dtype fp8_e4m3 (got kv_cache_dtype={self.kv_cache_dtype}); "
"use --dsa-prefill-backend flashmla_sparse for the bf16 path."
)
if self.device_sm_major != 9:
raise ValueError(
"--dsa-prefill-backend flashmla_sparse_q8 is SM90-only; got compute "
f"capability sm_{self.device_sm_major}x."
)
# `flashmla_sparse_q8` is prefill-only (FP8 decode goes through
# `flashmla_kv`); reject it as a decode backend, since argparse accepts it
# via the shared DSA_CHOICES list.
if self.dsa_decode_impl == "flashmla_sparse_q8":
raise ValueError(
"--dsa-decode-backend flashmla_sparse_q8 is not supported: "
"flashmla_sparse_q8 is a prefill-only backend. For FP8, use "
"--dsa-prefill-backend flashmla_sparse_q8 together with "
"--dsa-decode-backend flashmla_kv."
)
# Q8KV8 per-call device-tensor caches, populated lazily on the first
# Q8KV8 dispatch (no-ops for other backends).
self._q8kv8_identity_scale: Optional[torch.Tensor] = None
self._q8kv8_qpad_buf: Optional[torch.Tensor] = None
# Allocate global workspace buffer for TRT-LLM kernels (ragged attention on SM100/B200, or trtllm decode) # Allocate global workspace buffer for TRT-LLM kernels (ragged attention on SM100/B200, or trtllm decode)
if self.device_sm_major >= 10 or self.dsa_decode_impl == "trtllm": if self.device_sm_major >= 10 or self.dsa_decode_impl == "trtllm":
self.workspace_buffer = get_buffer( self.workspace_buffer = get_buffer(
@@ -1981,12 +2019,46 @@ class DeepseekSparseAttnBackend(
sm_scale=layer.scaling, sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim, v_head_dim=layer.v_head_dim,
) )
elif dsa_impl == "flashmla_sparse": elif dsa_impl in ("flashmla_sparse", "flashmla_sparse_q8"):
if q_rope is not None:
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
if topk_transform_method == TopkTransformMethod.RAGGED: if topk_transform_method == TopkTransformMethod.RAGGED:
if any(forward_batch.extend_prefix_lens_cpu): _has_prefix = any(forward_batch.extend_prefix_lens_cpu)
page_table_1 = topk_indices
# `flashmla_sparse_q8` = native FP8 sparse prefill (constructor
# guarantees fp8_e4m3 KV + SM90). The helper consumes q_nope/q_rope
# directly (fusing the concat with the bf16->fp8 cast), so no bf16
# q_all is materialized on this path. The prefix path hands over the
# paged fp8 KV as-is; the non-prefix path passes the gathered bf16 KV.
if dsa_impl == "flashmla_sparse_q8":
if _has_prefix:
page_table_1_flattened = (
self.forward_metadata.page_table_1_flattened
)
assert page_table_1_flattened is not None
return self._forward_flashmla_sparse_q8kv8(
q_nope=q_nope,
q_rope=q_rope,
kv_bf16=None,
paged_kv_cache=kv_cache,
page_table_1_flattened=page_table_1_flattened,
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
kv_cache = _cat([k, k_rope], dim=-1)
return self._forward_flashmla_sparse_q8kv8(
q_nope=q_nope,
q_rope=q_rope,
kv_bf16=kv_cache,
paged_kv_cache=None,
page_table_1_flattened=None,
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
# bf16 path (dsa_impl == "flashmla_sparse").
if _has_prefix:
page_table_1_flattened = ( page_table_1_flattened = (
self.forward_metadata.page_table_1_flattened self.forward_metadata.page_table_1_flattened
) )
@@ -1996,8 +2068,9 @@ class DeepseekSparseAttnBackend(
) )
else: else:
kv_cache = _cat([k, k_rope], dim=-1) kv_cache = _cat([k, k_rope], dim=-1)
page_table_1 = topk_indices
if q_rope is not None:
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_sparse( return self._forward_flashmla_sparse(
q_all=q_all, q_all=q_all,
kv_cache=kv_cache, kv_cache=kv_cache,
@@ -2292,6 +2365,131 @@ class DeepseekSparseAttnBackend(
return o return o
def _forward_flashmla_sparse_q8kv8(
self,
q_nope: torch.Tensor,
q_rope: torch.Tensor,
kv_bf16: Optional[torch.Tensor],
v_head_dim: int,
page_table_1: torch.Tensor,
sm_scale: float,
paged_kv_cache: Optional[torch.Tensor] = None,
page_table_1_flattened: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Native FP8 (q8 x kv8) sparse-prefill attention (SM90 JIT kernel).
Same contract as ``_forward_flashmla_sparse`` but executed through the
FP8 ``sparse_mla_q8kv8_prefill_fwd`` kernel. Identity per-tensor
scales (scalar 1.0) are used: a raw bf16->fp8 cast of q/kv is accurate
on real DeepSeek-V3 magnitudes, so no dynamic rescaling is applied.
The kernel runs via its fixed full-topk entry (``attn_sink`` /
``topk_length`` left None), keeping control flow identical across DP
ranks; -1 topk sentinels are clamped to distinct zero pad rows inside
the kernel.
Two KV paths:
* non-prefix extend: ``kv_bf16`` (the gathered bf16 KV) is cast into
a zero-padded fp8 buffer.
* prefix extend: ``paged_kv_cache`` (fp8, 656 B/token: nope_fp8 +
per-group scales + rope_bf16) is gathered, dequantized per group,
and requantized to per-tensor fp8 in one fused Triton kernel
(``gather_dequant_requant_fp8_paged``) — no intermediate bf16
materialization.
"""
from sglang.jit_kernel.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
num_tokens, num_heads, d_nope = q_nope.shape
head_dim = d_nope + q_rope.shape[-1]
dev = q_nope.device
# The SM90 kernel requires num_heads % 64 == 0; smaller head counts
# (high-TP splits) are zero-padded up to 64.
required_padding = 64
need_padding = num_heads % required_padding != 0
# Build the fp8 q. concat_and_cast_q_fp8_pad fuses the nope/rope
# concat with the bf16->fp8 cast in one Triton kernel (bit-exact vs
# concat + .to(fp8)); it requires power-of-two head/dim counts (a
# tl.arange constraint), so non-power-of-two head counts fall back to
# the generic concat + cast.
if need_padding:
if required_padding % num_heads != 0:
raise ValueError(
f"num_heads={num_heads} cannot be padded to {required_padding}; "
"this TP size is incompatible with flashmla_sparse_q8."
)
# Cached zero-padded fp8 q buffer: the pad rows [num_heads:64] are
# zero on first alloc and only ever read by the kernel; the active
# slice is overwritten each forward. Eager-mode DSA runs layers
# sequentially on one stream, so single-buffer reuse is safe.
# Grown on demand.
buf = self._q8kv8_qpad_buf
if buf is None or buf.shape[0] < num_tokens:
buf = torch.zeros(
(num_tokens, required_padding, head_dim),
dtype=torch.float8_e4m3fn,
device=dev,
)
self._q8kv8_qpad_buf = buf
q_fp8 = buf[:num_tokens]
# head counts that divide 64 are powers of two, so the fused
# concat-cast is always applicable here.
concat_and_cast_q_fp8_pad(q_fp8, q_nope, q_rope, num_heads)
elif (num_heads & (num_heads - 1)) == 0:
q_fp8 = q_nope.new_empty(
(num_tokens, num_heads, head_dim), dtype=torch.float8_e4m3fn
)
concat_and_cast_q_fp8_pad(q_fp8, q_nope, q_rope, num_heads)
else:
# Generic fallback for non-power-of-two head counts.
q_fp8 = concat_mla_absorb_q_general(q_nope, q_rope).to(torch.float8_e4m3fn)
# Identity per-tensor scale, cached: creating it per call is a
# host->device copy that synchronizes the stream.
identity_scale = self._q8kv8_identity_scale
if identity_scale is None:
identity_scale = torch.tensor([1.0], dtype=torch.float32, device=dev)
self._q8kv8_identity_scale = identity_scale
# KV: append `topk` trailing zero rows so the kernel's -1-sentinel
# clamp can map every padded topk slot to a DISTINCT zero row.
# Mapping many slots onto one shared row would serialize the kernel's
# KV gather; distinct zero rows are value-identical (zero KV
# contributes nothing to the softmax-weighted sum) at full speed.
topk = page_table_1.shape[-1]
if paged_kv_cache is not None:
kv_padded = gather_dequant_requant_fp8_paged(
paged_kv_cache,
page_table_1_flattened,
extra_rows=topk,
).view(-1, 1, head_dim)
else:
kv_padded = kv_bf16.new_zeros(
(kv_bf16.shape[0] + topk, *kv_bf16.shape[1:]),
dtype=torch.float8_e4m3fn,
)
kv_padded[: kv_bf16.shape[0]].copy_(kv_bf16)
kv_padded = kv_padded.view(-1, 1, head_dim)
o, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q_fp8,
kv=kv_padded,
indices=page_table_1.unsqueeze(1),
sm_scale=sm_scale,
q_scale=identity_scale,
kv_scale=identity_scale,
d_v=v_head_dim,
attn_sink=None,
topk_length=None,
)
# Trim the output back to the original head count if we padded.
if need_padding:
o = o[:, :num_heads, :]
return o
def _forward_flashmla_kv( def _forward_flashmla_kv(
self, self,
q_all: torch.Tensor, q_all: torch.Tensor,
@@ -2847,7 +3045,11 @@ class DeepseekSparseAttnBackend(
if ( if (
# disable for MTP # disable for MTP
self.dsa_kv_cache_store_fp8 self.dsa_kv_cache_store_fp8
and self.dsa_prefill_impl == "flashmla_sparse" # flashmla_sparse_q8 shares flashmla_sparse's RAGGED prefill routing — the q8
# dispatch lives inside the RAGGED branch of forward_extend; without this the
# transform is PAGED, the q8 path is skipped, and the bf16 kernel crashes on
# fp8 KV ("kv must have dtype kBFloat16").
and self.dsa_prefill_impl in ("flashmla_sparse", "flashmla_sparse_q8")
and forward_mode == ForwardMode.EXTEND and forward_mode == ForwardMode.EXTEND
): ):
topk_transform_method = TopkTransformMethod.RAGGED topk_transform_method = TopkTransformMethod.RAGGED
+1
View File
@@ -323,6 +323,7 @@ DEFAULT_LORA_EVICTION_POLICY = "lru"
DSA_CHOICES = [ DSA_CHOICES = [
"flashmla_sparse", "flashmla_sparse",
"flashmla_sparse_q8",
"flashmla_kv", "flashmla_kv",
"flashmla_auto", "flashmla_auto",
"fa3", "fa3",
@@ -9,7 +9,7 @@ import torch
from sglang.srt.utils import is_sm90_supported from sglang.srt.utils import is_sm90_supported
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="1-gpu-large") register_cuda_ci(est_time=240, stage="base-b-kernel-unit", runner_config="1-gpu-large")
DTYPE_FP8 = torch.float8_e4m3fn DTYPE_FP8 = torch.float8_e4m3fn
@@ -417,5 +417,190 @@ def test_sparse_mla_q8kv8_prefill_rejects_bad_buffers():
_call(d_v=256) _call(d_v=256)
# ---------------------------------------------------------------------------
# End-to-end-discovered corner-case gates.
#
# Everything ABOVE (matches_reference / corner_cases / precision /
# no_alias_between_calls / caller_owned_buffers / rejects_bad_buffers) is the
# original unit suite: small, all-valid (or topk_length-bounded) shapes checked
# against a reference that REPRODUCES the kernel's own clamp semantics. That
# suite is blind to three bug classes that only surface under real
# DeepSeek-V3.2 serving; the gates below reproduce them as standalone kernel
# tests:
#
# 1. masked -1-sentinel SEMANTICS on few-valid rows (ctx << topk): pad slots
# must contribute NOTHING to the softmax denominator. This needs a MASKED
# (-inf) fp32 reference -- a reference that mimics the kernel's own clamp
# is blind to the bug.
# 2. s_q ENVELOPE to 6144: first-band NaNs from an is_kv_valid data race that
# only appears past s_q=2048 (never exercised above).
# 3. LARGE S_KV (65536) / large index values: gathered multi-request buffers
# reach tens of thousands of rows in e2e; the suite above used s_kv<=1024.
#
# These use h=128 (the real DeepSeek head count) and large s_q/s_kv, so they are
# heavier than the suite above; same SM90 skipif. They also DOCUMENT that the
# kernel is run-to-run nondeterministic at the fp8 noise floor, so they compare
# against an fp32 reference (never bitwise / self-consistency).
# ---------------------------------------------------------------------------
_D_FULL = 576 # nope(512) + rope(64): the real DeepSeek MLA absorbed q/kv width
def _ref_masked_blocked(q, kv, indices, sm_scale, d_v, row_start, row_end):
"""fp32 reference with PROPER -1 masking (pad slots -> -inf), computed over a
block of query rows [row_start, row_end) to bound peak memory. Unlike
``_torch_sparse_attention_ref`` (which bounds validity via topk_length and so
reproduces the kernel's clamp), this masks every -1 index out of the softmax,
making it sensitive to the denominator-pollution bug."""
q_f = q.float()
kv_f = kv.float()[:, 0, :]
idx_block = indices[row_start:row_end, 0, :].long()
gathered = kv_f[idx_block.clamp(min=0)]
scores = torch.einsum("qhd,qkd->qhk", q_f[row_start:row_end], gathered) * sm_scale
scores = scores.masked_fill((idx_block < 0)[:, None, :], float("-inf"))
probs = torch.softmax(scores, dim=-1)
return torch.einsum("qhk,qkd->qhd", probs, gathered[:, :, :d_v])
@pytest.mark.skipif(
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
)
@pytest.mark.parametrize("s_q", [2048, 4096])
def test_sparse_mla_q8kv8_prefill_masked_sentinels(s_q: int):
"""NEW gate (bug class 1): causal -1 structure (row i has min(1+i, topk)
valid slots, the rest -1). The kernel must mask pad slots out of the softmax
denominator. Checked against a MASKED (-inf) fp32 reference; a reference
that reproduced the kernel's clamp would be blind to this.
Gate: per-band cos > 0.97 AND magnitude ratio > 0.9
(the denominator-pollution bug crushes magnitude 50-2000x, unmistakable even
under fp8 noise)."""
from sglang.jit_kernel.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
h, topk, band, n = 128, 2048, 512, 4608
s_kv = n + topk
g = torch.Generator(device="cuda").manual_seed(11)
q = torch.randn((s_q, h, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
kv = torch.zeros((s_kv, H_KV, _D_FULL), dtype=DTYPE_FP8, device="cuda")
kv[:n] = torch.randn((n, H_KV, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
idx = torch.full((s_q, H_KV, topk), -1, dtype=torch.int32, device="cuda")
slot = torch.arange(topk, device="cuda")
valid = torch.clamp(1 + torch.arange(s_q, device="cuda"), max=topk)
rnd = torch.randint(
0, n, (s_q, topk), dtype=torch.int32, device="cuda", generator=g
)
idx[:, 0, :] = torch.where(
slot[None, :] < valid[:, None], rnd, torch.full_like(rnd, -1)
)
one = torch.ones(1, dtype=torch.float32, device="cuda")
sm_scale = 1.0 / math.sqrt(_D_FULL)
out, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q, kv=kv, indices=idx, sm_scale=sm_scale, q_scale=one, kv_scale=one, d_v=D_V
)
torch.cuda.synchronize()
worst_cos, worst_mag = 1.0, 1.0
for s in range(0, s_q, band):
e = min(s + band, s_q)
ref = _ref_masked_blocked(q, kv, idx, sm_scale, D_V, s, e)
ob = out[s:e].float()
cos = torch.nn.functional.cosine_similarity(
ob.reshape(-1), ref.reshape(-1), dim=0
).item()
mag = (ob.norm() / ref.norm().clamp(min=1e-9)).item()
worst_cos = min(worst_cos, cos)
if mag < 1.0:
worst_mag = min(worst_mag, mag)
del ref, ob
torch.cuda.empty_cache()
print(
f"\n masked-sentinels s_q={s_q}: worst cos={worst_cos:.4f} "
f"worst |out|/|ref|={worst_mag:.3f}"
)
assert worst_cos > 0.97, f"cos {worst_cos:.4f} <= 0.97 (denominator pollution?)"
assert worst_mag > 0.9, f"mag {worst_mag:.3f} <= 0.9 (denominator pollution?)"
@pytest.mark.skipif(
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
)
@pytest.mark.parametrize("s_q", [2048, 4096, 6144])
def test_sparse_mla_q8kv8_prefill_sq_envelope(s_q: int):
"""NEW gate (bug class 2): all-valid correctness across the s_q envelope.
s_q=6144 previously produced first-band NaNs (an is_kv_valid data race that
only appears past s_q=2048)."""
from sglang.jit_kernel.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
h, topk, s_kv, band = 128, 2048, 8192, 1024
g = torch.Generator(device="cuda").manual_seed(7)
q = torch.randn((s_q, h, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
kv = torch.randn((s_kv, H_KV, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
idx = torch.randint(
0, s_kv, (s_q, H_KV, topk), dtype=torch.int32, device="cuda", generator=g
)
one = torch.ones(1, dtype=torch.float32, device="cuda")
sm_scale = 1.0 / math.sqrt(_D_FULL)
out, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q, kv=kv, indices=idx, sm_scale=sm_scale, q_scale=one, kv_scale=one, d_v=D_V
)
torch.cuda.synchronize()
has_nan = torch.isnan(out.float()).any().item()
worst_cos = 1.0
for s in range(0, s_q, band):
e = min(s + band, s_q)
ref = _ref_masked_blocked(q, kv, idx, sm_scale, D_V, s, e)
cos = torch.nn.functional.cosine_similarity(
out[s:e].float().reshape(-1), ref.reshape(-1), dim=0
).item()
worst_cos = min(worst_cos, cos)
del ref
torch.cuda.empty_cache()
print(f"\n s_q-envelope s_q={s_q}: nan={has_nan} worst cos={worst_cos:.4f}")
assert not has_nan, f"NaN in output at s_q={s_q} (is_kv_valid race)"
assert worst_cos > 0.99, f"cos {worst_cos:.4f} <= 0.99"
@pytest.mark.skipif(
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
)
def test_sparse_mla_q8kv8_prefill_large_skv():
"""NEW gate (bug class 3): large gathered buffers / large index values
(s_kv=65536, indices in [33000, 65536)). E2E multi-request gather buffers
reach tens of thousands of rows; the suite above used s_kv<=1024."""
from sglang.jit_kernel.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
h, topk, s_kv, s_q = 128, 2048, 65536, 2048
g = torch.Generator(device="cuda").manual_seed(13)
q = torch.randn((s_q, h, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
kv = torch.randn((s_kv, H_KV, _D_FULL), device="cuda", generator=g).to(DTYPE_FP8)
idx = torch.randint(
33000, s_kv, (s_q, H_KV, topk), dtype=torch.int32, device="cuda", generator=g
)
one = torch.ones(1, dtype=torch.float32, device="cuda")
sm_scale = 1.0 / math.sqrt(_D_FULL)
out, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q, kv=kv, indices=idx, sm_scale=sm_scale, q_scale=one, kv_scale=one, d_v=D_V
)
torch.cuda.synchronize()
ref = _ref_masked_blocked(q, kv, idx, sm_scale, D_V, 0, 1024)
cos = torch.nn.functional.cosine_similarity(
out[:1024].float().reshape(-1), ref.reshape(-1), dim=0
).item()
print(f"\n large-S_KV={s_kv}: band-0 cos={cos:.4f}")
assert cos > 0.99, f"cos {cos:.4f} <= 0.99"
if __name__ == "__main__": if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"])) sys.exit(pytest.main([__file__, "-v", "-s"]))