[DSA] Q8KV8 FP8 Sparse Prefill on GLM-5.2 & DeepSeek-V3.2: Q8-Path & Shared-Path Optimizations (#31888)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Ho-Ren (Jack) Chuang
2026-07-30 15:15:11 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 4f51dad1da
commit e4a40a71f8
30 changed files with 3837 additions and 52 deletions
@@ -0,0 +1,84 @@
/* Copyright 2026 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
// JIT dispatch entry for the SM90 Q8KV8 born-fp8 q-prep kernel.
#pragma once
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include "kernel.cuh"
#include <cstdint>
#include <cuda_runtime.h>
namespace {
// All strides are in elements; validation of dtypes/shapes/alignment happens
// in the Python wrapper (sglang/kernels/ops/attention/qprep_bf16_fp8_sm90.py).
void qprep_bf16_fp8_dispatch(
tvm::ffi::TensorView q_nope,
tvm::ffi::TensorView w_kc,
tvm::ffi::TensorView q_rope,
tvm::ffi::TensorView out,
int64_t num_tokens,
int64_t num_heads,
int64_t k_dim,
int64_t a_s0,
int64_t a_s1,
int64_t b_s0,
int64_t b_s2,
int64_t r_s0,
int64_t r_s1,
int64_t o_s0,
int64_t o_s1,
int64_t rope_vec16,
int64_t out_vec16,
int64_t cuda_stream) {
QprepBf16Fp8Sm90Params params;
params.num_tokens = (int)num_tokens;
params.num_heads = (int)num_heads;
params.q_nope = q_nope.data_ptr();
params.a_s0 = a_s0;
params.a_s1 = a_s1;
params.w_kc = w_kc.data_ptr();
params.b_s0 = b_s0;
params.b_s2 = b_s2;
params.q_rope = q_rope.data_ptr();
params.r_s0 = r_s0;
params.r_s1 = r_s1;
params.rope_vec16 = (bool)rope_vec16;
params.out = out.data_ptr();
params.o_s0 = o_s0;
params.o_s1 = o_s1;
params.out_vec16 = (bool)out_vec16;
DLDevice dev = q_nope.device();
cudaSetDevice(dev.device_id);
params.stream = reinterpret_cast<cudaStream_t>(cuda_stream);
switch (k_dim) {
case 128:
qprep_sm90::run_qprep_bf16_fp8_sm90<128>(params);
return;
case 192:
qprep_sm90::run_qprep_bf16_fp8_sm90<192>(params);
return;
default:
fprintf(stderr, "qprep_bf16_fp8_sm90: unsupported k_dim=%ld (must be 128 or 192)\n", (long)k_dim);
exit(1);
}
}
} // namespace
@@ -0,0 +1,516 @@
/* Copyright 2026 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
// SM90 (Hopper) Q8KV8 born-fp8 q-prep kernel.
//
// Computes, per head h:
// out[:, h, :512] = fp8_e4m3(bf16(fp32_accum(q_nope[:, h, :] @ w_kc[h])))
// out[:, h, 512:576] = fp8_e4m3(q_rope[:, h, :])
//
// This is the CUDA replacement for the Triton absorbed_bmm_concat_cast_q_fp8
// kernel (triton_ops/cache_ops.py). The epilogue keeps the exact rounding
// chain of the Triton variants: fp32 WGMMA accumulate -> bf16 round-to-nearest
// (cublas-equivalent output rounding) -> fp8_e4m3 rn/satfinite on store. The
// K dimension is consumed as one in-order chain of k=16 WGMMA steps into a
// single fp32 accumulator, i.e. the same fp32 add order as the Triton
// "two_dot"/"grouped" variants (128+64 chained tl.dot), so the nope half can
// come out bitwise identical to them.
//
// Phase-2 design (2 CTAs/SM + double-buffered B):
// grid = (ceil(T / 128), H); one CTA = two WGMMA warpgroups (256 threads)
// owning a 128-row m-tile of one head (warpgroup w computes rows
// [64w, 64w+64)). The A tile [128, K] bf16 is cp.async'd to smem once (L2
// evict_first: streamed) and the rope path runs under that load's wait.
// The N=512 output is produced in N_SLABS n-slabs of BN columns; the B
// slab [BN, K] bf16 is double-buffered (L2 evict_last: re-read by every
// CTA of the head) and prefetched one full round ahead. Per round, the
// fp8 stage-write -> barrier -> refill-issue -> coalesced-flush order
// makes one barrier serve both the stage handoff and the CTA-wide WGMMA
// drain of the buffer being refilled, and the flush plus the next round's
// gemm overlap the refill. BN is sized so that A + 2 B buffers + the fp8
// stage fit in half an SM's smem, keeping 2 CTAs co-resident per SM
// (register cap 128 via launch bounds; measured faster than every
// 1-CTA/SM variant tried, including wider CTAs and dual-accumulator
// cross-round software pipelines): K=192 -> BN=64 (104 KB), K=128 ->
// BN=128 (112 KB). The round loop is left un-unrolled when N_SLABS > 4:
// full unrolling blows the 128-register budget and spills to local.
#pragma once
#include <cute/tensor.hpp>
#include <cutlass/bfloat16.h>
#include "params.h"
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cuda_bf16.h>
#include <type_traits>
namespace qprep_sm90 {
using namespace cute;
using bf16 = cutlass::bfloat16_t;
#define QPREP_ASSERT(cond) \
do { \
if (!(cond)) { \
fprintf(stderr, "QPREP_ASSERT failed (%s:%d): %s\n", __FILE__, __LINE__, #cond); \
exit(1); \
} \
} while (0)
#define QPREP_CUDA_CHECK(call) \
do { \
cudaError_t err = (call); \
if (err != cudaSuccess) { \
fprintf(stderr, "CUDA error (%s:%d): %s\n", __FILE__, __LINE__, cudaGetErrorString(err)); \
exit(1); \
} \
} while (0)
__host__ __device__ __forceinline__ constexpr int ceil_div_i(int a, int b) {
return (a + b - 1) / b;
}
// ---------------------------------------------------------------------------
// Device helpers
// ---------------------------------------------------------------------------
// L2 eviction policies (same helpers as the sparse-prefill kernel): A/rope
// are streamed once (evict_first); the per-head w_kc slice is re-read from L2
// by every CTA of the head (evict_last).
__device__ __forceinline__ int64_t createpolicy_evict_last() {
int64_t res;
asm volatile("createpolicy.fractional.L2::evict_last.b64 %0, 1.0; \n\t" : "=l"(res) :);
return res;
}
__device__ __forceinline__ int64_t createpolicy_evict_first() {
int64_t res;
asm volatile("createpolicy.fractional.L2::evict_first.b64 %0, 1.0; \n\t" : "=l"(res) :);
return res;
}
// 16-byte cp.async.cg with an L2 cache policy, zero-filling when pred is
// false (same instruction family as the sparse-prefill producer).
__device__ __forceinline__ void
cp_async_16_zfill(void* smem_dst, const void* gmem_src, bool pred, int64_t cache_policy) {
uint32_t dst_addr = cute::cast_smem_ptr_to_uint(smem_dst);
asm volatile(
"cp.async.cg.shared.global.L2::cache_hint.L2::256B [%0], [%1], 16, %2, %3;\n" ::"r"(dst_addr),
"l"(gmem_src),
"r"(pred ? 16 : 0),
"l"(cache_policy));
}
__device__ __forceinline__ void cp_async_16(void* smem_dst, const void* gmem_src, int64_t cache_policy) {
uint32_t dst_addr = cute::cast_smem_ptr_to_uint(smem_dst);
asm volatile(
"cp.async.cg.shared.global.L2::cache_hint.L2::256B [%0], [%1], 16, %2;\n" ::"r"(dst_addr),
"l"(gmem_src),
"l"(cache_policy));
}
// Pack two fp32 into two fp8_e4m3 bytes with round-to-nearest + satfinite.
// PTX: cvt.rn.satfinite.e4m3x2.f32 d, a, b -> d[7:0] = cvt(b), d[15:8] = cvt(a).
__device__ __forceinline__ uint16_t f32x2_to_e4m3x2_rn_satfinite(float f_lo, float f_hi) {
uint16_t v;
asm volatile("cvt.rn.satfinite.e4m3x2.f32 %0, %1, %2;\n" : "=h"(v) : "f"(f_hi), "f"(f_lo));
return v;
}
// The exact Triton epilogue rounding chain for the nope half:
// fp32 accum -> bf16 (rn) -> fp32 (exact) -> fp8_e4m3 (rn, satfinite).
__device__ __forceinline__ uint16_t f32x2_to_bf16x2_to_e4m3x2(float f0, float f1) {
const __nv_bfloat162 b = __float22bfloat162_rn(make_float2(f0, f1));
return f32x2_to_e4m3x2_rn_satfinite(__low2float(b), __high2float(b));
}
// ---------------------------------------------------------------------------
// Kernel
// ---------------------------------------------------------------------------
template <typename Kernel>
__global__ void qprep_bf16_fp8_kernel(__grid_constant__ const QprepBf16Fp8Sm90Params params);
template <int K_DIM>
struct QprepBf16Fp8Kernel {
static constexpr int BM = 128; // m-tile rows (two WGMMA warpgroups)
// n-slab width: sized so that A + 2 B buffers + the fp8 stage fit in half
// an SM's smem -> 2 CTAs/SM (measured worth more than any intra-CTA
// pipelining): K=128 fits BN=128 (112 KB); K=192 needs BN=64 (104 KB).
static constexpr int BN = (K_DIM > 128) ? 64 : 128;
static constexpr int N_OUT = 512; // kv_lora_rank
static constexpr int ROPE = 64; // qk_rope_head_dim
static constexpr int NUM_THREADS = 256;
static constexpr int N_SLABS = N_OUT / BN;
static constexpr int LOAD_ROWS_PER_PASS = NUM_THREADS / 8; // 16B-chunk loaders
// 2 CTAs/SM co-residency (register cap 128 via launch bounds). Measured
// faster than every 1-CTA/SM variant tried (wider CTAs, dual-accumulator
// cross-round software pipelines).
static constexpr int MIN_CTAS = 2;
static_assert(K_DIM % 64 == 0, "K must tile the SW128 bf16 GMMA atom (64 cols)");
static_assert(N_OUT % BN == 0);
// K-major SW128 smem layouts for the SS WGMMA operands (bf16 atom = 8x64).
using SmemLayoutA =
decltype(tile_to_shape(GMMA::Layout_K_SW128_Atom<bf16>{}, Shape<Int<BM>, Int<K_DIM>>{}, Step<_1, _2>{}));
using SmemLayoutB =
decltype(tile_to_shape(GMMA::Layout_K_SW128_Atom<bf16>{}, Shape<Int<BN>, Int<K_DIM>>{}, Step<_1, _2>{}));
// Two warpgroups stacked along M: threads [128w, 128w+128) own rows
// [64w, 64w+64) of the m-tile. The atom's N width must match BN.
using MmaAtom_t = std::conditional_t<
BN == 128,
SM90_64x128x16_F32BF16BF16_SS<GMMA::Major::K, GMMA::Major::K>,
SM90_64x64x16_F32BF16BF16_SS<GMMA::Major::K, GMMA::Major::K>>;
using TiledMMA_t = decltype(make_tiled_mma(MmaAtom_t{}, Layout<Shape<_2, _1, _1>>{}));
struct SharedStorage {
array_aligned<bf16, cosize_v<SmemLayoutA>, 128> a; // resident A m-block
array_aligned<bf16, cosize_v<SmemLayoutB>, 128> b[2]; // double-buffered B slab
// fp8 output staging for one n-slab: scattered per-thread u16 epilogue
// writes land here, then leave as coalesced 16B global stores (the direct
// u16 global stores 4x-amplify the store sectors and throttle the LSU).
// Single buffer: the end-of-round B wait barrier separates one round's
// copy-out reads from the next round's stage writes.
array_aligned<uint8_t, BM * BN, 16> c_stage;
};
// -------------------------------------------------------------------------
// Loads: NUM_THREADS as (NUM_THREADS/8) row-threads x 8 chunk-threads, 16B
// per cp.async. Smem addresses go through the CUTE tensor so the SW128
// swizzle is applied (16B chunks stay contiguous under the swizzle).
// -------------------------------------------------------------------------
template <typename SmemT>
static __device__ __forceinline__ void
load_a_tile(SmemT& sA, const bf16* gA, int64_t a_s0, int m_residue, int tid, int64_t cache_policy) {
const int cthr = tid % 8, rthr = tid / 8;
CUTE_UNROLL
for (int mi = 0; mi < BM / LOAD_ROWS_PER_PASS; ++mi) {
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
const bool pred = row < m_residue; // zfill OOB rows: 0 * w == 0, never stored
const bf16* g = gA + (int64_t)row * a_s0;
CUTE_UNROLL
for (int ki = 0; ki < K_DIM / 64; ++ki) {
const int col = cthr * 8 + 64 * ki;
cp_async_16_zfill(&sA(row, col), g + col, pred, cache_policy);
}
}
}
template <typename SmemT>
static __device__ __forceinline__ void
load_b_slab(SmemT& sB, const bf16* gB_head, int64_t b_s2, int nb, int tid, int64_t cache_policy) {
const int cthr = tid % 8, rthr = tid / 8;
const bf16* g0 = gB_head + (int64_t)nb * BN * b_s2;
CUTE_UNROLL
for (int ni = 0; ni < BN / LOAD_ROWS_PER_PASS; ++ni) {
const int nrow = rthr + LOAD_ROWS_PER_PASS * ni;
const bf16* g = g0 + (int64_t)nrow * b_s2;
CUTE_UNROLL
for (int ki = 0; ki < K_DIM / 64; ++ki) {
const int col = cthr * 8 + 64 * ki;
cp_async_16(&sB(nrow, col), g + col, cache_policy);
}
}
}
// -------------------------------------------------------------------------
// SS WGMMA over the whole K extent as one in-order k=16 chain (clears the
// accumulator on the first step). Adapted from the sparse-prefill gemm_ss.
// -------------------------------------------------------------------------
template <typename TA, typename TB, typename TC>
static __device__ __forceinline__ void gemm_ss(TiledMMA_t& tiled_mma, TA const& sA, TB const& sB, TC& acc, int tid) {
ThrMMA thr_mma = tiled_mma.get_slice(tid);
Tensor sA_frag = thr_mma.partition_fragment_A(sA);
Tensor sB_frag = thr_mma.partition_fragment_B(sB);
static_assert(size<2>(sA_frag) == size<2>(sB_frag));
warpgroup_fence_operand(acc);
warpgroup_arrive();
tiled_mma.accumulate_ = GMMA::ScaleOut::Zero;
CUTE_UNROLL
for (int k = 0; k < size<2>(sA_frag); ++k) {
cute::gemm(tiled_mma, sA_frag(_, _, k), sB_frag(_, _, k), acc);
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
}
warpgroup_fence_operand(acc);
}
// -------------------------------------------------------------------------
// Epilogue for one n-slab: fp32 acc -> bf16 -> fp8, 2 adjacent columns per
// 16-bit store. WGMMA m64nN C layout: within its warpgroup, thread t holds
// rows (t/32)*16 + (t%32)/4 + {0,8} (plus 64 * warpgroup_idx here) and
// columns (t%4)*2 + 8j + {0,1}; fragment linear index
// i = 4j + 2*row_parity + col_parity.
// -------------------------------------------------------------------------
template <typename TC>
static __device__ __forceinline__ void
store_slab_direct(TC const& acc, uint8_t* gO, int64_t o_s0, int n0, int row_base, int col_base, int m_residue) {
CUTE_UNROLL
for (int rp = 0; rp < 2; ++rp) {
const int row = row_base + 8 * rp;
if (row >= m_residue) continue;
uint8_t* orow = gO + (int64_t)row * o_s0 + n0 + col_base;
CUTE_UNROLL
for (int j = 0; j < BN / 8; ++j) {
const float f0 = acc(j * 4 + rp * 2 + 0);
const float f1 = acc(j * 4 + rp * 2 + 1);
*reinterpret_cast<uint16_t*>(orow + 8 * j) = f32x2_to_bf16x2_to_e4m3x2(f0, f1);
}
}
}
// Staged variant: XOR-swizzle the 16B chunk index by the row so the u16
// stage writes (8 distinct rows per warp) spread across banks, while the
// 16B copy-out reads stay conflict-free row segments. The XOR must be
// masked to the chunks actually present in a BN-wide row.
static constexpr int STAGE_CHUNK_MASK = BN / 16 - 1;
static __device__ __forceinline__ int stage_off(int r, int c) {
const int phys = ((c >> 4) ^ r) & STAGE_CHUNK_MASK;
return r * BN + (phys << 4) + (c & 15);
}
// Stage-write half: fp32 acc -> bf16 -> fp8 u16 writes into the swizzled
// smem stage. Reads only the accumulator, so it can run while the next
// round's WGMMA chain and the B refill are in flight. The caller provides
// the __syncthreads() handoff before stage_flush.
template <typename TC>
static __device__ __forceinline__ void stage_write(TC const& acc, uint8_t* stage, int row_base, int col_base) {
CUTE_UNROLL
for (int rp = 0; rp < 2; ++rp) {
const int row = row_base + 8 * rp; // OOB rows staged but never copied out
CUTE_UNROLL
for (int j = 0; j < BN / 8; ++j) {
const float f0 = acc(j * 4 + rp * 2 + 0);
const float f1 = acc(j * 4 + rp * 2 + 1);
*reinterpret_cast<uint16_t*>(stage + stage_off(row, col_base + 8 * j)) = f32x2_to_bf16x2_to_e4m3x2(f0, f1);
}
}
}
// Copy-out half: coalesced 16B stores of the staged fp8 slab.
static __device__ __forceinline__ void
stage_flush(const uint8_t* stage, uint8_t* gO, int64_t o_s0, int n0, int m_residue, int tid) {
constexpr int CHUNKS_PER_ROW = BN / 16;
constexpr int NUM_CHUNKS = BM * BN / 16;
CUTE_UNROLL
for (int i = 0; i < NUM_CHUNKS / NUM_THREADS; ++i) {
const int chunk = tid + i * NUM_THREADS;
const int r = chunk / CHUNKS_PER_ROW;
const int c = (chunk % CHUNKS_PER_ROW) * 16;
if (r >= m_residue) continue;
const uint4 v = *reinterpret_cast<const uint4*>(stage + stage_off(r, c));
*reinterpret_cast<uint4*>(gO + (int64_t)r * o_s0 + n0 + c) = v;
}
}
// -------------------------------------------------------------------------
// Rope path: out[:, h, 512:576] = fp8(q_rope[:, h, :]). bf16 -> fp32
// (exact) -> fp8 rn/satfinite == the Triton store conversion, so this half
// is bit-exact vs concat_and_cast_q_fp8_pad. 8 bf16 per thread-chunk.
// -------------------------------------------------------------------------
static __device__ __forceinline__ void
rope_path(const bf16* gR, uint8_t* gO, const QprepBf16Fp8Sm90Params& p, int m_residue, int tid) {
const int cthr = tid % 8, rthr = tid / 8;
constexpr int PASSES = BM / LOAD_ROWS_PER_PASS;
if (p.rope_vec16 && p.out_vec16) {
// Fast path: batch-issue every row's uint4 load first so the load
// latencies pipeline (one exposed latency instead of PASSES chained
// load-use stalls), then convert + store.
uint4 raw[PASSES];
CUTE_UNROLL
for (int mi = 0; mi < PASSES; ++mi) {
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
if (row >= m_residue) continue;
raw[mi] = *reinterpret_cast<const uint4*>(gR + (int64_t)row * p.r_s0 + cthr * 8);
}
CUTE_UNROLL
for (int mi = 0; mi < PASSES; ++mi) {
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
if (row >= m_residue) continue;
const uint32_t* w = reinterpret_cast<const uint32_t*>(&raw[mi]);
uint16_t packed[4];
CUTE_UNROLL
for (int i = 0; i < 4; ++i) {
const __nv_bfloat162 v = *reinterpret_cast<const __nv_bfloat162*>(&w[i]);
packed[i] = f32x2_to_e4m3x2_rn_satfinite(__low2float(v), __high2float(v));
}
// out_vec16 guarantees 16B-aligned rows; N_OUT + 8*cthr keeps 8B
// alignment, so the 8-byte chunk goes out as one coalesced store.
*reinterpret_cast<uint64_t*>(gO + (int64_t)row * p.o_s0 + N_OUT + cthr * 8) =
*reinterpret_cast<const uint64_t*>(packed);
}
return;
}
// Unaligned fallback: element strides only guarantee 2B alignment.
CUTE_UNROLL
for (int mi = 0; mi < PASSES; ++mi) {
const int row = rthr + LOAD_ROWS_PER_PASS * mi;
if (row >= m_residue) continue;
const bf16* g = gR + (int64_t)row * p.r_s0 + cthr * 8;
uint8_t* o = gO + (int64_t)row * p.o_s0 + N_OUT + cthr * 8;
const __nv_bfloat16* gh = reinterpret_cast<const __nv_bfloat16*>(g);
CUTE_UNROLL
for (int i = 0; i < 4; ++i) {
const __nv_bfloat162 v = __nv_bfloat162(gh[2 * i], gh[2 * i + 1]);
*reinterpret_cast<uint16_t*>(o + 2 * i) = f32x2_to_e4m3x2_rn_satfinite(__low2float(v), __high2float(v));
}
}
}
// -------------------------------------------------------------------------
// Main device function
// -------------------------------------------------------------------------
static __device__ __forceinline__ void devfunc(const QprepBf16Fp8Sm90Params& p) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 900)
const int m0 = blockIdx.x * BM;
const int h = blockIdx.y;
const int tid = threadIdx.x;
const int m_residue = p.num_tokens - m0; // > 0 by grid construction
extern __shared__ char smem_raw[];
SharedStorage& ss = *reinterpret_cast<SharedStorage*>(smem_raw);
Tensor sA = make_tensor(make_smem_ptr(ss.a.data()), SmemLayoutA{});
Tensor sB0 = make_tensor(make_smem_ptr(ss.b[0].data()), SmemLayoutB{});
Tensor sB1 = make_tensor(make_smem_ptr(ss.b[1].data()), SmemLayoutB{});
const bf16* gA = reinterpret_cast<const bf16*>(p.q_nope) + (int64_t)m0 * p.a_s0 + (int64_t)h * p.a_s1;
const bf16* gB = reinterpret_cast<const bf16*>(p.w_kc) + (int64_t)h * p.b_s0;
const bf16* gR = reinterpret_cast<const bf16*>(p.q_rope) + (int64_t)m0 * p.r_s0 + (int64_t)h * p.r_s1;
uint8_t* gO = reinterpret_cast<uint8_t*>(p.out) + (int64_t)m0 * p.o_s0 + (int64_t)h * p.o_s1;
const int64_t policy_stream = createpolicy_evict_first();
const int64_t policy_keep = createpolicy_evict_last();
// Issue the A block + B slab 0 (group 0), then B slab 1 (group 1), then
// run the rope path over the in-flight async loads.
load_a_tile(sA, gA, p.a_s0, m_residue, tid, policy_stream);
load_b_slab(sB0, gB, p.b_s2, 0, tid, policy_keep);
cp_async_fence();
load_b_slab(sB1, gB, p.b_s2, 1, tid, policy_keep);
cp_async_fence();
// Rope in the prologue: its global-load latency hides under the wait for
// the A/B cp.async stream (measured better than placing it after the
// first gemm commit on the 1-CTA/SM K=192 path).
rope_path(gR, gO, p, m_residue, tid);
cp_async_wait<1>(); // A tile + B0 done; B1 still in flight
__syncthreads();
TiledMMA_t tiled_mma;
const int row_base = (tid / 128) * 64 + ((tid % 128) / 32) * 16 + ((tid % 32) / 4);
const int col_base = (tid % 4) * 2;
{
// Single accumulator: 2-CTA/SM co-residency covers the epilogue
// latency (measured faster than every cross-round dual-accumulator
// pipeline variant, which needs >128 regs and forfeits co-residency);
// the B double-buffer still prefetches slab nb+1 a full round ahead.
Tensor acc = partition_fragment_C(tiled_mma, Shape<Int<BM>, Int<BN>>{});
gemm_ss(tiled_mma, sA, sB0, acc, tid);
warpgroup_commit_batch();
auto round_body = [&](int nb) __attribute__((always_inline)) {
warpgroup_wait<0>(); // gemm(nb) drained
if (p.out_vec16) {
// Single barrier: stage handoff + CTA-wide WGMMA drain of B[nb%2].
stage_write(acc, ss.c_stage.data(), row_base, col_base);
__syncthreads();
if (nb + 2 < N_SLABS) {
load_b_slab((nb % 2 == 0) ? sB0 : sB1, gB, p.b_s2, nb + 2, tid, policy_keep);
cp_async_fence();
}
stage_flush(ss.c_stage.data(), gO, p.o_s0, nb * BN, m_residue, tid);
} else {
__syncthreads();
if (nb + 2 < N_SLABS) {
load_b_slab((nb % 2 == 0) ? sB0 : sB1, gB, p.b_s2, nb + 2, tid, policy_keep);
cp_async_fence();
}
store_slab_direct(acc, gO, p.o_s0, nb * BN, row_base, col_base, m_residue);
}
if (nb + 1 < N_SLABS) {
// Slab nb+1 resident (leave the nb+2 refill in flight, if any),
// then commit the next round's gemm. The barrier also separates
// this round's stage_flush reads from the next stage_write.
if (nb + 2 < N_SLABS) {
cp_async_wait<1>();
} else {
cp_async_wait<0>();
}
__syncthreads();
gemm_ss(tiled_mma, sA, (nb % 2 == 0) ? sB1 : sB0, acc, tid);
warpgroup_commit_batch();
}
};
if constexpr (N_SLABS <= 4) {
CUTE_UNROLL
for (int nb = 0; nb < N_SLABS; ++nb) {
round_body(nb);
}
} else {
// Fully unrolling 8 rounds blows the 128-register budget (2 CTAs/SM
// launch bound) and spills to local memory.
CUTE_NO_UNROLL
for (int nb = 0; nb < N_SLABS; ++nb) {
round_body(nb);
}
}
}
#else
if (cute::thread0()) {
CUTE_INVALID_CONTROL_PATH("qprep_bf16_fp8_sm90 only supports sm90");
}
#endif
}
// -------------------------------------------------------------------------
// Host-side launch
// -------------------------------------------------------------------------
static void run(const QprepBf16Fp8Sm90Params& p) {
QPREP_ASSERT(p.num_tokens > 0);
QPREP_ASSERT(p.num_heads > 0);
auto kernel = &qprep_bf16_fp8_kernel<QprepBf16Fp8Kernel<K_DIM>>;
constexpr size_t smem_size = sizeof(SharedStorage);
static bool attr_set = [&]() {
QPREP_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
return true;
}();
(void)attr_set;
dim3 grid(ceil_div_i(p.num_tokens, BM), p.num_heads, 1);
kernel<<<grid, NUM_THREADS, smem_size, p.stream>>>(p);
QPREP_CUDA_CHECK(cudaGetLastError());
}
};
template <typename Kernel>
__global__ void __launch_bounds__(Kernel::NUM_THREADS, Kernel::MIN_CTAS)
qprep_bf16_fp8_kernel(__grid_constant__ const QprepBf16Fp8Sm90Params params) {
Kernel::devfunc(params);
}
template <int K_DIM>
void run_qprep_bf16_fp8_sm90(const QprepBf16Fp8Sm90Params& params) {
QprepBf16Fp8Kernel<K_DIM>::run(params);
}
} // namespace qprep_sm90
@@ -0,0 +1,52 @@
/* Copyright 2026 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
// Parameters for the SM90 Q8KV8 born-fp8 q-prep kernel (absorbed-q bmm +
// nope/rope concat + fp32 -> bf16 -> fp8_e4m3 cast). All strides are in
// ELEMENTS of the respective tensor's dtype (fp8 strides == byte strides).
#pragma once
#include <cstdint>
#include <cuda_runtime.h>
struct QprepBf16Fp8Sm90Params {
int num_tokens; // T (runtime; m-tiles are masked)
int num_heads; // H (grid dim)
// q_nope: [T, H, K] bf16 (strided view OK; innermost dim contiguous)
const void* q_nope;
int64_t a_s0, a_s1;
// w_kc: [H, K, N] bf16 with K contiguous (stride(1) == 1; production layout
// is (K*N, 1, K), i.e. the N-major absorbed weight)
const void* w_kc;
int64_t b_s0, b_s2;
// q_rope: [T, H, R] bf16 (strided view OK; innermost dim contiguous)
const void* q_rope;
int64_t r_s0, r_s1;
// 16B-aligned rope rows (base pointer and both strides) -> uint4 loads
bool rope_vec16;
// out: [T, pad_heads, N + R] fp8_e4m3; only [:, :H, :] is written
void* out;
int64_t o_s0, o_s1;
// 16B-aligned out rows (base pointer and both strides) -> smem-staged
// coalesced uint4 stores for the nope half (else direct u16 stores)
bool out_vec16;
cudaStream_t stream;
};
@@ -199,4 +199,45 @@ void sparse_prefill_q8kv8_dispatch_full(
_run_q8kv8(params, true, true);
}
void sparse_prefill_q8kv8_dispatch_topk_length(
tvm::ffi::TensorView q,
tvm::ffi::TensorView kv,
tvm::ffi::TensorView indices,
tvm::ffi::TensorView q_scale,
tvm::ffi::TensorView kv_scale,
tvm::ffi::TensorView topk_length,
tvm::ffi::TensorView out,
tvm::ffi::TensorView max_logits,
tvm::ffi::TensorView lse,
int64_t s_q_val,
int64_t s_kv_val,
int64_t h_q_val,
int64_t h_kv_val,
int64_t d_qk_val,
int64_t d_v_val,
int64_t topk_val,
double sm_scale_val,
int64_t cuda_stream) {
SparseMlaQ8Kv8PrefillParams params = _make_common_params(
q,
kv,
indices,
q_scale,
kv_scale,
out,
max_logits,
lse,
s_q_val,
s_kv_val,
h_q_val,
h_kv_val,
d_qk_val,
d_v_val,
topk_val,
sm_scale_val,
cuda_stream);
params.topk_length = static_cast<int*>(topk_length.data_ptr());
_run_q8kv8(params, true, false);
}
} // namespace
@@ -1,3 +1,5 @@
from typing import Optional
import torch
import triton
import triton.language as tl
@@ -285,11 +287,19 @@ def _dequantize_k_cache_paged_kernel(
tl.store(dst_ptr, data, mask=mask)
# Tokens handled by one program of the vectorized gather kernel. 4 tokens
# x 512 fp8 nope elements = 2048 elements per program: with num_warps=4
# (128 threads) that is 16 fp8 elements per thread, which Triton emits as
# a single 16-byte vectorized load/store per thread.
_GATHER_TOKENS_PER_PROG = 4
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,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Gather paged fp8 KV tokens and re-pack into flat [576] fp8 layout.
@@ -300,6 +310,13 @@ def gather_dequant_requant_fp8_paged(
Rope is cast bf16->fp8. The whole operation is fused into a single
Triton kernel to avoid allocating an intermediate bf16 buffer.
The kernel writes EVERY byte of rows [0, num_tokens) and zero-fills
rows [num_tokens, num_tokens + extra_rows) (the -1-sentinel landing
pad required by the SM90 sparse MLA Q8KV8 kernel, which clamps each
-1 topk slot ``offs`` to distinct row ``num_tokens + offs``). The
destination therefore needs no pre-zeroing, which allows passing a
persistent (dirty) buffer via ``out``.
Args:
quant_k_cache: [total_num_tokens, 1, 656] fp8_e4m3fn
page_table_1_flattened: [num_tokens] int32
@@ -308,6 +325,9 @@ def gather_dequant_requant_fp8_paged(
the end of the output (used by the SM90 sparse MLA Q8KV8
kernel which over-reads past end-of-buffer for masked
indices)
out: optional pre-allocated destination of shape
[num_tokens + extra_rows, 1, 576] (or [.., 576]) fp8_e4m3fn,
contiguous. Contents may be arbitrary (fully overwritten).
Returns:
output: [num_tokens + extra_rows, 1, 576] fp8_e4m3fn
"""
@@ -323,11 +343,162 @@ def gather_dequant_requant_fp8_paged(
out_dim = dim_nope + dim_rope # 576
assert num_tiles * group_size == dim_nope
total_rows = num_tokens + extra_rows
if out is None:
# No zero-fill needed: the kernel overwrites every byte of the
# data rows and zero-fills the pad rows itself.
output = torch.empty(
(total_rows, 1, out_dim),
dtype=torch.float8_e4m3fn,
device=quant_k_cache.device,
)
else:
assert out.dtype == torch.float8_e4m3fn
assert out.device == quant_k_cache.device
assert out.is_contiguous()
assert out.numel() == total_rows * out_dim, (
f"out buffer has {out.numel()} elements, expected "
f"{total_rows} x {out_dim} = {total_rows * out_dim}"
)
output = out.view(total_rows, 1, out_dim)
if total_rows == 0:
return output
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)
grid = (triton.cdiv(total_rows, _GATHER_TOKENS_PER_PROG),)
_gather_dequant_requant_fp8_paged_vec_kernel[grid](
output,
input_nope_q,
input_nope_s,
input_rope,
page_table_1_flattened,
num_tokens,
total_rows,
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,
TOKENS_PER_PROG=_GATHER_TOKENS_PER_PROG,
num_warps=4,
)
return output
@triton.jit
def _gather_dequant_requant_fp8_paged_vec_kernel(
output_ptr,
input_nope_q_ptr,
input_nope_s_ptr,
input_rope_ptr,
page_table_1_ptr,
num_tokens: int,
total_rows: int,
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,
TOKENS_PER_PROG: tl.constexpr,
):
"""Vectorized fused gather + dequant(per-group) + requant(per-tensor).
One program handles TOKENS_PER_PROG consecutive output rows (full
576-byte rows each), instead of the legacy one-program-per-(token,
128-elem-slice) layout, so each thread moves 16 contiguous fp8 bytes
per load/store. Rows >= num_tokens (the -1-sentinel landing pad) are
zero-filled without touching the KV cache. Per-element math is
bit-identical to the legacy kernel: fp8 -> f32, * f32 group scale,
-> fp8 (nope); bf16 -> fp8 (rope).
"""
pid = tl.program_id(0)
offs_t = pid * TOKENS_PER_PROG + tl.arange(0, TOKENS_PER_PROG) # [T]
row_in_range = offs_t < total_rows
is_real = offs_t < num_tokens
# Masked lanes (pad rows) never touch memory; `other=0` keeps the
# address arithmetic in-bounds-irrelevant.
paged = tl.load(page_table_1_ptr + offs_t, mask=is_real, other=0).to(tl.int64)
# 64-bit output row offsets: total_rows * 576 can exceed int32 for
# very large gathered buffers.
offs_t64 = offs_t.to(tl.int64)
offs_g = tl.arange(0, NUM_NOPE_BLOCKS) # [G] dequant groups
offs_i = tl.arange(0, GROUP_SIZE) # [I] elems within a group
# a. nope: [T, G, I] fp8 block; the (G, I) plane spans the contiguous
# DIM_NOPE bytes of one cache row.
ptr_q = (
input_nope_q_ptr
+ paged[:, None, None] * input_nope_q_stride_0
+ offs_g[None, :, None] * GROUP_SIZE
+ offs_i[None, None, :]
)
y_q = tl.load(ptr_q, mask=is_real[:, None, None], other=0.0).to(tl.float32)
ptr_s = input_nope_s_ptr + paged[:, None] * input_nope_s_stride_0 + offs_g[None, :]
y_s = tl.load(ptr_s, mask=is_real[:, None], other=0.0)
# dequant -> f32 -> requant to fp8; pad rows: (0 * 0) -> +0 -> byte 0x00
y = (y_q * y_s[:, :, None]).to(tl.float8e4nv)
dst_q = (
output_ptr
+ offs_t64[:, None, None] * output_stride_0
+ offs_g[None, :, None] * GROUP_SIZE
+ offs_i[None, None, :]
)
tl.store(dst_q, y, mask=row_in_range[:, None, None])
# b. rope: [T, R] bf16 -> fp8; pad rows: 0.0 -> byte 0x00
offs_r = tl.arange(0, DIM_ROPE)
src_r = input_rope_ptr + paged[:, None] * input_rope_stride_0 + offs_r[None, :]
data = tl.load(src_r, mask=is_real[:, None], other=0.0).to(tl.float8e4nv)
dst_r = (
output_ptr + offs_t64[:, None] * output_stride_0 + DIM_NOPE + offs_r[None, :]
)
tl.store(dst_r, data, mask=row_in_range[:, None])
def gather_dequant_requant_fp8_paged_legacy(
quant_k_cache: torch.Tensor,
page_table_1_flattened: torch.Tensor,
group_size: int = 128,
extra_rows: int = 0,
) -> torch.Tensor:
"""Legacy (pre-vectorization) gather + dequant + requant.
Kept as the bit-exactness / performance reference for
``gather_dequant_requant_fp8_paged`` (see
``benchmark/kernels/deepseek/benchmark_q8kv8_kv_gather.py``). Allocates and zero-fills the
full destination each call, then launches one program per
(token, 128-elem slice).
"""
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.
# num_tokens for masked indices).
output = torch.zeros(
(total_rows, 1, out_dim),
dtype=torch.float8_e4m3fn,
@@ -419,3 +590,67 @@ def _gather_dequant_requant_fp8_paged_kernel(
if __name__ == "__main__":
raise Exception("UT is in quant_k_cache.py")
@triton.jit
def _concat_cast_kv_fp8_pad_kernel(
out_ptr,
k_ptr,
kr_ptr,
num_tokens,
k_stride,
kr_stride,
NOPE: tl.constexpr,
ROPE: tl.constexpr,
):
"""Row program: real rows write cast(k)||cast(k_rope); pad-band rows
write zeros (the -1-sentinel landing pad the kernel's clamp maps to)."""
row = tl.program_id(0).to(tl.int64)
offs_n = tl.arange(0, NOPE)
offs_r = tl.arange(0, ROPE)
head = NOPE + ROPE
if row < num_tokens:
v_n = tl.load(k_ptr + row * k_stride + offs_n)
tl.store(out_ptr + row * head + offs_n, v_n.to(tl.float8e4nv))
v_r = tl.load(kr_ptr + row * kr_stride + offs_r)
tl.store(out_ptr + row * head + NOPE + offs_r, v_r.to(tl.float8e4nv))
else:
zero_n = tl.zeros([NOPE], dtype=tl.float32).to(tl.float8e4nv)
zero_r = tl.zeros([ROPE], dtype=tl.float32).to(tl.float8e4nv)
tl.store(out_ptr + row * head + offs_n, zero_n)
tl.store(out_ptr + row * head + NOPE + offs_r, zero_r)
def concat_cast_kv_fp8_pad(
out: torch.Tensor,
k: torch.Tensor,
k_rope: torch.Tensor,
num_tokens: int,
) -> torch.Tensor:
"""Fused non-prefix Q8KV8 KV prep: cast-concat k (nope latent) and k_rope
directly into the persistent fp8 kv_buf and zero the trailing pad band —
replaces the bf16 `_cat` materialization + `.copy_` cast + `.zero_()`
tail (3 kernels + one [tokens, 576] bf16 alloc). Same bf16->fp8
store-cast the gather kernel uses (bit-identical bytes).
``out``: [total_rows, 576] fp8 slice (total_rows = num_tokens + pad band);
``k``: [num_tokens, NOPE] bf16 view; ``k_rope``: [num_tokens, ROPE] bf16.
"""
total_rows, head = out.shape
nope = k.shape[-1]
rope = k_rope.shape[-1]
assert head == nope + rope and out.dtype == torch.float8_e4m3fn
k2 = k.view(num_tokens, nope)
kr2 = k_rope.view(num_tokens, rope)
assert k2.stride(-1) == 1 and kr2.stride(-1) == 1
_concat_cast_kv_fp8_pad_kernel[(total_rows,)](
out,
k2,
kr2,
num_tokens,
k2.stride(0),
kr2.stride(0),
NOPE=nope,
ROPE=rope,
)
return out
@@ -7,6 +7,36 @@ from collections.abc import Callable
import torch
def _restore_row_stride(logits: torch.Tensor) -> torch.Tensor:
"""Undo PyTorch's DLPack stride normalization on single-row DeepGEMM logits.
DeepGEMM returns paged-MQA logits as a row-padded view:
``torch.empty(num_rows, aligned_len)[:, :max_len]`` with ``aligned_len``
256-element (1024-byte) aligned. tvm-ffi builds of DeepGEMM (sgl-deep-gemm
>= 0.1.x) round-trip that view through DLPack on return, and PyTorch's
DLPack *exporter* rewrites the stride of every size<2 dim to 1 whenever it
differs from the packed expectation (pytorch/pytorch#83158). A one-row
result whose row is actually padded (``max_len % 256 != 0``, e.g. any
``model context_len + 4`` page-table width at bs=1 decode capture) therefore
arrives with ``stride() == (1, 1)`` instead of ``(aligned_len, 1)``, which
violates the fused top-k v2 kernel ABI (``score_stride % 4 == 0``, enforced
both in ``dsa_topk_backend._topk_transform_v2_paged`` and by the kernel's
own RuntimeCheck).
For ``num_rows <= 1`` the row stride is semantically arbitrary (row 0 is
the only row ever addressed and ``(num_rows - 1) * stride(0)`` contributes
nothing to the storage extent), so restoring a 16-byte-aligned value is a
pure metadata rewrite: same storage, same data pointer, no copy and no
kernel launch -- trivially CUDA-graph-capture-safe. Multi-row results keep
their true strides through DLPack (no size<2 dim) and pass through
untouched.
"""
if logits.shape[0] <= 1 and logits.stride(0) % 4 != 0:
width = logits.shape[1]
logits = logits.as_strided((logits.shape[0], width), ((width + 3) // 4 * 4, 1))
return logits
def deepgemm_paged_mqa_logits_native(
fp8_paged_mqa_logits_fn: Callable[..., torch.Tensor],
q_fp8: torch.Tensor,
@@ -23,7 +53,7 @@ def deepgemm_paged_mqa_logits_native(
) -> torch.Tensor:
# block_tables[::next_n] de-expands the caller's repeat_interleave without a
# copy (DeepGEMM only checks `stride(1) == 1`).
return fp8_paged_mqa_logits_fn(
logits = fp8_paged_mqa_logits_fn(
q_fp8[:q_offset].view(B, next_n, q_fp8.shape[1], q_fp8.shape[2]),
kv_cache_fp8,
weights[:q_offset],
@@ -33,6 +63,7 @@ def deepgemm_paged_mqa_logits_native(
max_seq_len,
clean_logits=False,
)
return _restore_row_stride(logits)
def deepgemm_paged_mqa_logits_split(
@@ -48,7 +79,7 @@ def deepgemm_paged_mqa_logits_split(
q_offset: int,
) -> torch.Tensor:
q_fp8 = q_fp8.unsqueeze(1)
return fp8_paged_mqa_logits_fn(
logits = fp8_paged_mqa_logits_fn(
q_fp8[:q_offset],
kv_cache_fp8,
weights[:q_offset],
@@ -58,6 +89,7 @@ def deepgemm_paged_mqa_logits_split(
max_seq_len,
clean_logits=False,
)
return _restore_row_stride(logits)
def aiter_paged_mqa_logits(
@@ -0,0 +1,134 @@
"""JIT-compiled SM90 (Hopper) kernel for the Q8KV8 born-fp8 q-prep.
Fuses the per-head absorbed-q bmm (q_nope [T, H, K] bf16 x w_kc [H, K, N]
bf16, fp32 accumulate), the nope/rope concat, and the bf16 -> fp8_e4m3 cast
into one hand-written WGMMA kernel. CUDA replacement for the Triton
``absorbed_bmm_concat_cast_q_fp8`` (triton_ops/cache_ops.py) with the
identical fp32 -> bf16 -> fp8 epilogue rounding chain; the rope half is
bit-exact vs ``concat_and_cast_q_fp8_pad``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.kernel_api_logging import debug_kernel_api
from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch
if TYPE_CHECKING:
from tvm_ffi.module import Module
N_LORA = 512 # kv_lora_rank (nope output dim)
ROPE_DIM = 64 # qk_rope_head_dim
@cache_once
def _jit_qprep_bf16_fp8_module() -> Module:
if torch.cuda.get_device_capability()[0] != 9:
raise RuntimeError("qprep_bf16_fp8_sm90 requires an SM90 (Hopper) GPU")
with override_jit_cuda_arch(9, 0, "a"):
return load_jit(
"qprep_bf16_fp8_sm90",
cuda_files=["qprep_bf16_fp8_sm90/entry.cuh"],
cuda_wrappers=[("dispatch", "qprep_bf16_fp8_dispatch")],
# Same minimal flag set as the sparse_mla_q8kv8_prefill_sm90 JIT
# build (per-flag ablation there showed the rest are no-ops).
extra_cuda_cflags=[
"-O3",
"-DNDEBUG",
"-DCUTE_USE_PACKED_TUPLE=1",
"-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1",
"--use_fast_math",
],
extra_dependencies=["cutlass"],
)
# torch._C._cuda_getCurrentRawStream returns the cudaStream_t pointer expected
# by the JIT wrapper (see sparse_mla_q8kv8_prefill_sm90.py).
_get_current_stream_raw = torch._C._cuda_getCurrentRawStream
@debug_kernel_api
def q8kv8_qprep_fwd(
q_fp8_pad: torch.Tensor,
q_nope: torch.Tensor,
w_kc: torch.Tensor,
q_rope: torch.Tensor,
num_heads: int,
) -> None:
"""Fused absorbed-q bmm + nope/rope concat + bf16->fp8 cast ("born fp8" q).
Mirrors the contract of ``absorbed_bmm_concat_cast_q_fp8``:
* ``q_fp8_pad``: [num_tokens, pad_heads, N + ROPE] fp8_e4m3 destination;
only ``[:, :num_heads, :]`` is written.
* ``q_nope``: [num_tokens, H, K] bf16 pre-absorb q (strided views OK).
* ``w_kc``: [H, K, N] bf16 absorbed weight with K contiguous
(``stride(1) == 1``, the production N-major layout).
* ``q_rope``: [num_tokens, H, ROPE] bf16 post-rope q (strided views OK).
K (``qk_nope_head_dim``) must be 128 or 192. Extra restrictions vs the
Triton kernel (all satisfied by the production layouts): 16-byte aligned
q_nope/w_kc base pointers, q_nope/w_kc strides that are multiples of 8
elements, and even q_fp8_pad row/head strides.
"""
num_tokens, _, k_dim = q_nope.shape
n_dim = w_kc.shape[-1]
rope_dim = q_rope.shape[-1]
assert q_fp8_pad.dtype == torch.float8_e4m3fn
assert q_nope.dtype == torch.bfloat16 and w_kc.dtype == torch.bfloat16
assert q_rope.dtype == torch.bfloat16
assert q_nope.is_cuda and w_kc.is_cuda and q_rope.is_cuda and q_fp8_pad.is_cuda
assert q_nope.shape[1] == num_heads and q_rope.shape[1] == num_heads
assert w_kc.shape[0] == num_heads and w_kc.shape[1] == k_dim
assert q_fp8_pad.shape[0] >= num_tokens and q_fp8_pad.shape[1] >= num_heads
assert q_fp8_pad.shape[2] == n_dim + rope_dim
assert k_dim in (128, 192), "CUDA q-prep supports K in {128, 192}"
assert n_dim == N_LORA and rope_dim == ROPE_DIM
# Innermost-contiguous requirements (same as the Triton kernel).
assert q_nope.stride(2) == 1 and q_rope.stride(2) == 1
assert q_fp8_pad.stride(2) == 1
# CUDA-kernel-specific layout requirements (production layouts satisfy
# all of these; the Triton kernel stays the general-strides fallback).
assert w_kc.stride(1) == 1, "w_kc must have K contiguous (N-major layout)"
assert q_nope.data_ptr() % 16 == 0 and w_kc.data_ptr() % 16 == 0
assert q_nope.stride(0) % 8 == 0 and q_nope.stride(1) % 8 == 0
assert w_kc.stride(0) % 8 == 0 and w_kc.stride(2) % 8 == 0
assert q_fp8_pad.stride(0) % 2 == 0 and q_fp8_pad.stride(1) % 2 == 0
rope_vec16 = (
q_rope.data_ptr() % 16 == 0
and q_rope.stride(0) % 8 == 0
and q_rope.stride(1) % 8 == 0
)
out_vec16 = (
q_fp8_pad.data_ptr() % 16 == 0
and q_fp8_pad.stride(0) % 16 == 0
and q_fp8_pad.stride(1) % 16 == 0
)
module = _jit_qprep_bf16_fp8_module()
module.dispatch(
q_nope,
w_kc,
q_rope,
q_fp8_pad,
num_tokens,
num_heads,
k_dim,
q_nope.stride(0),
q_nope.stride(1),
w_kc.stride(0),
w_kc.stride(2),
q_rope.stride(0),
q_rope.stride(1),
q_fp8_pad.stride(0),
q_fp8_pad.stride(1),
int(rope_vec16),
int(out_vec16),
_get_current_stream_raw(q_nope.device.index),
)
@@ -68,6 +68,7 @@ def _jit_sparse_mla_q8kv8_prefill_module() -> Module:
cuda_wrappers=[
("dispatch", "sparse_prefill_q8kv8_dispatch"),
("dispatch_full", "sparse_prefill_q8kv8_dispatch_full"),
("dispatch_topk_length", "sparse_prefill_q8kv8_dispatch_topk_length"),
],
extra_cuda_cflags=_q8kv8_cuda_flags(),
extra_dependencies=["cutlass"],
@@ -86,6 +87,7 @@ def _get_entries() -> tuple:
_resolved_entries = (
m["dispatch"],
m["dispatch_full"],
m["dispatch_topk_length"],
)
return _resolved_entries
@@ -146,7 +148,7 @@ def _sparse_mla_q8kv8_prefill_op(
sm_scale: float,
cuda_stream: int,
) -> None:
dispatch_fn, _ = _get_entries()
dispatch_fn, _, _ = _get_entries()
dispatch_fn(
q,
kv,
@@ -193,7 +195,7 @@ def _sparse_mla_q8kv8_prefill_full_op(
sm_scale: float,
cuda_stream: int,
) -> None:
_, dispatch_full_fn = _get_entries()
_, dispatch_full_fn, _ = _get_entries()
dispatch_full_fn(
q,
kv,
@@ -217,6 +219,53 @@ def _sparse_mla_q8kv8_prefill_full_op(
)
@register_custom_op(
op_name="sparse_mla_q8kv8_prefill_topk_length",
mutates_args=["out", "max_logits", "lse"],
)
def _sparse_mla_q8kv8_prefill_topk_length_op(
q: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
q_scale: torch.Tensor,
kv_scale: torch.Tensor,
topk_length: torch.Tensor,
out: torch.Tensor,
max_logits: torch.Tensor,
lse: torch.Tensor,
s_q: int,
s_kv: int,
h_q: int,
h_kv: int,
d_qk: int,
d_v: int,
topk: int,
sm_scale: float,
cuda_stream: int,
) -> None:
_, _, dispatch_topk_length_fn = _get_entries()
dispatch_topk_length_fn(
q,
kv,
indices,
q_scale,
kv_scale,
topk_length,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
@debug_kernel_api
def sparse_mla_q8kv8_prefill_fwd(
q: torch.Tensor, # [s_q, h_q, d_qk], float8_e4m3fn
@@ -256,8 +305,8 @@ def sparse_mla_q8kv8_prefill_fwd(
f"sparse_mla_q8kv8_prefill_fwd only supports d_v=512, got {d_v}"
)
if (attn_sink is None) != (topk_length is None):
raise ValueError("attn_sink and topk_length must be provided together")
if attn_sink is not None and topk_length is None:
raise ValueError("attn_sink requires topk_length to be provided as well")
device = q.device
if out is None:
@@ -305,6 +354,27 @@ def sparse_mla_q8kv8_prefill_fwd(
sm_scale,
cuda_stream,
)
elif topk_length is not None:
_sparse_mla_q8kv8_prefill_topk_length_op(
q,
kv,
indices,
q_scale,
kv_scale,
topk_length,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
else:
_sparse_mla_q8kv8_prefill_op(
q,
@@ -30,6 +30,9 @@ 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,
)
from sglang.kernels.ops.kvcache.cache_ops import (
q8kv8_topk_length_from_indices as q8kv8_topk_length_from_indices,
)
from sglang.kernels.ops.kvcache.cache_ops import (
reshape_and_cache_flash as reshape_and_cache_flash,
)
@@ -322,6 +322,396 @@ def concat_and_cast_q_fp8_pad(q_fp8_pad, q_nope, q_rope, num_heads):
)
@triton.jit
def absorbed_bmm_concat_cast_q_fp8_kernel(
qout_ptr, # [num_tokens, pad_heads, N+ROPE] fp8 (dst; only [:, :H, :] written)
a_ptr, # q_nope (pre-absorb) [num_tokens, H, K] bf16
b_ptr, # w_kc [H, K, N] bf16 (any strides; typically N-major)
rope_ptr, # q_rope (post-rope) [num_tokens, H, ROPE] bf16
T, # num_tokens (runtime; masked)
qout_s0,
qout_s1,
a_s0,
a_s1,
b_s0,
b_s1,
b_s2,
rope_s0,
rope_s1,
K: tl.constexpr,
N: tl.constexpr,
ROPE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
K_MODE: tl.constexpr,
):
# One program per (token-block, head): q_out[m, h, :N] = fp8(bf16(fp32(
# q_nope[m, h, :K] @ w_kc[h, :K, :N]))) and q_out[m, h, N:] =
# fp8(q_rope[m, h, :ROPE]). This makes q "born fp8": the absorbed bmm,
# the nope/rope concat, and the bf16->fp8 cast collapse into one kernel,
# so neither the bf16 q_nope_out ([H, T, N], written by cublas and re-read
# by the concat-cast) nor the standalone concat-cast launch exist anymore.
#
# K handling (K_MODE selects the codegen for the nope-gemm K dimension;
# every mode keeps the same fp32-accumulator -> bf16 -> fp8 epilogue):
# 0 "single": BLOCK_K == K. Preload the whole [BLOCK_M, K] a-tile once,
# one tl.dot per N-block — identical codegen to the original
# power-of-2-only kernel (DeepSeek K=128). For non-power-of-2 K this
# only compiles if the Triton build allows non-power-of-2 tl.arange
# (Triton <= 3.5.x does NOT: "arange's range must be a power of 2").
# 1 "loop": split-K loop, K % BLOCK_K == 0, BLOCK_K power of 2 >= 16
# (e.g. K=192 with BLOCK_K=64 -> 3 iterations). The a-tile is
# re-loaded per (N-block, K-block); slices are L1/L2-resident after
# the first N-block, but the load/dot interleave costs bandwidth
# (measured ~1372 GB/s vs ~2x that for mode 0 at K=128).
# 2 "two_dot": K = BLOCK_K + (K - BLOCK_K), both power-of-2 halves
# (192 = 128 + 64). Both a-tiles preload once before the N-loop;
# each N-block issues two chained tl.dot into one fp32 accumulator.
# No K-loop, no a re-reads — the direct generalization of mode 0.
# 3 "three_dot": K = 3 * BLOCK_K (192 = 3 x 64). Same as mode 2 with
# three preloaded a-tiles / three chained tl.dot per N-block; the
# hoisted-loads analogue of mode 1 (identical fp32 add order).
# 4 "pad": BLOCK_K = next_pow2(K) > K, k-masked loads (zero fill).
# Single tl.dot per N-block; the padded zeros are exact fp32
# additive identities so the result matches a K-wide single dot,
# at the cost of BLOCK_K/K (e.g. 256/192 = 1.33x) extra MMA work.
#
# Rounding contract: the fp32 accumulator is rounded to bf16 first (the
# same output rounding stage as the cublas bf16 bmm) and then converted
# bf16->fp8 by the same implicit-store conversion the fused concat-cast
# kernel uses. The split-K accumulator stays fp32 across all K-blocks,
# so the rounding stages are identical in both layouts. The rope half is
# a bit-exact copy of that kernel (loads the post-rope bf16, converts on
# store). The nope half is NOT guaranteed bit-exact vs the default path:
# tl.dot accumulates fp32 in a different order than cublas, so last-ulp
# fp32 differences can occasionally flip the bf16 (and hence fp8)
# rounding.
pid_m = tl.program_id(0)
h = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
m_mask = offs_m < T
# token-row offsets in int64: T * row-stride can exceed int32 (e.g. 128
# heads x 576 dims x tens of thousands of tokens).
offs_m64 = offs_m.to(tl.int64)
qout_head = qout_ptr + offs_m64[:, None] * qout_s0 + h * qout_s1
a_row = a_ptr + offs_m64[:, None] * a_s0 + h * a_s1
b_head = b_ptr + h * b_s0
if K_MODE == 0:
# single-dot path (original kernel): BLOCK_K == K, preload a once.
offs_k = tl.arange(0, BLOCK_K)
a = tl.load(a_row + offs_k[None, :], mask=m_mask[:, None], other=0.0)
for nb in tl.static_range(N // BLOCK_N):
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
b = tl.load(b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2)
acc = tl.dot(a, b) # fp32 accumulator
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
elif K_MODE == 2:
# two-dot preload: K split as BLOCK_K + (K - BLOCK_K), no K-loop.
offs_k0 = tl.arange(0, BLOCK_K)
offs_k1 = BLOCK_K + tl.arange(0, K - BLOCK_K)
a0 = tl.load(a_row + offs_k0[None, :], mask=m_mask[:, None], other=0.0)
a1 = tl.load(a_row + offs_k1[None, :], mask=m_mask[:, None], other=0.0)
for nb in tl.static_range(N // BLOCK_N):
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
b0 = tl.load(b_head + offs_k0[:, None] * b_s1 + offs_n[None, :] * b_s2)
b1 = tl.load(b_head + offs_k1[:, None] * b_s1 + offs_n[None, :] * b_s2)
acc = tl.dot(a0, b0) # fp32 accumulator
acc = tl.dot(a1, b1, acc) # chained: stays fp32 across both dots
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
elif K_MODE == 3:
# three-dot preload: K = 3 * BLOCK_K, a-tiles hoisted out of the N-loop.
offs_k0 = tl.arange(0, BLOCK_K)
offs_k1 = BLOCK_K + offs_k0
offs_k2 = 2 * BLOCK_K + offs_k0
a0 = tl.load(a_row + offs_k0[None, :], mask=m_mask[:, None], other=0.0)
a1 = tl.load(a_row + offs_k1[None, :], mask=m_mask[:, None], other=0.0)
a2 = tl.load(a_row + offs_k2[None, :], mask=m_mask[:, None], other=0.0)
for nb in tl.static_range(N // BLOCK_N):
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
b0 = tl.load(b_head + offs_k0[:, None] * b_s1 + offs_n[None, :] * b_s2)
b1 = tl.load(b_head + offs_k1[:, None] * b_s1 + offs_n[None, :] * b_s2)
b2 = tl.load(b_head + offs_k2[:, None] * b_s1 + offs_n[None, :] * b_s2)
acc = tl.dot(a0, b0) # fp32 accumulator
acc = tl.dot(a1, b1, acc)
acc = tl.dot(a2, b2, acc) # same fp32 add order as the K_MODE=1 loop
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
elif K_MODE == 4:
# padded single dot: BLOCK_K = next_pow2(K), zero-fill the k tail.
offs_k = tl.arange(0, BLOCK_K)
k_mask = offs_k < K
a = tl.load(
a_row + offs_k[None, :],
mask=m_mask[:, None] & k_mask[None, :],
other=0.0,
)
for nb in tl.static_range(N // BLOCK_N):
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
b = tl.load(
b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2,
mask=k_mask[:, None],
other=0.0,
)
acc = tl.dot(a, b) # fp32 accumulator (padded zeros add exactly 0)
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
else:
# K_MODE == 1: split-K loop (K % BLOCK_K == 0, e.g. K=192, BLOCK_K=64).
for nb in tl.static_range(N // BLOCK_N):
offs_n = nb * BLOCK_N + tl.arange(0, BLOCK_N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for kb in tl.static_range(K // BLOCK_K):
offs_k = kb * BLOCK_K + tl.arange(0, BLOCK_K)
a = tl.load(a_row + offs_k[None, :], mask=m_mask[:, None], other=0.0)
b = tl.load(b_head + offs_k[:, None] * b_s1 + offs_n[None, :] * b_s2)
acc = tl.dot(a, b, acc) # fp32 accumulator across K-blocks
val = acc.to(tl.bfloat16) # cublas-equivalent bf16 output rounding
# implicit bf16 -> fp8 conversion on store (same as the concat-cast)
tl.store(qout_head + offs_n[None, :], val, mask=m_mask[:, None])
offs_r = tl.arange(0, ROPE)
r = tl.load(
rope_ptr + offs_m64[:, None] * rope_s0 + h * rope_s1 + offs_r[None, :],
mask=m_mask[:, None],
other=0.0,
)
tl.store(qout_head + N + offs_r[None, :], r, mask=m_mask[:, None])
# Non-power-of-2-K variant used by variant="auto" (power-of-2 K always takes
# the single-dot fast path). Set to the winner of the K=192 A/B in
# benchmark/kernels/deepseek/benchmark_q8kv8_q_prep.py; "loop" = the pre-A/B split-K behavior.
_AUTO_NONPOW2_VARIANT = "two_dot"
def _qprep_env_variant():
from sglang.srt.environ import envs
return envs.SGLANG_OPT_Q8KV8_QPREP_VARIANT.get()
# Resolved once at import (matches the module-constant style above). "auto"
# keeps the per-K dispatch; "cuda" routes every shape to the hand-written
# SM90 WGMMA kernel (bitwise-identical to two_dot; 1.16-1.38x faster).
_ENV_QPREP_VARIANT = None
def absorbed_bmm_concat_cast_q_fp8(
q_fp8_pad: "torch.Tensor",
q_nope: "torch.Tensor",
w_kc: "torch.Tensor",
q_rope: "torch.Tensor",
num_heads: int,
block_m: int = 128,
block_n: int = 64,
variant: str = "auto",
block_k: int = 0,
num_warps: int = 8,
num_stages: int = 0,
):
"""Fused absorbed-q bmm + nope/rope concat + bf16->fp8 cast ("born fp8" q).
Replaces ``torch.bmm(q_nope.transpose(0, 1), w_kc).transpose(0, 1)``
followed by ``concat_and_cast_q_fp8_pad`` on the Q8KV8 sparse-prefill
path, writing the active ``[:, :num_heads, :]`` slice of the padded fp8 q
buffer directly. Inputs:
* ``q_fp8_pad``: [num_tokens, pad_heads, N + ROPE] fp8_e4m3 destination.
* ``q_nope``: [num_tokens, H, K] bf16 pre-absorb q (strided views OK).
* ``w_kc``: [H, K, N] bf16 absorbed weight (any strides).
* ``q_rope``: [num_tokens, H, ROPE] bf16 post-rope q (strided views OK).
The rope half is bit-exact vs ``concat_and_cast_q_fp8_pad``. The nope
half keeps the same rounding stages (fp32 accum -> bf16 -> fp8) but a
different fp32 accumulation order than cublas, so it is near- but not
guaranteed bit-exact; keep this path behind
``SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q``.
K (``qk_nope_head_dim``) supports any multiple of 16 in [16, 256].
Power-of-2 K (DeepSeek 128) always takes the preload-once
single-``tl.dot`` path. For other K (GLM 192), ``variant`` selects the
K-dimension codegen (every variant keeps the identical fp32 -> bf16 ->
fp8 epilogue):
* ``"auto"``: the current production choice (see
``_AUTO_NONPOW2_VARIANT``).
* ``"loop"``: split-K accumulator loop, ``BLOCK_K`` = ``block_k`` or the
largest power-of-2 divisor of K capped at 128 (192 -> 64 x 3).
* ``"two_dot"``: preload a as two power-of-2 tiles (192 = 128 + 64), two
chained ``tl.dot`` per N-block, no K-loop.
* ``"three_dot"``: preload a as three K/3 tiles (192 = 3 x 64), three
chained ``tl.dot`` per N-block; same fp32 add order as ``"loop"``.
* ``"pad"``: single ``tl.dot`` with ``BLOCK_K`` = next_pow2(K) (192 ->
256) and zero-masked k tails.
* ``"single_k"``: single ``tl.dot`` with ``BLOCK_K`` == K. Only
compiles if the Triton build supports non-power-of-2 ``tl.arange``
(Triton <= 3.5.x raises "arange's range must be a power of 2").
``block_m`` / ``block_n`` / ``num_warps`` / ``num_stages`` are tuning
knobs for the microbench sweep (0 = Triton default for ``num_stages``).
"""
num_tokens, _, k_dim = q_nope.shape
n_dim = w_kc.shape[-1]
rope_dim = q_rope.shape[-1]
assert q_fp8_pad.dtype == torch.float8_e4m3fn
assert q_nope.dtype == torch.bfloat16 and w_kc.dtype == torch.bfloat16
assert q_rope.dtype == torch.bfloat16
assert q_nope.shape[1] == num_heads and q_rope.shape[1] == num_heads
assert w_kc.shape[0] == num_heads and w_kc.shape[1] == k_dim
assert q_fp8_pad.shape[0] >= num_tokens and q_fp8_pad.shape[1] >= num_heads
assert q_fp8_pad.shape[2] == n_dim + rope_dim
# tl.arange / tl.dot constraints
assert (
k_dim % 16 == 0 and 16 <= k_dim <= 256
), "K must be a multiple of 16 in [16, 256]"
assert (rope_dim & (rope_dim - 1)) == 0, "ROPE must be a power of two"
assert n_dim % block_n == 0, "N must be a multiple of block_n"
assert q_nope.stride(2) == 1 and q_rope.stride(2) == 1
assert q_fp8_pad.stride(2) == 1
# Env override for production dispatch (SGLANG_OPT_Q8KV8_QPREP_VARIANT):
# "auto" (default) keeps the per-K Triton dispatch; "cuda" routes every
# shape to the WGMMA kernel below.
global _ENV_QPREP_VARIANT
if _ENV_QPREP_VARIANT is None:
_ENV_QPREP_VARIANT = _qprep_env_variant()
if variant == "auto" and _ENV_QPREP_VARIANT != "auto":
variant = _ENV_QPREP_VARIANT
# Hand-written SM90 WGMMA kernel (opt-in only; "auto" never routes here).
# Same fp32 -> bf16 -> fp8 epilogue; bitwise identical to "two_dot" on
# SM90. Requires K in {128, 192} and the production N-major w_kc layout
# (see the wrapper's asserts); the Triton variants remain the
# general-strides fallback.
_valid = ("auto", "cuda", "loop", "two_dot", "three_dot", "pad", "single_k")
if variant not in _valid:
raise ValueError(
f"unknown q-prep variant {variant!r} "
f"(SGLANG_OPT_Q8KV8_QPREP_VARIANT); valid: {_valid}"
)
if variant == "cuda":
from sglang.kernels.ops.attention.qprep_bf16_fp8_sm90 import q8kv8_qprep_fwd
q8kv8_qprep_fwd(q_fp8_pad, q_nope, w_kc, q_rope, num_heads)
return
# Resolve (K_MODE, BLOCK_K) from the variant; see the kernel's K-handling
# comment for what each mode compiles to.
if k_dim & (k_dim - 1) == 0:
# power-of-2 K: every variant collapses to the single-dot fast path.
k_mode, blk_k = 0, k_dim
else:
v = _AUTO_NONPOW2_VARIANT if variant == "auto" else variant
if v == "loop":
# Largest power-of-2 divisor of K, capped at 128 (K % 16 == 0
# makes this >= 16), unless the caller pinned block_k.
blk_k = block_k or min(k_dim & -k_dim, 128)
assert (
k_dim % blk_k == 0 and blk_k & (blk_k - 1) == 0 and blk_k >= 16
), "loop needs BLOCK_K a power-of-2 divisor of K >= 16"
k_mode = 1
elif v == "two_dot":
blk_k = 1 << (k_dim.bit_length() - 1) # largest power of 2 < K
k1 = k_dim - blk_k
assert (
k1 & (k1 - 1) == 0 and k1 >= 16
), "two_dot needs K = pow2 + pow2 with both halves >= 16"
k_mode = 2
elif v == "three_dot":
blk_k = k_dim // 3
assert (
k_dim % 3 == 0 and blk_k & (blk_k - 1) == 0 and blk_k >= 16
), "three_dot needs K = 3 * pow2 with pow2 >= 16"
k_mode = 3
elif v == "pad":
blk_k = 1 << k_dim.bit_length() # next power of 2 above K
k_mode = 4
elif v == "single_k":
# Non-power-of-2 BLOCK_K == K: compiles only on Triton builds
# that allow non-power-of-2 tl.arange (not 3.5.x).
blk_k = k_dim
k_mode = 0
else:
raise ValueError(f"unknown absorbed-bmm K variant: {variant!r}")
extra = {"num_stages": num_stages} if num_stages else {}
grid = (triton.cdiv(num_tokens, block_m), num_heads)
absorbed_bmm_concat_cast_q_fp8_kernel[grid](
q_fp8_pad,
q_nope,
w_kc,
q_rope,
num_tokens,
q_fp8_pad.stride(0),
q_fp8_pad.stride(1),
q_nope.stride(0),
q_nope.stride(1),
w_kc.stride(0),
w_kc.stride(1),
w_kc.stride(2),
q_rope.stride(0),
q_rope.stride(1),
K=k_dim,
N=n_dim,
ROPE=rope_dim,
BLOCK_M=block_m,
BLOCK_N=block_n,
BLOCK_K=blk_k,
K_MODE=k_mode,
num_warps=num_warps,
**extra,
)
@triton.jit
def q8kv8_topk_length_backscan_kernel(
indices_ptr,
out_ptr,
stride_row,
topk,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
base = indices_ptr + row * stride_row
off = topk
length = 1
found = 0
while (found == 0) & (off > 0):
off -= BLOCK
idx = off + tl.arange(0, BLOCK)
vals = tl.load(base + idx)
pos = tl.max(tl.where(vals >= 0, idx, -1), axis=0)
found = tl.where(pos >= 0, 1, found)
length = tl.where(pos >= 0, pos + 1, length)
tl.store(out_ptr + row, length)
def q8kv8_topk_length_from_indices(indices: torch.Tensor) -> torch.Tensor:
"""Per-row valid-topk count = last non-negative position + 1 (min 1).
``indices``: [s_q, topk] int32 topk output whose pad slots are -1.
Backward block scan per row: the loop exits at the first block holding a
valid entry, so the cost is proportional to the trailing pad run — one
block (~topk/4 elements) for rows with a full topk, which dominate long
contexts. Semantics match the unfused ``(indices >= 0) * ramp).amax``
derivation exactly, including all-pad rows (length 1: one pad-only block
keeps the kernel on its clamp+mask path, contributing zero).
"""
s_q, topk = indices.shape
assert indices.dtype == torch.int32 and indices.stride(1) == 1
out = torch.empty(s_q, dtype=torch.int32, device=indices.device)
block = 512 if topk % 512 == 0 else (256 if topk % 256 == 0 else 128)
q8kv8_topk_length_backscan_kernel[(s_q,)](
indices,
out,
indices.stride(0),
topk,
BLOCK=block,
)
return out
# ---------------------------------------------------------------------------
# Decode Context Parallel (DCP) helpers.
#
@@ -4,10 +4,12 @@ from typing import Optional, Tuple
import torch
import triton
from sglang.srt.environ import envs
from sglang.srt.utils import ceil_div, is_cuda, is_musa
logger = logging.getLogger(__name__)
_is_cuda = is_cuda()
_is_musa = is_musa()
@@ -1539,6 +1541,25 @@ def moe_ep_deepgemm_preprocess(
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
# For masked grouped GEMM, shape M should be multiple of the block M (current block M: {block_m}) https://github.com/deepseek-ai/DeepGEMM/blob/main/deep_gemm/jit_kernels/m_grouped_gemm.py#L165
m_max = (hidden_states.size(0) // 256 + 1) * 256
if (
envs.SGLANG_OPT_DG_MASKED_M_CAP.get()
and not torch.cuda.is_current_stream_capturing()
):
# (capture guard: decode CUDA-graph capture also routes through this
# preprocess; the D2H sync is illegal mid-capture, and decode batches
# are small enough that the uncapped m_max is harmless there.)
# m_max reserves capacity for ALL rank tokens in EVERY local expert:
# the [num_local_experts, m_max, *] masked-GEMM intermediates reach
# 7+ GiB per 32k-token chunk and OOM saturated serving. The hottest
# expert only ever holds max(masked_m) rows, so cap the padded
# capacity there (rounded up to the DeepGEMM block-M). Costs one
# probe dispatch-index launch + one D2H sync per MoE layer;
# correctness is unconditional (m_cap >= max(masked_m) by
# construction, and the final src2dst below is built with the same
# capped stride).
masked_m_probe, _ = fused_moe_dispatch_index(topk_ids, num_local_experts, m_max)
m_cap = (int(masked_m_probe.max().item()) + 255) // 256 * 256
m_max = min(m_max, max(m_cap, 256))
expected_m = (topk_ids.numel() - 1) // num_local_experts + 1
masked_m, src2dst = fused_moe_dispatch_index(topk_ids, num_local_experts, m_max)
@@ -768,7 +768,13 @@ def invoke_fused_moe_kernel(
# activation block-wise fp8 quantization
assert len(block_shape) == 2
block_n, block_k = block_shape[0], block_shape[1]
if _is_cuda:
if A.dtype == torch.float8_e4m3fn:
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): the
# caller already ran the per-token-group quant; A_scale holds
# the matching scales (row- or column-major, strides are
# passed to the kernel below).
assert A_scale is not None
elif _is_cuda:
A, A_scale = sglang_per_token_group_quant_fp8(A, block_k)
else:
A, A_scale = per_token_group_quant_fp8(A, block_k)
+57
View File
@@ -560,6 +560,14 @@ class Envs:
# symmetric-memory kernel), OFF elsewhere (would fall back to RCCL); override
# explicitly to force on/off on any platform.
SGLANG_DP_USE_REDUCE_SCATTER = EnvBool(_default_hip)
# Quantize the variable-length DP-MoE gather payload (SGLANG_DP_USE_GATHERV
# path, prefill/extend only) to fp8-e4m3 with per-token-group-128 scales:
# halves the gathered hidden-state bytes over NCCL; the combine
# (reduce_scatterv) leg stays bf16 (NCCL SUM cannot run on fp8). Lossy on
# the wire — same group quantization the MoE expert GEMMs apply to their
# input anyway, but router/shared-expert reads see rounded values, so this
# stays accuracy-gated and default OFF.
SGLANG_ENABLE_DP_GATHER_FP8 = EnvBool(False)
SGLANG_USE_AITER_UNIFIED_ATTN = EnvBool(False)
# Select the gate/up tile layout for AITER MoE: True -> interleave
# (matches FlyDSL `gate_mode="interleave"` kernels), False -> separated
@@ -689,6 +697,18 @@ class Envs:
# DeepGemm
SGLANG_ENABLE_JIT_DEEPGEMM = EnvBool(True)
# Cap the DeepGEMM masked grouped-GEMM per-expert padded capacity at
# round_up(max(masked_m), 256) instead of round_up(rank_tokens, 256):
# shrinks the [num_local_experts, m, *] MoE intermediates ~4x under
# load imbalance (they otherwise OOM saturated --moe-runner-backend
# deep_gemm serving). Costs one D2H sync per MoE layer.
SGLANG_OPT_DG_MASKED_M_CAP = EnvBool(False)
# Drop dp-attention MAX_LEN pad rows from MoE dispatch (StandardDispatcher
# post-translation topk_ids -> -1): pad rows otherwise run the router on
# stale hidden values and burn expert compute whose outputs are discarded;
# colliding pad top-ks also inflate the DeepGEMM masked-GEMM workspace to
# OOM at saturation. Capture-safe (reads only global_num_tokens_gpu).
SGLANG_OPT_MASK_DP_PAD_MOE = EnvBool(False)
SGLANG_JIT_DEEPGEMM_PRECOMPILE = EnvBool(True)
SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False)
SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4)
@@ -733,6 +753,35 @@ class Envs:
SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False)
SGLANG_DSA_TOPK_BROADCAST = EnvBool(False)
SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False)
# Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: fuse the
# absorbed q bmm with the nope/rope concat + fp8 cast so q is written
# directly in fp8 ("born fp8") and the standalone concat-cast kernel
# disappears. Not bit-exact vs the default path (same rounding stages,
# different GEMM accumulation order), hence default OFF until accuracy-
# gated (oracle + full-set gsm8k).
SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q = EnvBool(False)
# Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: pass a
# per-row valid-topk count (derived from the trailing -1 pad run of the
# topk indices) so the kernel skips whole pad-only topk blocks instead of
# computing masked zero contributions. Bit-exact by construction: skipped
# blocks contain only -1 pads, and -1 entries inside the consumed range
# still take the in-kernel clamp+mask path.
SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH = EnvBool(False)
# Opt-in: run the born-fp8 q-prep (absorbed bmm + concat + fp8 cast,
# ~173us/layer-call) on alt_stream underneath the DSA indexer — the two
# chains fork independently from the q_a_layernorm output. Requires
# SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q; eager-prefill-only via the born
# predicate. Coarse per-layer join keeps the single-slot born-q buffer
# WAR-safe.
SGLANG_ENABLE_DSA_Q8KV8_QPREP_OVERLAP = EnvBool(False)
# Opt-in: fuse the Q8KV8 non-prefix KV prep — cast-concat k/k_rope
# directly into the persistent fp8 kv buffer and zero the pad band in one
# Triton kernel (replaces bf16 _cat + copy_ cast + zero_ tail).
SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION = EnvBool(False)
# Q8KV8 born-fp8 q-prep codegen: "auto" = per-K Triton dispatch (default);
# "cuda" = the hand-written SM90 WGMMA kernel (bitwise identical to the
# Triton two_dot variant, 1.16-1.38x faster across GLM/DS shapes).
SGLANG_OPT_Q8KV8_QPREP_VARIANT = EnvStr("auto")
# sgl-kernel
SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK = EnvBool(False)
@@ -1114,6 +1163,14 @@ class Envs:
SGLANG_OPT_USE_JIT_EP_ACTIVATION = EnvBool(True)
SGLANG_OPT_FUSE_WQA_WKV = EnvBool(True)
SGLANG_OPT_SWIGLU_CLAMP_FUSION = EnvBool(True)
# DeepSeek/GLM MoE (deepseek_v2.py): quantize the (dp-gathered) MoE input
# to per-token-group-128 fp8 ONCE and feed both the fused shared-expert
# GEMM (cutlass w8a8 linear) and the routed experts' triton fused runner,
# instead of quantizing the same [T, hidden] tensor twice with different
# scale layouts. Only engages on CUDA with fp8 block-128 weights, the
# standard dispatcher, and the triton MoE runner; falls back silently
# otherwise.
SGLANG_OPT_MOE_QUANT_ONCE = EnvBool(False)
# Cache / overlap
SGLANG_OPT_USE_FUSED_STORE_CACHE = EnvBool(True)
@@ -19,6 +19,7 @@ from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__)
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
concat_cast_kv_fp8_pad,
dequantize_k_cache_paged,
gather_dequant_requant_fp8_paged,
)
@@ -30,6 +31,7 @@ from sglang.kernels.ops.attention.dsa.transform_index import (
from sglang.kernels.ops.attention.utils import (
concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8,
q8kv8_topk_length_from_indices,
seqlens_expand_triton,
)
from sglang.kernels.ops.kvcache.cache_ops import concat_and_cast_q_fp8_pad
@@ -495,6 +497,41 @@ class DeepseekSparseAttnBackend(
# Q8KV8 dispatch (no-ops for other backends).
self._q8kv8_identity_scale: Optional[torch.Tensor] = None
self._q8kv8_qpad_buf: Optional[torch.Tensor] = None
# Persistent (grow-only) fp8 KV destination for the Q8KV8 prefill
# gather: [capacity_rows, 576]. Avoids a fresh torch.zeros
# (alloc + full-buffer FillFunctor) per layer per call; only the
# `topk` -1-sentinel landing-pad rows need zeroing each call, and
# the gather kernel fuses that in. Same single-stream reuse
# argument as `_q8kv8_qpad_buf`.
self._q8kv8_kv_buf: Optional[torch.Tensor] = None
# Per-row valid-topk early-exit (SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH):
# rows whose topk indices end in a -1 pad run skip whole topk blocks
# in-kernel.
self._q8kv8_topk_length_enabled: bool = (
envs.SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH.get()
)
# Persistent (grow-only) kernel-output buffers (out/max_logits/lse).
self._q8kv8_out_bufs: Optional[tuple] = None
# Fused non-prefix KV prep (cast-concat k/k_rope directly into the
# fp8 buffer; SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION).
self._q8kv8_kv_cat_fusion: bool = (
envs.SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION.get()
)
# Born-fp8 q handshake (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): when the
# model's q-prep decides (via q8kv8_born_fp8_q_eligible) that this
# batch's forward_extend is guaranteed to hit
# _forward_flashmla_sparse_q8kv8, it writes the padded fp8 q directly
# (fused absorbed-bmm + concat + cast) into _q8kv8_born_q_buf and
# stashes (num_tokens, layer_id); the helper consumes the stash
# instead of rebuilding q_fp8. Same single-stream reuse argument as
# _q8kv8_qpad_buf. The bf16 q that flows through the attention API in
# that mode is a NaN-poisoned sentinel: any code path that reads it by
# mistake fails loudly instead of producing silently wrong output.
self._q8kv8_born_q_buf: Optional[torch.Tensor] = None
self._q8kv8_born_q_stash: Optional[Tuple[int, int]] = None
self._q8kv8_born_q_sentinel: Optional[torch.Tensor] = None
self._q8kv8_born_q_tbo = model_runner.server_args.enable_two_batch_overlap
from sglang.kernels.ops.attention.flash_mla_sm120 import (
_validate_flashinfer_sparse_mla_backend,
@@ -2083,6 +2120,24 @@ class DeepseekSparseAttnBackend(
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
layer_id=layer.layer_id,
)
if self._q8kv8_kv_cat_fusion:
# Fused path: no bf16 concat materialization — k and
# k_rope are cast-concatenated straight into the fp8
# buffer inside the helper.
return self._forward_flashmla_sparse_q8kv8(
q_nope=q_nope,
q_rope=q_rope,
kv_bf16=None,
kv_k=k,
kv_k_rope=k_rope,
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,
layer_id=layer.layer_id,
)
kv_cache = _cat([k, k_rope], dim=-1)
return self._forward_flashmla_sparse_q8kv8(
@@ -2094,6 +2149,7 @@ class DeepseekSparseAttnBackend(
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
layer_id=layer.layer_id,
)
# bf16 path (dsa_impl == "flashmla_sparse").
@@ -2428,6 +2484,97 @@ class DeepseekSparseAttnBackend(
return o
def q8kv8_born_fp8_q_eligible(
self, forward_batch: ForwardBatch, num_heads: int
) -> bool:
"""True iff this batch's forward_extend is guaranteed to consume q via
``_forward_flashmla_sparse_q8kv8`` (born-fp8 q handshake precondition).
Must stay in lockstep with the forward_extend dispatch: a True here
while dispatch takes any other branch would leak the NaN sentinel into
a real attention kernel (loud NaNs, not silent corruption, but still a
failed forward).
"""
if self.dsa_prefill_impl != "flashmla_sparse_q8":
return False
# RAGGED routing requires exactly EXTEND (excludes decode/idle, MIXED,
# target-verify and draft-extend, which use dsa_decode_impl anyway).
if forward_batch.forward_mode != ForwardMode.EXTEND:
return False
# Per-batch dense fallback (il <= threshold) reads bf16 q directly.
if self.use_mha:
return False
if self.hisparse_coordinator is not None:
return False
# TBO interleaves two micro-batches through one backend instance; the
# single-slot stash handshake is not safe there.
if self._q8kv8_born_q_tbo:
return False
if is_dsa_enable_prefill_cp():
return False
if (
self.get_topk_transform_method(forward_batch.forward_mode)
!= TopkTransformMethod.RAGGED
):
return False
# Mirror the helper's head-padding compatibility check.
if num_heads % 64 != 0 and 64 % num_heads != 0:
return False
return True
def q8kv8_acquire_born_q_buffer(
self, num_tokens: int, num_heads: int, head_dim: int, device: torch.device
) -> torch.Tensor:
"""Padded fp8 q destination for the born-fp8 kernel (grow-only).
Pad rows [num_heads:pad_heads] are zeroed at allocation and never
written afterwards (the fused kernel only writes the active heads),
matching the _q8kv8_qpad_buf invariant the SM90 kernel relies on.
"""
pad = 64
padded_heads = num_heads if num_heads % pad == 0 else pad
buf = self._q8kv8_born_q_buf
if (
buf is None
or buf.shape[0] < num_tokens
or buf.shape[1] != padded_heads
or buf.shape[2] != head_dim
):
buf = torch.zeros(
(num_tokens, padded_heads, head_dim),
dtype=torch.float8_e4m3fn,
device=device,
)
self._q8kv8_born_q_buf = buf
return buf[:num_tokens]
def q8kv8_stash_born_q(self, num_tokens: int, layer_id: int) -> None:
if self._q8kv8_born_q_stash is not None:
raise RuntimeError(
"q8kv8 born-fp8 q stash was never consumed (previous stash "
f"{self._q8kv8_born_q_stash}, new ({num_tokens}, {layer_id})): "
"the eligibility predicate fired but forward_extend dispatched "
"away from _forward_flashmla_sparse_q8kv8."
)
self._q8kv8_born_q_stash = (num_tokens, layer_id)
def q8kv8_born_q_sentinel(
self, num_tokens: int, num_heads: int, v_head_dim: int, device: torch.device
) -> torch.Tensor:
"""NaN-poisoned bf16 stand-in for q_nope_out in born-fp8 mode.
Only its shape/dtype/device are ever legitimately used downstream; a
NaN payload turns any accidental read into loud NaN output.
"""
numel = num_tokens * num_heads * v_head_dim
buf = self._q8kv8_born_q_sentinel
if buf is None or buf.numel() < numel:
buf = torch.full(
(numel,), float("nan"), dtype=torch.bfloat16, device=device
)
self._q8kv8_born_q_sentinel = buf
return buf[:numel].view(num_tokens, num_heads, v_head_dim)
def _forward_flashmla_sparse_q8kv8(
self,
q_nope: torch.Tensor,
@@ -2438,6 +2585,9 @@ class DeepseekSparseAttnBackend(
sm_scale: float,
paged_kv_cache: Optional[torch.Tensor] = None,
page_table_1_flattened: Optional[torch.Tensor] = None,
layer_id: Optional[int] = None,
kv_k: Optional[torch.Tensor] = None,
kv_k_rope: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Native FP8 (q8 x kv8) sparse-prefill attention (SM90 JIT kernel).
@@ -2472,12 +2622,37 @@ class DeepseekSparseAttnBackend(
required_padding = 64
need_padding = num_heads % required_padding != 0
# Born-fp8 fast path (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): the model's
# q-prep already wrote the padded fp8 q (fused absorbed-bmm + concat +
# cast); consume the stash instead of rebuilding it. q_nope here is
# the NaN sentinel (shape-only); q_rope's bf16 content is valid but
# unused.
born = self._q8kv8_born_q_stash
if born is not None:
self._q8kv8_born_q_stash = None
born_tokens, born_layer_id = born
if born_tokens != num_tokens or (
layer_id is not None and born_layer_id != layer_id
):
raise RuntimeError(
"q8kv8 born-fp8 q stash mismatch: stashed "
f"(num_tokens={born_tokens}, layer_id={born_layer_id}) but "
f"consuming (num_tokens={num_tokens}, layer_id={layer_id})."
)
q_fp8 = self._q8kv8_born_q_buf[:num_tokens]
expected_heads = required_padding if need_padding else num_heads
if q_fp8.shape[1] != expected_heads or q_fp8.shape[2] != head_dim:
raise RuntimeError(
"q8kv8 born-fp8 q buffer shape mismatch: got "
f"{tuple(q_fp8.shape)}, expected (*, {expected_heads}, "
f"{head_dim})."
)
# 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:
elif need_padding:
if required_padding % num_heads != 0:
raise ValueError(
f"num_heads={num_heads} cannot be padded to {required_padding}; "
@@ -2521,21 +2696,85 @@ class DeepseekSparseAttnBackend(
# 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.
#
# The destination is a persistent grow-only buffer instead of a fresh
# torch.zeros: rows [0, num_kv_tokens) are fully overwritten every
# call (gather kernel / cast-copy), so only the pad rows
# [num_kv_tokens, num_kv_tokens + topk) - exactly the rows the SM90
# kernel's -1 clamp (pad_base + slot) can read - need zeroing, and
# they need it EVERY call because a previous, larger call may have
# left real KV data there. The gather kernel fuses the pad-row
# zeroing; the bf16 path zeroes the tail explicitly.
topk = page_table_1.shape[-1]
if paged_kv_cache is not None:
num_kv_tokens = page_table_1_flattened.shape[0]
elif kv_k is not None:
num_kv_tokens = kv_k.shape[0]
else:
num_kv_tokens = kv_bf16.shape[0]
total_kv_rows = num_kv_tokens + topk
kv_buf = self._q8kv8_kv_buf
if kv_buf is None or kv_buf.shape[0] < total_kv_rows:
kv_buf = torch.empty(
(total_kv_rows, head_dim),
dtype=torch.float8_e4m3fn,
device=dev,
)
self._q8kv8_kv_buf = kv_buf
if paged_kv_cache is not None:
kv_padded = gather_dequant_requant_fp8_paged(
paged_kv_cache,
page_table_1_flattened,
extra_rows=topk,
out=kv_buf[:total_kv_rows],
).view(-1, 1, head_dim)
elif kv_k is not None:
# Fused non-prefix KV prep (SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION):
# cast-concat k/k_rope straight into the fp8 buffer + zero the pad
# band in ONE kernel — the bf16 _cat materialization, the copy_
# cast and the zero_ tail all disappear. Same store-cast as the
# gather kernel (bit-identical bytes).
kv_padded = concat_cast_kv_fp8_pad(
kv_buf[:total_kv_rows], kv_k, kv_k_rope, num_kv_tokens
).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_buf[:total_kv_rows]
# bf16 -> fp8 cast copy, same op as the previous fresh-buffer
# path (bit-identical bytes).
kv_padded[:num_kv_tokens].copy_(kv_bf16.view(num_kv_tokens, head_dim))
kv_padded[num_kv_tokens:].zero_()
kv_padded = kv_padded.view(-1, 1, head_dim)
# Per-row valid-topk count = last non-pad position + 1. Bit-exact
# vs topk_length=None: the skipped tail blocks contain only -1 pads
# (masked to zero contribution today), and -1 entries inside the
# consumed range still take the kernel's clamp+mask path. The
# backscan's cost is proportional to the trailing pad run, so rows
# with a full topk (all rows at long context) pay ~one block read.
topk_length = None
if self._q8kv8_topk_length_enabled:
topk_length = q8kv8_topk_length_from_indices(page_table_1)
# Persistent kernel-output buffers (out / max_logits / lse): the
# wrapper otherwise torch.empty's all three per layer-call. The
# kernel fully overwrites the active [:s_q] rows and everything runs
# on one stream, so reuse is safe — same argument as _q8kv8_qpad_buf.
s_q, pad_heads = q_fp8.shape[0], q_fp8.shape[1]
out_bufs = self._q8kv8_out_bufs
if (
out_bufs is None
or out_bufs[0].shape[0] < s_q
or out_bufs[0].shape[1] != pad_heads
):
out_bufs = (
torch.empty(
s_q, pad_heads, v_head_dim, dtype=torch.bfloat16, device=dev
),
torch.empty(s_q, pad_heads, dtype=torch.float32, device=dev),
torch.empty(s_q, pad_heads, dtype=torch.float32, device=dev),
)
self._q8kv8_out_bufs = out_bufs
o, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q_fp8,
kv=kv_padded,
@@ -2545,7 +2784,10 @@ class DeepseekSparseAttnBackend(
kv_scale=identity_scale,
d_v=v_head_dim,
attn_sink=None,
topk_length=None,
topk_length=topk_length,
out=out_bufs[0][:s_q],
max_logits=out_bufs[1][:s_q],
lse=out_bufs[2][:s_q],
)
# Trim the output back to the original head count if we padded.
+165 -1
View File
@@ -1,11 +1,14 @@
from __future__ import annotations
import functools
import logging
from contextlib import contextmanager
from enum import IntEnum, auto
from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
import triton
import triton.language as tl
from sglang.srt.distributed import (
GroupCoordinator,
@@ -136,6 +139,7 @@ class _DpGatheredBufferWrapper:
_local_dp_buffer_len: int = 0
_dp_max_padding: bool = False
_global_num_tokens: Optional[List[int]] = None
_global_num_tokens_gpu: Optional[torch.Tensor] = None
@classmethod
def set_metadata(cls, hidden_size: int, dtype: torch.dtype, device: torch.device):
@@ -153,11 +157,13 @@ class _DpGatheredBufferWrapper:
local_dp_buffer_len: int,
dp_max_padding: bool,
global_num_tokens: Optional[List[int]] = None,
global_num_tokens_gpu: Optional[torch.Tensor] = None,
):
cls._global_dp_buffer_len = global_dp_buffer_len
cls._local_dp_buffer_len = local_dp_buffer_len
cls._dp_max_padding = dp_max_padding
cls._global_num_tokens = global_num_tokens
cls._global_num_tokens_gpu = global_num_tokens_gpu
@classmethod
def get_global_dp_buffer(cls, group: GroupCoordinator) -> torch.Tensor:
@@ -201,6 +207,10 @@ class _DpGatheredBufferWrapper:
def get_dp_global_num_tokens(cls) -> List[int]:
return cls._global_num_tokens
@classmethod
def get_dp_global_num_tokens_gpu(cls) -> Optional[torch.Tensor]:
return cls._global_num_tokens_gpu
@classmethod
def get_dp_hidden_size(cls) -> int:
from sglang.srt.runtime_context import get_flags
@@ -229,9 +239,14 @@ def set_dp_buffer_len(
local_dp_buffer_len: int,
dp_max_padding: bool,
global_num_tokens: Optional[List[int]] = None,
global_num_tokens_gpu: Optional[torch.Tensor] = None,
):
_DpGatheredBufferWrapper.set_dp_buffer_len(
global_dp_buffer_len, local_dp_buffer_len, dp_max_padding, global_num_tokens
global_dp_buffer_len,
local_dp_buffer_len,
dp_max_padding,
global_num_tokens,
global_num_tokens_gpu,
)
@@ -505,6 +520,142 @@ def _dp_gather_via_all_gather(
# tp_size==dp_size (attn_tp_size==1) case is supported for now (e.g. tp8dp8).
_USE_DP_GATHERV = get_bool_env_var("SGLANG_DP_USE_GATHERV")
_DP_GATHER_FP8_GROUP = 128
# Grow-only gathered fp8 payload / scales buffers, keyed by device.
_dp_gather_fp8_bufs: dict = {}
@functools.lru_cache(maxsize=1)
def _use_dp_gather_fp8() -> bool:
from sglang.srt.environ import envs
return envs.SGLANG_ENABLE_DP_GATHER_FP8.get()
def _get_dp_gather_fp8_bufs(rows: int, hidden: int, device: torch.device):
key = str(device)
bufs = _dp_gather_fp8_bufs.get(key)
if bufs is None or bufs[0].shape[0] < rows:
bufs = (
torch.empty((rows, hidden), dtype=torch.uint8, device=device),
torch.empty(
(rows, hidden // _DP_GATHER_FP8_GROUP),
dtype=torch.float32,
device=device,
),
)
_dp_gather_fp8_bufs[key] = bufs
return bufs[0][:rows], bufs[1][:rows]
@triton.jit
def _dequant_per_token_group_fp8_kernel(
q_ptr,
s_ptr,
out_ptr,
HIDDEN: tl.constexpr,
NGROUPS: tl.constexpr,
GROUP: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
# HIDDEN may not be a multiple of BLOCK (e.g. DeepSeek 7168 vs BLOCK
# 2048): the tail iteration must be masked or it reads/writes up to
# BLOCK-1 elements past the row (cross-row corruption + OOB on the last
# row). HIDDEN is constexpr, so the mask folds away when it divides.
for start in tl.static_range(0, HIDDEN, BLOCK):
offs = start + tl.arange(0, BLOCK)
mask = offs < HIDDEN
qv = tl.load(q_ptr + row * HIDDEN + offs, mask=mask, other=0.0).to(tl.float32)
sv = tl.load(s_ptr + row * NGROUPS + offs // GROUP, mask=mask, other=0.0)
tl.store(out_ptr + row * HIDDEN + offs, (qv * sv).to(tl.bfloat16), mask=mask)
@triton.jit
def _mask_dp_pad_topk_ids_kernel(
topk_ids_ptr,
counts_ptr,
max_len,
TOPK: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
rank = row // max_len
pos = row % max_len
valid = pos < tl.load(counts_ptr + rank)
if valid == 0:
offs = tl.arange(0, BLOCK)
tl.store(topk_ids_ptr + row * TOPK + offs, -1, mask=offs < TOPK)
def mask_dp_pad_moe_topk_ids(topk_ids: torch.Tensor) -> None:
"""Set MAX_LEN pad rows' (post-translation, local) topk_ids to -1 in place.
Under dp-attention MAX_LEN padding the gathered MoE buffer is
[dp_size * max_len, hidden] with rank r's real rows at
[r*max_len, r*max_len + global_num_tokens[r]); the pad rows carry stale
hidden values, run the router, and get dispatched into experts whose
outputs are then discarded by the post-reorder scatter — pure wasted
compute, and a masked-grouped-GEMM workspace blow-up when they collide
on the same top-k. -1 is the drop sentinel both the triton fused_moe
(filter_expert) and the DeepGEMM EP preprocess honor; it must be applied
AFTER the local_expert_mapping gather (a pre-translation -1 aliases to
the mapping table's last entry). Capture-safe: per-batch state is read
only from the replay-updated global_num_tokens_gpu tensor.
"""
counts = _DpGatheredBufferWrapper.get_dp_global_num_tokens_gpu()
if counts is None:
return
max_len = _DpGatheredBufferWrapper.get_local_dp_buffer_len()
rows, topk = topk_ids.shape
if max_len <= 0 or rows != counts.shape[0] * max_len:
# Layout mismatch (e.g. non-DP or logits-path caller): do nothing.
return
_mask_dp_pad_topk_ids_kernel[(rows,)](
topk_ids,
counts,
max_len,
TOPK=topk,
BLOCK=triton.next_power_of_2(topk),
)
def _dp_gather_via_all_gatherv_fp8(
global_tokens: torch.Tensor,
local_real: torch.Tensor,
sizes: List[int],
):
"""fp8 wire format for the variable-length DP gather: quantize the local
rows per-token-group (the SAME group-128 quantization the MoE expert GEMMs
apply to their input downstream), gather payload (as uint8 — NCCL has no
fp8 dtype; the gatherv leg is broadcast-only so a byte view is safe) and
scales in two output-buffered gatherv calls, then dequantize into the
bf16 global buffer. Zero pad rows quantize to (q=0, s=eps) and so
dequantize back to exact zeros — the MoE-tail invariant is preserved.
The combine leg (reduce_scatterv) stays bf16: NCCL SUM cannot run on fp8."""
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
rows = global_tokens.shape[0]
hidden = global_tokens.shape[-1]
q, s = sglang_per_token_group_quant_fp8(
local_real.contiguous(), _DP_GATHER_FP8_GROUP
)
gq, gs = _get_dp_gather_fp8_bufs(rows, hidden, global_tokens.device)
tp_group = get_tp_group()
tp_group.all_gatherv(q.view(torch.uint8), sizes=sizes, output=gq)
tp_group.all_gatherv(s, sizes=sizes, output=gs)
_dequant_per_token_group_fp8_kernel[(rows,)](
gq.view(torch.float8_e4m3fn),
gs,
global_tokens,
HIDDEN=hidden,
NGROUPS=hidden // _DP_GATHER_FP8_GROUP,
GROUP=_DP_GATHER_FP8_GROUP,
BLOCK=2048,
)
def is_dp_gatherv_active() -> bool:
"""Variable-length DP-MoE gather/scatter (all_gatherv + reduce_scatterv) is
@@ -568,6 +719,19 @@ def _dp_gather_via_all_gatherv(
# falls back to all_reduce). Pass global_tokens as the NCCL output buffer so
# the gather writes directly into it -- avoids the previous extra full-buffer
# torch.cat + copy_ (two ~sum(sizes)*hidden DtoD copies, ~700us/layer at c512).
# NOTE: the fp8 branch condition must be identical on EVERY DP rank (all
# ranks must issue the same NCCL op sequence) — env/dtype/hidden are
# rank-uniform; never gate on per-rank state like forward_mode (ranks can
# be extend/idle-mixed within one global forward). Prefill-only is
# already structural: the gatherv path runs only under SUM_LEN padding,
# which decode-only steps and CUDA-graph capture never select.
if (
_use_dp_gather_fp8()
and global_tokens.dtype == torch.bfloat16
and global_tokens.shape[-1] % _DP_GATHER_FP8_GROUP == 0
):
_dp_gather_via_all_gatherv_fp8(global_tokens, local_real, sizes)
return
get_tp_group().all_gatherv(local_real, sizes=sizes, output=global_tokens)
@@ -1332,7 +1332,12 @@ class FusedMoE(torch.nn.Module):
f"Unsupported weight_name {weight_name} for FusedMoE weight_loader_fused. Nothing is loaded."
)
def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
def forward(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
):
if self._use_ascend_fuseep:
from sglang.srt.hardware_backend.npu.moe.fuseep import forward_fuseep
@@ -1360,11 +1365,20 @@ class FusedMoE(torch.nn.Module):
)
else:
# Make sure there is torch lib op registration for the whole moe layer
return self.forward_impl(hidden_states, topk_output)
return self.forward_impl(
hidden_states, topk_output, pre_quant_input=pre_quant_input
)
else:
return self.forward_impl(hidden_states, topk_output)
return self.forward_impl(
hidden_states, topk_output, pre_quant_input=pre_quant_input
)
def forward_impl(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
def forward_impl(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
):
origin_hidden_states_dim = hidden_states.shape[-1]
assert self.quant_method is not None
@@ -1375,6 +1389,18 @@ class FusedMoE(torch.nn.Module):
dispatch_output = self.dispatcher.dispatch(
hidden_states=hidden_states, topk_output=topk_output
)
if (
pre_quant_input is not None
and dispatch_output.format.is_standard()
and dispatch_output.hidden_states_scale is None
):
# SGLANG_OPT_MOE_QUANT_ONCE: the standard dispatch was a pure
# passthrough, so the caller's pre-quantized (q, scale) pair still
# matches dispatch_output.hidden_states; attach it for the triton
# fused runner to skip its own activation quant.
dispatch_output = dispatch_output._replace(
hidden_states_pre_quant=pre_quant_input
)
combine_input = self.run_moe_core(
dispatch_output=dispatch_output,
@@ -1,5 +1,6 @@
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
@@ -8,6 +9,9 @@ import torch
from sglang.kernels.ops.attention.dsv4 import silu_and_mul_masked_post_quant
from sglang.kernels.ops.quantization import per_token_group_quant
logger = logging.getLogger(__name__)
from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
@@ -448,9 +452,20 @@ class DeepGemmRunnerCore(MoeRunnerCore):
num_groups, m, k = hidden_states.shape
n = w13_weight.size(1)
gateup_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
try:
gateup_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
except torch.OutOfMemoryError:
logger.error(
"Masked grouped-GEMM workspace allocation failed "
"(num_groups=%d m=%d n=%d). If this happens under saturated "
"dp-attention prefill, try SGLANG_OPT_DG_MASKED_M_CAP=1.",
num_groups,
m,
n,
)
raise
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked(
(hidden_states, hidden_states_scale),
(w13_weight, w13_scale),
@@ -217,6 +217,15 @@ def fused_experts_none_to_triton(
fused_experts,
)
# SGLANG_OPT_MOE_QUANT_ONCE: use the caller's pre-quantized activation
# (per-token-group-128 fp8 q + scales) instead of re-quantizing inside
# invoke_fused_moe_kernel.
pre_quant = dispatch_output.hidden_states_pre_quant
if pre_quant is not None:
a1_q, a1_scale = pre_quant
else:
a1_q, a1_scale = None, quant_info.a13_scale
output = fused_experts(
hidden_states=dispatch_output.hidden_states,
w1=quant_info.w13_weight,
@@ -234,9 +243,10 @@ def fused_experts_none_to_triton(
w2_scale=quant_info.w2_scale,
w1_zp=quant_info.w13_zp,
w2_zp=quant_info.w2_zp,
a1_scale=quant_info.a13_scale,
a1_scale=a1_scale,
a2_scale=quant_info.a2_scale,
block_shape=quant_info.block_shape,
a1_q=a1_q,
)
return StandardCombineInput(
@@ -128,6 +128,7 @@ def inplace_fused_experts(
filter_expert: bool = True,
swiglu_limit: Optional[float] = None,
gate_up_interleaved: bool = True,
a1_q: Optional[torch.Tensor] = None,
) -> None:
fused_experts_impl(
hidden_states,
@@ -160,6 +161,7 @@ def inplace_fused_experts(
filter_expert,
swiglu_limit=swiglu_limit,
gate_up_interleaved=gate_up_interleaved,
a1_q=a1_q,
)
@@ -194,6 +196,7 @@ def outplace_fused_experts(
filter_expert: bool = True,
swiglu_limit: Optional[float] = None,
gate_up_interleaved: bool = True,
a1_q: Optional[torch.Tensor] = None,
) -> torch.Tensor:
return fused_experts_impl(
hidden_states,
@@ -226,6 +229,7 @@ def outplace_fused_experts(
filter_expert=filter_expert,
swiglu_limit=swiglu_limit,
gate_up_interleaved=gate_up_interleaved,
a1_q=a1_q,
)
@@ -249,6 +253,7 @@ def fused_experts(
a1_scale: Optional[torch.Tensor] = None,
a2_scale: Optional[torch.Tensor] = None,
block_shape: Optional[List[int]] = None,
a1_q: Optional[torch.Tensor] = None,
):
topk_weights, topk_ids, _ = topk_output
filter_expert = (
@@ -286,6 +291,7 @@ def fused_experts(
filter_expert,
swiglu_limit=moe_runner_config.swiglu_limit,
gate_up_interleaved=moe_runner_config.gate_up_interleaved,
a1_q=a1_q,
)
return hidden_states
else:
@@ -319,6 +325,7 @@ def fused_experts(
filter_expert=filter_expert,
swiglu_limit=moe_runner_config.swiglu_limit,
gate_up_interleaved=moe_runner_config.gate_up_interleaved,
a1_q=a1_q,
)
@@ -461,12 +468,19 @@ def _fused_moe_kernel_sequence(
hooks: Optional[Any] = None,
swiglu_limit: Optional[float] = None,
gate_up_interleaved: bool = True,
a1_q: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Run the MoE kernel/activation/kernel/combine sequence in a single shot.
Inputs are already aligned and the block-size config is already resolved.
Supports optional LoRA hooks that fire between the two kernels and before
combine. Returns ``out_hidden_states``.
``a1_q`` (SGLANG_OPT_MOE_QUANT_ONCE): optional pre-quantized fp8 view of
``hidden_states`` for the gate-up GEMM (per-token-group ``block_shape[1]``
quant, ``a1_scale`` holds the matching scales, rows may exceed
``num_tokens`` due to 4-row padding). ``hidden_states`` stays bf16 and is
still used for output dtype/shape and the inplace combine.
"""
num_tokens = hidden_states.shape[0]
E, N, _ = w1.shape
@@ -479,6 +493,17 @@ def _fused_moe_kernel_sequence(
if hooks and (hooks.after_gate_up is not None or hooks.after_down is not None):
down_moe_use_tma = False
if a1_q is not None:
assert (
use_fp8_w8a8
and block_shape is not None
and a1_scale is not None
and a1_q.dtype == torch.float8_e4m3fn
and a1_q.is_contiguous()
and a1_q.shape[0] >= num_tokens
and a1_q.shape[1] == hidden_states.shape[1]
), "a1_q requires block-wise fp8 with matching pre-quantized activation"
padded_tokens = (
min(num_tokens * topk, E + 1) * (config["BLOCK_SIZE_M"] - 1)
if down_moe_use_tma
@@ -520,7 +545,7 @@ def _fused_moe_kernel_sequence(
)
invoke_fused_moe_kernel(
hidden_states,
a1_q if a1_q is not None else hidden_states,
w1,
b1,
intermediate_cache1,
@@ -866,6 +891,7 @@ def fused_experts_impl(
filter_expert: bool = True,
swiglu_limit: Optional[float] = None,
gate_up_interleaved: bool = True,
a1_q: Optional[torch.Tensor] = None,
):
padded_size = padding_size
if not (use_fp8_w8a8 or use_int8_w8a8) or block_shape is not None or _use_aiter:
@@ -942,6 +968,7 @@ def fused_experts_impl(
hooks=None,
swiglu_limit=swiglu_limit,
gate_up_interleaved=gate_up_interleaved,
a1_q=a1_q,
)
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, NamedTuple, Optional
from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple
import torch
@@ -14,6 +14,8 @@ from sglang.srt.layers.dp_attention import (
get_dp_global_num_tokens,
get_local_dp_buffer,
is_allocation_symmetric,
is_dp_max_padding,
mask_dp_pad_moe_topk_ids,
)
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.token_dispatcher.base import (
@@ -39,6 +41,10 @@ from sglang.srt.utils.common import (
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
from sglang.srt.environ import envs as _envs
_MASK_DP_PAD_MOE = _envs.SGLANG_OPT_MASK_DP_PAD_MOE.get()
if TYPE_CHECKING:
from sglang.srt.layers.moe.topk import TopKOutput
@@ -62,6 +68,11 @@ class StandardDispatchOutput(NamedTuple):
hidden_states: torch.Tensor
hidden_states_scale: Optional[torch.Tensor]
topk_output: TopKOutput
# SGLANG_OPT_MOE_QUANT_ONCE: optional pre-quantized (q, scale) pair for
# ``hidden_states`` (per-token-group-128 fp8, q rows possibly padded to a
# multiple of 4). Consumed by the standard->triton fused runner so it can
# skip its own activation quant; ``hidden_states`` itself stays bf16.
hidden_states_pre_quant: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
@property
def format(self) -> DispatchOutputFormat:
@@ -213,9 +224,18 @@ class StandardDispatcher(BaseDispatcher):
)
elif not self.use_aiter_moe_runner:
if TopKOutputChecker.format_is_standard(topk_output):
topk_output = topk_output._replace(
topk_ids=self.local_expert_mapping[topk_output.topk_ids]
)
topk_ids_local = self.local_expert_mapping[topk_output.topk_ids]
# Drop dp-attention MAX_LEN pad rows from the dispatch:
# pad rows carry stale hidden through the router and
# their expert outputs are discarded downstream — pure
# wasted compute (and a masked-grouped-GEMM workspace
# blow-up when they collide on the same top-k). Must
# run POST-translation (a pre-translation -1 aliases to
# the mapping table's last entry); -1 is the drop
# sentinel both the triton and deep_gemm runners honor.
if _MASK_DP_PAD_MOE and is_dp_max_padding():
mask_dp_pad_moe_topk_ids(topk_ids_local)
topk_output = topk_output._replace(topk_ids=topk_ids_local)
elif TopKOutputChecker.format_is_triton_kernels(topk_output):
raise NotImplementedError()
@@ -789,11 +789,28 @@ def cutlass_w8a8_block_fp8_linear_with_fallback(
input_scale: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert input_scale is None
# TODO: add more robust shape check here
shape_supported = weight.shape[0] % 128 == 0 and weight.shape[1] % 128 == 0
if input_scale is not None:
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): ``input`` is
# the fp8 per-token-group-128 q (rows possibly padded to a multiple
# of 4), ``input_scale`` the matching column-major scales
# (stride(0) == 1). Output keeps the (padded) row count; the caller
# slices back to the true token count.
assert shape_supported, (
"pre-quantized fp8 input requires cutlass-supported weight shapes "
f"(got {tuple(weight.shape)})"
)
assert input.dtype == torch.float8_e4m3fn
input_2d = input.view(-1, input.shape[-1])
output = fp8_blockwise_scaled_mm(
input_2d, weight.T, input_scale, weight_scale.T, out_dtype=torch.bfloat16
)
if bias is not None:
output += bias
return output.view(*input.shape[:-1], weight.shape[0])
if not shape_supported:
# fallback to triton
return triton_w8a8_block_fp8_linear(
@@ -829,7 +846,33 @@ def deepgemm_w8a8_block_fp8_linear_with_fallback(
input_scale: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert input_scale is None
if input_scale is not None:
# Pre-quantized activation (SGLANG_OPT_MOE_QUANT_ONCE): ``input`` is
# the fp8 per-token-group-128 q with rows padded to a multiple of 4
# and ``input_scale`` the matching column-major fp32 scales
# (stride == (1, padded_rows)) -- identical to the MN-major
# TMA-aligned layout this path's own quant would produce below.
# Output keeps the padded row count; the caller slices back.
# UE8M0 packed scales (Blackwell DeepGEMM) use a different layout;
# the caller gates on it.
assert not deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
assert input.dtype == torch.float8_e4m3fn
assert weight.shape[0] % 64 == 0 and weight.shape[1] % 128 == 0, (
"pre-quantized fp8 input requires DeepGEMM-supported weight shapes "
f"(got {tuple(weight.shape)})"
)
input_2d = input.view(-1, input.shape[-1])
output = w8a8_block_fp8_matmul_deepgemm(
input_2d,
weight,
input_scale,
weight_scale,
block_size,
output_dtype=torch.bfloat16,
)
if bias is not None:
output += bias
return output.view(*input.shape[:-1], weight.shape[0])
output_dtype = input.dtype
dtype_supported = output_dtype == torch.bfloat16
@@ -1282,6 +1282,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
num_tokens,
dp_padding_mode.is_max_len(),
global_num_tokens,
self.global_num_tokens_gpu,
)
set_is_extend_in_batch(self.is_extend_in_batch)
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Optional
import torch
from sglang.kernels.ops.kvcache.cache_ops import absorbed_bmm_concat_cast_q_fp8
from sglang.kernels.ops.quantization.fp8_kernel import (
fp8_dtype,
per_tensor_quant_mla_fp8,
@@ -52,6 +53,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
is_in_breakable_cuda_graph,
)
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
get_tc_piecewise_forward_context,
is_in_tc_piecewise_cuda_graph,
)
from sglang.srt.models.deepseek_common.utils import (
@@ -75,6 +77,8 @@ from sglang.srt.utils.custom_op import register_custom_op
logger = logging.getLogger(__name__)
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
_ENABLE_DSA_Q8KV8_BORN_FP8_Q = envs.SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q.get()
_ENABLE_DSA_Q8KV8_QPREP_OVERLAP = envs.SGLANG_ENABLE_DSA_Q8KV8_QPREP_OVERLAP.get()
if TYPE_CHECKING:
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
@@ -254,6 +258,84 @@ class DeepseekMLAForwardMixin:
attn_output_buf=attn_output_buf,
)
def _q8kv8_born_fp8_q_backend(
self: DeepseekV2AttentionMLA,
forward_batch: ForwardBatch,
llama_4_scaling: Optional[torch.Tensor],
):
"""Return the DSA backend iff the born-fp8 q fast path can run.
Gated by SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q (checked by the caller).
When this returns a backend, the bf16 absorbed bmm + the standalone
concat_and_cast_q_fp8_pad are replaced by one fused kernel that writes
the fp8 q directly into the backend's q8kv8 buffer; q_nope_out becomes
a NaN sentinel. Every condition here must therefore guarantee that
forward_extend consumes q via _forward_flashmla_sparse_q8kv8 and that
nothing else reads q_nope_out's payload.
"""
from sglang.srt.model_executor.runner import get_is_capture_mode
if llama_4_scaling is not None:
return None
if _is_hip or _is_cpu:
return None
if self.current_attention_backend not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
return None
if self.use_deep_gemm_bmm:
return None
w_kc = self.w_kc
if w_kc is None or w_kc.dtype != torch.bfloat16:
return None
if is_kv_b_lora_active(self) or _SGLANG_EXPERIMENTAL_LORA_OPTI:
return None
# The fused kernel consumes the post-rope q_pe, so the eager rope
# apply below must run (mirror of its condition).
if self.rotary_emb is None:
return None
if self._fuse_rope_for_trtllm_mla(forward_batch):
return None
if self._skip_rope_for_dsa_tilelang_fused():
return None
if self._skip_rope_for_aiter_fused_mla():
return None
if _use_aiter and _is_gfx95_supported and not self.use_dsa:
return None
# Graph/compile surfaces run their own dispatch; the python-side
# stash handshake is eager-only.
if is_graph_dsa_split_op_surface(forward_batch):
return None
if get_tc_piecewise_forward_context() is not None:
return None
if is_in_breakable_cuda_graph():
return None
if get_is_capture_mode():
return None
if get_parallel().dcp_enabled:
return None
# Context-parallel prefill reshuffles the KV side; keep the handshake
# out of those paths.
if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch):
return None
# Kernel shape constraints (tl.arange / tl.dot / block tiling). K
# (qk_nope_head_dim) needs only K % 16 == 0 and K <= 256: power-of-2
# K (DeepSeek 128) takes the kernel's preload-once path, other K
# (GLM-5 192) its split-K loop.
k_dim = self.qk_nope_head_dim
rope_dim = self.qk_rope_head_dim
if k_dim < 16 or k_dim > 256 or k_dim % 16 != 0:
return None
if rope_dim <= 0 or (rope_dim & (rope_dim - 1)) != 0:
return None
if self.kv_lora_rank % 128 != 0:
return None
if tuple(w_kc.shape) != (self.num_local_heads, k_dim, self.kv_lora_rank):
return None
backend = get_attn_backend()
eligible = getattr(backend, "q8kv8_born_fp8_q_eligible", None)
if eligible is None or not eligible(forward_batch, self.num_local_heads):
return None
return backend
def forward_absorb_prepare(
self: DeepseekV2AttentionMLA,
positions: torch.Tensor,
@@ -265,6 +347,11 @@ class DeepseekMLAForwardMixin:
):
from sglang.srt.model_executor.runner import get_is_capture_mode
# Q8KV8 q-prep/indexer overlap handshake (see the fork site below):
# True between the alt-stream fork and its consumption in the born
# block; also suppresses the duplicate split/rope on that path.
self._q8kv8_qprep_overlap_pending = False
fuse_bmm_attention = (
self.q_lora_rank is not None
and self._can_fuse_bmm_into_attention(forward_batch)
@@ -422,6 +509,47 @@ class DeepseekMLAForwardMixin:
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
fusion_plan = self._make_mla_bmm_fusion_plan(q, q_nope)
# Q8KV8 q-prep/indexer overlap (opt-in): the born-fp8 q-prep
# chain (split -> rope -> fused absorbed-bmm+cast, ~173us)
# and the indexer chain both fork from the q_a_layernorm
# output and never touch each other's tensors, so the q-prep
# can run on alt_stream underneath the indexer. The fork
# must be enqueued BEFORE the indexer (a later wait_stream
# would serialize behind it). The born predicate itself
# guarantees eager-only and the plain-rope branch (all fused
# /skip-rope variants make it return None), so applying rope
# here is exactly what the skipped block below would do.
if (
_ENABLE_DSA_Q8KV8_QPREP_OVERLAP
and _ENABLE_DSA_Q8KV8_BORN_FP8_Q
and fusion_plan is None
and self.alt_stream is not None
and q_lora is not None
and self.rotary_emb is not None
):
_born_backend_early = self._q8kv8_born_fp8_q_backend(
forward_batch, llama_4_scaling
)
if _born_backend_early is not None:
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
_q_fp8 = _born_backend_early.q8kv8_acquire_born_q_buffer(
q_nope.shape[0],
self.num_local_heads,
self.kv_lora_rank + self.qk_rope_head_dim,
q_nope.device,
)
self.alt_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(self.alt_stream):
absorbed_bmm_concat_cast_q_fp8(
_q_fp8,
q_nope,
self.w_kc,
q_pe,
self.num_local_heads,
)
self._q8kv8_qprep_overlap_pending = True
if q_lora is not None:
if self.should_run_indexer(prev_topk_indices):
topk_indices = self.indexer(
@@ -456,6 +584,15 @@ class DeepseekMLAForwardMixin:
q_nope, q_pe, k_pe = self._split_q_nope_pe(q, latent_cache)
_kvb_q = None
born_q_backend = None
if (
_ENABLE_DSA_Q8KV8_BORN_FP8_Q
and fusion_plan is None
and q_nope.dtype == torch.bfloat16
):
born_q_backend = self._q8kv8_born_fp8_q_backend(
forward_batch, llama_4_scaling
)
if q_replicate_active:
# full-head absorb with the pre-gathered w_kc (q_nope already full-head)
q_nope_out = (
@@ -467,6 +604,11 @@ class DeepseekMLAForwardMixin:
# The composite split op fills q_nope_out_buf and attention reads
# this transposed alias directly.
q_nope_out = fusion_plan.q_nope_out_view
elif born_q_backend is not None:
# Born-fp8 q: skip the bf16 absorbed bmm entirely; the fused
# bmm+concat+cast kernel (launched after rope below) writes the
# fp8 q directly into the q8kv8 backend buffer.
q_nope_out = None
else:
if _SGLANG_EXPERIMENTAL_LORA_OPTI:
# Fork the kv_b q-correction A-step onto the LoRA side stream to overlap the bmm.
@@ -591,9 +733,41 @@ class DeepseekMLAForwardMixin:
or self.use_dsa
or self.current_attention_backend == "triton"
)
# Already applied at the q-prep/indexer overlap fork.
and not self._q8kv8_qprep_overlap_pending
):
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
if born_q_backend is not None:
# Born-fp8 q (SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q): one fused
# kernel replaces bmm -> bf16 q_nope_out ->
# concat_and_cast_q_fp8_pad. q_nope is the pre-absorb bf16 view
# (rope only touched the disjoint q_pe columns) and q_pe carries
# the post-rope values. The stash is consumed by
# _forward_flashmla_sparse_q8kv8; q_nope_out becomes a
# NaN-poisoned shape-only sentinel.
num_tokens = q_nope.shape[0]
if self._q8kv8_qprep_overlap_pending:
# q_fp8 was produced on alt_stream at the fork above; join so
# everything downstream (incl. the next layer's fork, which
# reuses the single born-q slot) orders after it.
torch.cuda.current_stream().wait_stream(self.alt_stream)
self._q8kv8_qprep_overlap_pending = False
else:
q_fp8 = born_q_backend.q8kv8_acquire_born_q_buffer(
num_tokens,
self.num_local_heads,
self.kv_lora_rank + self.qk_rope_head_dim,
q_nope.device,
)
absorbed_bmm_concat_cast_q_fp8(
q_fp8, q_nope, self.w_kc, q_pe, self.num_local_heads
)
born_q_backend.q8kv8_stash_born_q(num_tokens, self.attn_mqa.layer_id)
q_nope_out = born_q_backend.q8kv8_born_q_sentinel(
num_tokens, self.num_local_heads, self.kv_lora_rank, q_nope.device
)
dsa_prefill_cp = dsa_use_prefill_cp(forward_batch)
mla_prefill_cp = mla_use_prefill_cp(forward_batch)
defer_kv_gather_until_after_rope = _should_defer_dsa_cp_kv_gather(
+180 -18
View File
@@ -241,6 +241,9 @@ from sglang.kernels.ops.gemm.fused_a_gemm import (
logger = logging.getLogger(__name__)
# One-time SGLANG_OPT_MOE_QUANT_ONCE engagement log (see _moe_quant_once_enabled).
_moe_quant_once_logged = False
_enable_pcg_dsv2_dual_stream = (
_is_cuda and envs.SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM.get()
)
@@ -307,6 +310,7 @@ class DeepseekV2MLP(nn.Module):
x,
forward_batch=None,
gemm_output_zero_allocator: BumpAllocator = None,
gateup_pre_quant: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
):
if (self.tp_size == 1) and x.shape[0] == 0:
return x
@@ -336,17 +340,24 @@ class DeepseekV2MLP(nn.Module):
out, _ = self.down_proj((out_fp4, out_scale))
return out
if (
gemm_output_zero_allocator is not None
and x.shape[0] <= 256
and self.gate_up_proj.weight.dtype == torch.uint8
):
y = gemm_output_zero_allocator.allocate(
x.shape[0] * self.gate_up_proj.output_size_per_partition
).view(x.shape[0], self.gate_up_proj.output_size_per_partition)
x = (x, None, y)
if gateup_pre_quant is not None:
# SGLANG_OPT_MOE_QUANT_ONCE: reuse the caller's per-token-group-128
# fp8 (q, scale) of x for the gate_up GEMM instead of re-quantizing
# inside the fp8 linear method. q rows may be padded to a multiple
# of 4; the caller slices the MLP output back.
gate_up, _ = self.gate_up_proj(gateup_pre_quant)
else:
if (
gemm_output_zero_allocator is not None
and x.shape[0] <= 256
and self.gate_up_proj.weight.dtype == torch.uint8
):
y = gemm_output_zero_allocator.allocate(
x.shape[0] * self.gate_up_proj.output_size_per_partition
).view(x.shape[0], self.gate_up_proj.output_size_per_partition)
x = (x, None, y)
gate_up, _ = self.gate_up_proj(x)
gate_up, _ = self.gate_up_proj(x)
# Fast path: fused silu+clamp+fp8_quant+deepgemm when conditions met.
# Only valid when down_proj does NOT need an all-reduce and its weights
# are fp8 (uint8 storage with weight_scale_inv).
@@ -822,6 +833,9 @@ class DeepseekV2MoE(nn.Module):
or get_moe_a2a_backend().is_flashinfer()
)
self._fuse_shared_experts_inside_sbo = SboFlags.fuse_shared_experts_inside_sbo()
# SGLANG_OPT_MOE_QUANT_ONCE eligibility, resolved lazily on first
# forward (weights and runner are final by then). None = undecided.
self._moe_quant_once: Optional[bool] = None
def get_moe_weights(self):
# EPLB only rebalances physical routed experts. Fused shared expert
@@ -933,6 +947,13 @@ class DeepseekV2MoE(nn.Module):
# deep_gemm does not free hidden_states, which the shared expert reads on the alt stream.
use_flashinfer_trtllm_bypass = get_forward().flashinfer_trtllm_bypass
current_stream = torch.cuda.current_stream()
# Quantize-once (SGLANG_OPT_MOE_QUANT_ONCE) must happen on the main
# stream BEFORE the alt-stream fork so both consumers see it.
pre_quant_input = (
None
if use_flashinfer_trtllm_bypass
else self._maybe_quant_moe_input_once(hidden_states)
)
self.alt_stream.wait_stream(current_stream)
has_shared_output = (
hidden_states.shape[0] > 0 and self.num_fused_shared_experts == 0
@@ -975,6 +996,10 @@ class DeepseekV2MoE(nn.Module):
)
elif use_flashinfer_trtllm_bypass:
final_hidden_states = self.experts.forward_impl(hidden_states, topk_output)
elif pre_quant_input is not None:
final_hidden_states = self.experts(
hidden_states, topk_output, pre_quant_input=pre_quant_input
)
else:
final_hidden_states = self.experts(hidden_states, topk_output)
if (
@@ -988,7 +1013,9 @@ class DeepseekV2MoE(nn.Module):
# Shared expert on alt stream, issued AFTER the main (routed) branch. See note above.
with torch.cuda.stream(self.alt_stream):
shared_output = self._forward_shared_experts(
hidden_states, gemm_output_zero_allocator
hidden_states,
gemm_output_zero_allocator,
pre_quant_input=pre_quant_input,
)
current_stream.wait_stream(self.alt_stream)
@@ -1044,13 +1071,22 @@ class DeepseekV2MoE(nn.Module):
# reduce_scatterv. When set, never compute/add it here (on the global buffer).
shared_output = None
if hidden_states.shape[0] > 0:
# Quantize-once (SGLANG_OPT_MOE_QUANT_ONCE): only worthwhile when
# the shared expert also runs here on the same tensor.
pre_quant_input = (
None
if skip_shared_experts
else self._maybe_quant_moe_input_once(hidden_states)
)
if (
not defer_shared
and not self._fuse_shared_experts_inside_sbo
and not skip_shared_experts
):
shared_output = self._forward_shared_experts(
hidden_states, gemm_output_zero_allocator
hidden_states,
gemm_output_zero_allocator,
pre_quant_input=pre_quant_input,
)
# router_logits: (num_tokens, n_experts)
router_logits = self.gate(hidden_states, gemm_output_zero_allocator)
@@ -1066,6 +1102,7 @@ class DeepseekV2MoE(nn.Module):
**topk_kwargs,
)
else:
pre_quant_input = None
shared_output = None
topk_output = self.topk.empty_topk_output(
hidden_states.device, layer_id=self.layer_id
@@ -1101,10 +1138,17 @@ class DeepseekV2MoE(nn.Module):
self.experts.dispatcher.register_post_combine_hook(_post_combine_hook)
)
final_hidden_states = self.experts(
hidden_states,
topk_output,
)
if pre_quant_input is not None:
final_hidden_states = self.experts(
hidden_states,
topk_output,
pre_quant_input=pre_quant_input,
)
else:
final_hidden_states = self.experts(
hidden_states,
topk_output,
)
if (
not _is_cuda
and not _is_musa
@@ -1122,7 +1166,9 @@ class DeepseekV2MoE(nn.Module):
and not skip_shared_experts
):
shared_output = self._forward_shared_experts(
hidden_states, gemm_output_zero_allocator
hidden_states,
gemm_output_zero_allocator,
pre_quant_input=pre_quant_input,
)
final_hidden_states = maybe_fuse_routed_scale_and_shared_add(
@@ -1426,15 +1472,131 @@ class DeepseekV2MoE(nn.Module):
return final_hidden_states
def _forward_shared_experts(
self, hidden_states, gemm_output_zero_allocator: BumpAllocator = None
self,
hidden_states,
gemm_output_zero_allocator: BumpAllocator = None,
pre_quant_input: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
):
if (hidden_states.shape[0] > 0) and (self.num_fused_shared_experts == 0):
if pre_quant_input is not None:
# SGLANG_OPT_MOE_QUANT_ONCE: (q, s) rows may be padded to a
# multiple of 4; the padded rows flow through the MLP (all ops
# are row-local) and are sliced off here.
out = self.shared_experts(
hidden_states, gateup_pre_quant=pre_quant_input
)
return out[: hidden_states.shape[0]]
return self.shared_experts(
hidden_states, gemm_output_zero_allocator=gemm_output_zero_allocator
)
else:
return None
def _moe_quant_once_enabled(self) -> bool:
"""SGLANG_OPT_MOE_QUANT_ONCE: quantize the (dp-gathered) MoE input to
per-token-group-128 fp8 once per layer and feed both the fused shared
expert's fp8 GEMM (cutlass or deepgemm w8a8 linear) and the routed
experts' triton fused runner, instead of quantizing the same
[T, hidden] tensor twice with different scale layouts."""
if self._moe_quant_once is None:
self._moe_quant_once, reason = self._compute_moe_quant_once_enabled()
global _moe_quant_once_logged
if envs.SGLANG_OPT_MOE_QUANT_ONCE.get() and not _moe_quant_once_logged:
_moe_quant_once_logged = True
logger.info(
"SGLANG_OPT_MOE_QUANT_ONCE: %s (layer %s)",
"ENGAGED" if self._moe_quant_once else f"INELIGIBLE: {reason}",
self.layer_id,
)
return self._moe_quant_once
def _compute_moe_quant_once_enabled(self) -> Tuple[bool, str]:
"""Returns (eligible, reason); reason names the first failing check."""
from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatcher
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod, Fp8MoEMethod
from sglang.srt.layers.quantization.fp8_utils import (
cutlass_w8a8_block_fp8_linear_with_fallback,
deepgemm_w8a8_block_fp8_linear_with_fallback,
)
if not envs.SGLANG_OPT_MOE_QUANT_ONCE.get():
return False, "env off"
if not _is_cuda:
return False, "not CUDA"
if self._enable_a2a_moe or self._fuse_shared_experts_inside_sbo:
return False, "a2a MoE or SBO shared-expert fusion"
# Shared-expert side: fp8 block-128 weights served by a w8a8 linear
# backend taught to accept a pre-quantized (q, scale) tuple: cutlass
# or deepgemm (fp32 scales only, i.e. not UE8M0/Blackwell).
if self.num_fused_shared_experts != 0 or not hasattr(self, "shared_experts"):
return False, "no separate shared experts"
if not self.shared_experts_is_fp8:
return False, "shared experts not fp8"
if self.shared_experts_weight_block_size != [128, 128]:
return False, "shared weight block size != [128, 128]"
gate_up = self.shared_experts.gate_up_proj
if not isinstance(gate_up.quant_method, Fp8LinearMethod):
return False, "shared gate_up quant method not Fp8LinearMethod"
linear_fn = gate_up.quant_method.w8a8_block_fp8_linear
if linear_fn is cutlass_w8a8_block_fp8_linear_with_fallback:
if gate_up.weight.shape[0] % 128 != 0 or gate_up.weight.shape[1] % 128 != 0:
return False, "gate_up weight shape unsupported by cutlass"
elif linear_fn is deepgemm_w8a8_block_fp8_linear_with_fallback:
from sglang.srt.layers import deep_gemm_wrapper
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
return False, "DeepGEMM UE8M0 scales (Blackwell) unsupported"
if gate_up.weight.shape[0] % 64 != 0 or gate_up.weight.shape[1] % 128 != 0:
return False, "gate_up weight shape unsupported by deepgemm"
else:
return False, f"w8a8 linear backend {linear_fn.__name__} unsupported"
# Routed side: standard dispatcher + triton fused func with dynamic
# per-token-group-128 fp8 activation quant.
experts = self.experts
if not isinstance(experts, FusedMoE):
return False, "experts not FusedMoE"
quant_method = experts.quant_method
if not isinstance(quant_method, Fp8MoEMethod):
return False, "experts quant method not Fp8MoEMethod"
if not quant_method.block_quant or quant_method.use_mxfp8:
return False, "experts not block-quant fp8"
if quant_method.quant_config.weight_block_size != [128, 128]:
return False, "experts weight block size != [128, 128]"
# Fp8MoEMethod only sets .runner for runner backends it drives itself.
runner = getattr(quant_method, "runner", None)
if runner is None or not runner.runner_backend.is_triton():
return False, "MoE runner backend not triton"
if runner.fused_func is None or runner.lora_enabled:
return False, "triton fused func unavailable (or LoRA enabled)"
if not isinstance(experts.dispatcher, StandardDispatcher):
return False, "dispatcher not StandardDispatcher"
if experts.moe_runner_config.apply_router_weight_on_input:
return False, "apply_router_weight_on_input"
if experts.w13_input_scale is not None:
return False, "static w13 input scale"
return True, "ok"
def _maybe_quant_moe_input_once(
self, hidden_states: torch.Tensor
) -> Optional[Tuple[torch.Tensor, torch.Tensor]]:
"""Quantize hidden_states once (per-token-group-128 fp8, rows padded to
a multiple of 4, column-major scales) for both the shared-expert GEMM
and the routed dispatch, or return None when ineligible."""
if hidden_states.shape[0] == 0 or hidden_states.dtype != torch.bfloat16:
return None
if not self._moe_quant_once_enabled():
return None
if is_in_tc_piecewise_cuda_graph():
# The piecewise MoE op quantizes internally; a pre-quant here
# would be dead work.
return None
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8_row_padded,
)
q, s = sglang_per_token_group_quant_fp8_row_padded(hidden_states, 128)
return q, s
def op_gate(self, state):
if state.hidden_states_mlp_input.shape[0] > 0:
# router_logits: (num_tokens, n_experts)