[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:
co-authored by
Claude Opus 4.8
Xiaoyu Zhang
parent
cce5fe7696
commit
b8ec544946
@@ -14,11 +14,11 @@ limitations under the License.
|
||||
==============================================================================*/
|
||||
|
||||
// SM90 FP8 native sparse MLA prefill kernel.
|
||||
//
|
||||
|
||||
// Algorithm inspired by DeepSeek FlashMLA
|
||||
// (https://github.com/deepseek-ai/FlashMLA); the kernel itself is a
|
||||
// clean-room re-implementation targeting the Q8KV8 sparse prefill path.
|
||||
//
|
||||
|
||||
// Design: Native fp8 GMMA path
|
||||
// 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)
|
||||
@@ -137,15 +137,36 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
array_aligned<fp8_t, cosize_v<SmemLayoutQ>> q; // B_H * D_Q fp8
|
||||
array_aligned<bf16, cosize_v<SmemLayoutO>> o; // B_H * D_V/2 bf16
|
||||
} q_o;
|
||||
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
|
||||
array_aligned<fp8_t, cosize_v<SmemLayoutK>> k[2]; // 2x K double-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.
|
||||
|
||||
bool is_kv_valid[2][B_TOPK];
|
||||
float2 sM[32];
|
||||
// double-buffer the VALIDITY DATA by iteration-pair
|
||||
// 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];
|
||||
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;
|
||||
// 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.
|
||||
@@ -187,7 +208,10 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
plan.bar_k0_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
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
// 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 +
|
||||
q_h_idx * B_H * (int64_t)params.stride_q_h_q;
|
||||
|
||||
// Vectorized Q loading via cp.async.cg (16 bytes per op)
|
||||
constexpr int Q_GROUP_SIZE = 8;
|
||||
// Vectorized Q loading via cp.async.cg (16 bytes per op).
|
||||
// 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_ROWS_PER_GROUP = B_H / Q_NUM_GROUPS;
|
||||
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 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 rP = partition_fragment_C(TiledMMA_QK{}, Shape<Int<B_H>, Int<B_TOPK>>{});
|
||||
cute::fill(rO, 0.0f);
|
||||
@@ -259,6 +292,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
Tensor rP_fp8_local = make_tensor<fp8_t>(rP_fp8_layout_t{});
|
||||
|
||||
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 Warpgroup1 {};
|
||||
|
||||
@@ -272,28 +306,39 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
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>;
|
||||
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
|
||||
for (int row_idx = 0; row_idx < 2; ++row_idx) {
|
||||
CUTE_UNROLL
|
||||
for (int i = row_idx * 2; i < size(rP); i += 4) {
|
||||
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][col + 1]) rP(i + 1) = -INFINITY;
|
||||
if (!plan.is_kv_valid[IS_WG1][_kv_buf][col]) rP(i) = -INFINITY; // parity-buffered
|
||||
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
|
||||
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.
|
||||
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
|
||||
const float scale = qk_combined_scale_div_log2;
|
||||
float r_sM[2];
|
||||
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];
|
||||
CUTE_UNROLL
|
||||
@@ -307,6 +352,12 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
cur_max = max(cur_max, __shfl_xor_sync(0xffffffff, cur_max, 2));
|
||||
cur_max *= scale;
|
||||
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]);
|
||||
CUTE_UNROLL
|
||||
for (int i = row_idx * 2; i < size(rO); i += 4) {
|
||||
@@ -326,7 +377,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
}
|
||||
__syncwarp();
|
||||
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[1] = new_maxs[1];
|
||||
@@ -441,32 +492,10 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
}
|
||||
};
|
||||
|
||||
auto undo_v_transpose_col_permutation = [&]() {
|
||||
// Undo the column permutation from the fp8 V transpose before writing O.
|
||||
// CLayout_64x256: col bit0 = t1_bit0 (thread), col bit3 = v1 (register).
|
||||
// V transpose introduces bit0<->bit3 swap. Fix by cross-thread exchange:
|
||||
// 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;
|
||||
}
|
||||
}
|
||||
};
|
||||
// No output un-permutation is needed here: the PV output columns are
|
||||
// produced by the Vt B-operand, not by the permute_Cregs_fp8-rearranged P
|
||||
// A-operand (the permute only reorders the P contraction/K dimension,
|
||||
// never the output N).
|
||||
|
||||
// ============================================================
|
||||
// WG0 Pipeline -- native fp8
|
||||
@@ -508,7 +537,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
CUTE_NO_UNROLL
|
||||
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
|
||||
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) {
|
||||
pipelined_wait_and_qkt_gemm_l();
|
||||
@@ -518,11 +547,16 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
plan.bar_k0_free.arrive();
|
||||
}
|
||||
|
||||
mask_rP(Warpgroup0{});
|
||||
online_softmax_and_rescale_o(Warpgroup0{});
|
||||
mask_rP(Warpgroup0{}, block_idx); // pass block_idx for kv_buf parity
|
||||
online_softmax_and_rescale_o(Warpgroup0{}, (block_idx >> 1) & 1);
|
||||
|
||||
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).
|
||||
// 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.
|
||||
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg1_bunch_0_ready);
|
||||
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
|
||||
for (int i = 0; i < 2; ++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(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);
|
||||
warpgroup_commit_batch();
|
||||
|
||||
@@ -579,8 +613,6 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
}
|
||||
}
|
||||
|
||||
undo_v_transpose_col_permutation();
|
||||
|
||||
reduce_L();
|
||||
store_O();
|
||||
|
||||
@@ -614,7 +646,7 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
CUTE_NO_UNROLL
|
||||
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
|
||||
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) {
|
||||
pipelined_wait_and_qkt_gemm_r_wg1();
|
||||
@@ -624,13 +656,20 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
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);
|
||||
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());
|
||||
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)
|
||||
NamedBarrier::arrive_and_wait(256, vt1_for_wg1);
|
||||
@@ -643,6 +682,18 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
warpgroup_fence_operand(rO);
|
||||
warpgroup_fence_operand(rP_fp8_local);
|
||||
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());
|
||||
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.
|
||||
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);
|
||||
warpgroup_commit_batch();
|
||||
|
||||
@@ -658,8 +709,17 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
cur_bar_wait_phase ^= 1;
|
||||
// Overlap: start next-iteration QK-right while PV drains
|
||||
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);
|
||||
// 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);
|
||||
plan.bar_vt_free[0].arrive();
|
||||
pipelined_wait_and_qkt_gemm_l_wg1();
|
||||
@@ -667,14 +727,21 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
warpgroup_fence_operand(rP);
|
||||
plan.bar_k1_free.arrive();
|
||||
} else {
|
||||
warpgroup_wait<0>();
|
||||
warpgroup_wait<0>(); // drains the peer PV
|
||||
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();
|
||||
}
|
||||
}
|
||||
|
||||
undo_v_transpose_col_permutation();
|
||||
|
||||
reduce_L();
|
||||
store_O();
|
||||
|
||||
@@ -720,6 +787,9 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
int64_t token_indices[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) {
|
||||
CUTE_UNROLL
|
||||
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) {
|
||||
int offs = (block_idx + buf_idx) * B_TOPK + local_row * NUM_GROUPS + group_idx;
|
||||
int t = __ldg(gIndices + offs);
|
||||
bool is_cur_token_valid = t >= 0 && t < params.s_kv;
|
||||
if constexpr (HAVE_TOPK_LENGTH) {
|
||||
is_cur_token_valid &= offs < topk_length;
|
||||
bool is_cur_token_valid;
|
||||
if constexpr (!HAVE_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;
|
||||
is_token_valid[buf_idx][local_row] = is_cur_token_valid;
|
||||
@@ -766,23 +858,28 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
SmemTransposeV smem_transpose_v;
|
||||
using SmemLayoutTransposeV_t = typename SmemTransposeV::SmemLayoutTransposeV;
|
||||
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.
|
||||
// but same composition-based framework as before.
|
||||
auto transpose_v_half = [&](int smem_k_buf, int vt_buf, int tile_start, int tile_end) {
|
||||
auto transpose_v_to_half = [&](int smem_k_buf, fp8_t* vt_dst, int tile_start, int tile_end) {
|
||||
Tensor sV_src = as_position_independent_swizzle_tensor(
|
||||
make_tensor(make_smem_ptr(plan.k[smem_k_buf].data()), SmemLayoutTransposeV_t{}));
|
||||
Tensor sVt_dst = as_position_independent_swizzle_tensor(
|
||||
make_tensor(make_smem_ptr(plan.vt[vt_buf].data()), SmemLayoutTransposeVt_t{}));
|
||||
Tensor sVt_dst =
|
||||
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");
|
||||
CUTE_UNROLL
|
||||
for (int j = tile_start; j < tile_end; j += 2) {
|
||||
smem_transpose_v.transpose_pair(
|
||||
flatten(sV_src(_, 0, j)),
|
||||
flatten(sVt_dst(_, 0, j)),
|
||||
flatten(sVt_dst(_, 0, j - tile_start)),
|
||||
flatten(sV_src(_, 0, j + 1)),
|
||||
flatten(sVt_dst(_, 0, j + 1)));
|
||||
flatten(sVt_dst(_, 0, j + 1 - tile_start)));
|
||||
}
|
||||
asm volatile("" ::: "memory");
|
||||
};
|
||||
@@ -797,7 +894,6 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
CUTE_NO_UNROLL
|
||||
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.
|
||||
|
||||
plan.bar_k0_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)
|
||||
CUTE_UNROLL
|
||||
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.bar_is_kv_valid_ready.arrive();
|
||||
plan.is_kv_valid[buf_idx][(block_idx >> 1) & 1][local_row * NUM_GROUPS + group_idx] =
|
||||
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);
|
||||
@@ -830,7 +927,9 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
commit_to_mbar(plan.bar_k1_ready[0]);
|
||||
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" ::);
|
||||
// fence.proxy.async: make cp.async data visible through generic 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);
|
||||
}
|
||||
|
||||
transpose_v_half(0, 0, 0, 4);
|
||||
NamedBarrier::arrive(256, vt0_left_ready);
|
||||
transpose_v_to_half(0, plan.vt_loc[0].data(), 0, 4); // block0-left -> WG0 local
|
||||
// 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.
|
||||
// WG0 is on the critical path (feeds WG1 via sM/wg0_bunch).
|
||||
// WG0 waits for vt1_for_wg0 (V[1]-LEFT) for PV-remote.
|
||||
// Moving V[1]-LEFT earlier (2nd instead of 4th) reduces WG0
|
||||
// critical-path stall by ~768 cycles per iteration.
|
||||
//
|
||||
// v52 CRASH FIX: K[1]-left tiles 0-3 are in cp.async group-1,
|
||||
// critical-path stall.
|
||||
|
||||
// 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
|
||||
// CTA counts (512+), memory bandwidth saturation delays group-1
|
||||
// completion past the V[0]-LEFT transpose timing margin, causing
|
||||
@@ -870,20 +972,25 @@ struct SparseMlaQ8Kv8PrefillKernel {
|
||||
// generic proxy for LDSM reads in V transpose
|
||||
fence_view_async_shared();
|
||||
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
|
||||
if (block_idx > 0) {
|
||||
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);
|
||||
|
||||
// V[0]-RIGHT: tiles 4-7 from K[0]
|
||||
transpose_v_half(0, 0, 4, 8);
|
||||
NamedBarrier::arrive(256, vt0_right_ready);
|
||||
|
||||
transpose_v_to_half(0, plan.vt_rem[1].data(), 4, 8); // block0-right -> WG1 remote
|
||||
// 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);
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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__":
|
||||
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 (
|
||||
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 (
|
||||
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),
|
||||
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
|
||||
|
||||
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.transform_index import (
|
||||
transform_index_page_table_decode,
|
||||
@@ -29,6 +32,7 @@ from sglang.kernels.ops.attention.utils import (
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
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.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import (
|
||||
@@ -323,7 +327,7 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
||||
|
||||
|
||||
_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.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)
|
||||
if self.device_sm_major >= 10 or self.dsa_decode_impl == "trtllm":
|
||||
self.workspace_buffer = get_buffer(
|
||||
@@ -1981,12 +2019,46 @@ class DeepseekSparseAttnBackend(
|
||||
sm_scale=layer.scaling,
|
||||
v_head_dim=layer.v_head_dim,
|
||||
)
|
||||
elif dsa_impl == "flashmla_sparse":
|
||||
if q_rope is not None:
|
||||
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
|
||||
|
||||
elif dsa_impl in ("flashmla_sparse", "flashmla_sparse_q8"):
|
||||
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 = (
|
||||
self.forward_metadata.page_table_1_flattened
|
||||
)
|
||||
@@ -1996,8 +2068,9 @@ class DeepseekSparseAttnBackend(
|
||||
)
|
||||
else:
|
||||
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(
|
||||
q_all=q_all,
|
||||
kv_cache=kv_cache,
|
||||
@@ -2292,6 +2365,131 @@ class DeepseekSparseAttnBackend(
|
||||
|
||||
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(
|
||||
self,
|
||||
q_all: torch.Tensor,
|
||||
@@ -2847,7 +3045,11 @@ class DeepseekSparseAttnBackend(
|
||||
if (
|
||||
# disable for MTP
|
||||
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
|
||||
):
|
||||
topk_transform_method = TopkTransformMethod.RAGGED
|
||||
|
||||
@@ -323,6 +323,7 @@ DEFAULT_LORA_EVICTION_POLICY = "lru"
|
||||
|
||||
DSA_CHOICES = [
|
||||
"flashmla_sparse",
|
||||
"flashmla_sparse_q8",
|
||||
"flashmla_kv",
|
||||
"flashmla_auto",
|
||||
"fa3",
|
||||
|
||||
Reference in New Issue
Block a user