Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
bdaf2e2
dsv4.1: extract Top-k kernels and candidate helpers
hnyls2002 Sep 15, 2026
9b17c0d
dsv4.1: skip unwritten Top-k plans in metadata comparisons
hnyls2002 Sep 16, 2026
f51894f
fix dsv4 top-k edge cases and trim tests
BBuf Sep 16, 2026
ed1c649
dsv4.1: extract communication kernels and wrappers
hnyls2002 Sep 15, 2026
5873fba
dsv4.1: extract vocab gather and sharded greedy selection
hnyls2002 Sep 15, 2026
8c02375
fix sharded greedy test suite
hnyls2002 Sep 16, 2026
b297ac2
clarify sharded greedy docstring; split test cases; fix stale usage path
hnyls2002 Sep 16, 2026
d3d9a85
fix communication guards and cover ragged collectives
BBuf Sep 16, 2026
5a8166b
remove standalone nvlink communication test
BBuf Sep 16, 2026
3c5130f
dsv4.1: extract compression and metadata kernels
hnyls2002 Sep 15, 2026
b14c741
dsv4.1: extract KV store and dequantization paths
hnyls2002 Sep 15, 2026
a0ae65a
dsv4.1: preserve 64K prefill planner indices and reject sentinel coll…
hnyls2002 Sep 15, 2026
e36a479
fix c2 padding and trim metadata tests
BBuf Sep 16, 2026
1b94c71
dsv4.1: restore C2 padding test entry point
hnyls2002 Sep 16, 2026
ed29816
remove standalone dsv4 metadata tests
BBuf Sep 16, 2026
cde9da0
merge main; drop stale topk lineage
hnyls2002 Sep 16, 2026
313f98d
merge dsv4.1-communication
hnyls2002 Sep 16, 2026
be34bc2
merge main
hnyls2002 Sep 16, 2026
56efc4a
drop norm-only c2 entry and aliases; move small metadata into dsv4; n…
hnyls2002 Sep 16, 2026
c888f9e
drop unused fp4 torch reference quantizer
hnyls2002 Sep 16, 2026
f33acae
c2: drop the dead duplicate freqs_cis load
BBuf Sep 16, 2026
8af4e43
merge c1/c2 wrappers; split small metadata into its homes; move torch…
hnyls2002 Sep 16, 2026
fdae609
trim comments
hnyls2002 Sep 16, 2026
3b2108b
Merge branch 'main' into dsv4.1-metadata
hnyls2002 Sep 16, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
312 changes: 312 additions & 0 deletions python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh
Original file line number Diff line number Diff line change
@@ -0,0 +1,312 @@
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>

#include <sgl_kernel/math.cuh>
#include <sgl_kernel/type.cuh>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <sgl_kernel/warp.cuh>

#include <sgl_kernel/deepseek_v4/fp4_utils.cuh>
#include <sgl_kernel/deepseek_v4/fp8_utils.cuh>
#include <sgl_kernel/deepseek_v4/kv_layout.cuh>

#include <tvm/ffi/container/tensor.h>

#include <bit>
#include <cstdint>

namespace sglang {

/// \brief Ratio-1 decode compressor: RMSNorm and the whole main-KV write.
///
/// At ratio 1 the latent stands for the token itself, so the kernel's input is the
/// `wkv` GEMM output and its RoPE position is `positions`, not `positions - 1`.
/// `kv_output` is the pre-RoPE latent, for the index-K branch's `wk` projection.
struct Compress1DecodeParams {
const bf16_t* __restrict__ kv_input; // [num_tokens, kHeadDim] bf16
bf16_t* __restrict__ kv_output; // [num_tokens, kHeadDim] bf16, pre-RoPE
const bf16_t* __restrict__ norm_weight; // [kHeadDim] bf16
const float* __restrict__ freqs_cis; // [max_pos, kRopeDim] fp32, real/imag interleaved
const void* __restrict__ positions; // [num_tokens] PosT
const void* __restrict__ out_loc; // [num_tokens] LocT compressed slot; 0 marks a padded row
uint8_t* __restrict__ kvcache; // [npages, kPageBytes] uint8
float eps;
};

/// Elements per thread; 256 threads per token measured fastest on B200 decode batches.
/// At (512, 64) it also keeps the nope/rope split warp-aligned, as the fp8 amax reduction requires.
constexpr uint32_t kC1VecSize = 2;

/// \brief RMSNorm + RoPE tail + fp4 fake-quant + the FlashMLA store.
///
/// One CTA per token, `kHeadDim / kC1VecSize` threads over the row.
///
/// The three reductions have different widths and are not interchangeable: the RMSNorm
/// statistic spans the row, an fp8 store scale 64 elements, an fp4 block 16.
///
/// kLayout is the cache's page format: V4 and V41 store the fake-quantized value; V41_FP4
/// stores the e2m1 codes and their e4m3 scales directly, so the fp4 rounding happens once.
template <
int64_t kHeadDim,
int64_t kRopeDim,
int32_t kPageBits,
typename PosT,
typename LocT,
deepseek_v4::KVLayout kLayout,
bool kUsePDL>
__global__ __launch_bounds__(kHeadDim / kC1VecSize) void flash_c1_decode_kernel(
const __grid_constant__ Compress1DecodeParams params) {
using namespace device;
using deepseek_v4::KVLayout;
using deepseek_v4::fp8::cast_to_ue8m0;
using deepseek_v4::fp8::inv_scale_ue8m0;
using deepseek_v4::fp8::pack_fp8;

/// Threads over one token; the leading kNopeLanes carry the fp8 nope part, the rest the bf16 RoPE tail.
constexpr uint32_t kVecSize = kC1VecSize;
constexpr uint32_t kRowLanes = kHeadDim / kVecSize;
constexpr uint32_t kNopeLanes = (kHeadDim - kRopeDim) / kVecSize;
constexpr uint32_t kRowWarps = kRowLanes / kWarpThreads;
constexpr uint32_t kFp8Lanes = 64 / kVecSize;
constexpr uint32_t kFp4Lanes = deepseek_v4::fp4::kCompressedKVBlockSize / kVecSize;
using Paged = deepseek_v4::PagedKV<kLayout, kPageBits>;
static_assert(kHeadDim == 512 && kRopeDim == 64, "the FlashMLA layouts require (512, 64)");
static_assert(kHeadDim % kVecSize == 0 && kVecSize % 2 == 0);
static_assert(kRowLanes % kWarpThreads == 0, "a token owns a whole number of warps");
static_assert(kNopeLanes % kFp8Lanes == 0, "the nope part must end on an fp8 scale block");
static_assert(
(kHeadDim - kRopeDim) % deepseek_v4::fp4::kCompressedKVBlockSize == 0,
"no fp4 block may straddle the nope/rope seam");
static_assert(kFp8Lanes <= kWarpThreads && kFp4Lanes <= kWarpThreads);

using bf16_vec_t = AlignedVector<bf16x2_t, kVecSize / 2>;
using fp8_vec_t = AlignedVector<fp8x2_e4m3_t, kVecSize / 2>;
using freq_vec_t = AlignedVector<float, kVecSize>;

const uint32_t tx = threadIdx.x;
const uint32_t row = blockIdx.x;

// `out_loc` and `positions` are step metadata, independent of the PDL producer.
// Slots fit in int32; padded rows are suppressed at the cache store.
const auto out_loc = static_cast<int32_t>(static_cast<const LocT*>(params.out_loc)[row]);
const auto position = static_cast<int64_t>(static_cast<const PosT*>(params.positions)[row]);
PDLWaitPrimary<kUsePDL>();

float data[kVecSize];
bf16_vec_t latent;
{
bf16_vec_t input, weight;
input.load(params.kv_input + row * kHeadDim, tx);
weight.load(params.norm_weight, tx);

// `project` already returns bf16 at ratio 1, so `finish`'s `.to(bfloat16)`
// is a no-op and the statistic is taken over the loaded values as they are.
float local_sqrsum = 0.0f;
#pragma unroll
for (uint32_t j = 0; j < kVecSize / 2; ++j) {
const auto [x, y] = cast<fp32x2_t>(input[j]);
local_sqrsum += x * x;
local_sqrsum += y * y;
data[j * 2 + 0] = x;
data[j * 2 + 1] = y;
}

__shared__ float s_warp_sum[kRowWarps];
s_warp_sum[tx / kWarpThreads] = warp::reduce_sum(local_sqrsum);
__syncthreads();
float sqrsum = 0.0f;
#pragma unroll
for (uint32_t i = 0; i < kRowWarps; ++i) {
sqrsum += s_warp_sum[i];
}
constexpr float kInvHeadDim = 1.0f / static_cast<float>(kHeadDim);
const auto norm_factor = math::rsqrt(sqrsum * kInvHeadDim + params.eps);

#pragma unroll
for (uint32_t j = 0; j < kVecSize / 2; ++j) {
const auto [wx, wy] = cast<fp32x2_t>(weight[j]);
const auto x = data[j * 2 + 0] * norm_factor * wx;
const auto y = data[j * 2 + 1] * norm_factor * wy;
latent[j] = cast<bf16x2_t>(fp32x2_t{x, y});
}
}

latent.store(params.kv_output + row * kHeadDim, tx);
PDLTriggerSecondary<kUsePDL>();

// Match finish()'s bf16 rounding before the main-KV RoPE.
#pragma unroll
for (uint32_t j = 0; j < kVecSize / 2; ++j) {
const auto [x, y] = cast<fp32x2_t>(latent[j]);
data[j * 2 + 0] = x;
data[j * 2 + 1] = y;
}

if (tx >= kNopeLanes) {
// Match rope_tail()'s bf16 rounding: it ends in `.to(x.dtype)` before the fake-quant.
freq_vec_t freq;
freq.load(params.freqs_cis + position * kRopeDim, tx - kNopeLanes);
#pragma unroll
for (uint32_t j = 0; j < kVecSize / 2; ++j) {
const auto k = j * 2;
const auto x_real = data[k + 0];
const auto x_imag = data[k + 1];
const auto f_real = freq[k + 0];
const auto f_imag = freq[k + 1];
const auto rotated =
cast<bf16x2_t>(fp32x2_t{x_real * f_real - x_imag * f_imag, x_real * f_imag + x_imag * f_real});
const auto [r0, r1] = cast<fp32x2_t>(rotated);
data[k + 0] = r0;
data[k + 1] = r1;
}
}

if constexpr (kLayout == KVLayout::V41_FP4) {
// The fp4 cache takes the rotated bf16 value as is: its row quantizer is the fake quant, minus the dequant.
if (out_loc <= 0) return;
const auto kv_row = Paged::row(params.kvcache, out_loc);
return deepseek_v4::v41::store_row<kLayout>(kv_row.data, kv_row.scale, tx, data);
}

// FP4/E4M3 fake-quant over 16 elements, i.e. kFp4Lanes threads.
{
float amax = fabsf(data[0]);
#pragma unroll
for (uint32_t i = 1; i < kVecSize; ++i) {
amax = fmaxf(amax, fabsf(data[i]));
}
amax = warp::reduce_max<kFp4Lanes>(amax);
const auto scale = deepseek_v4::fp4::compressed_kv_scale(amax);
#pragma unroll
for (uint32_t i = 0; i < kVecSize / 2; ++i) {
const auto [x, y] = deepseek_v4::fp4::fake_quant_compressed_kv_x2({data[i * 2 + 0], data[i * 2 + 1]}, scale);
data[i * 2 + 0] = x;
data[i * 2 + 1] = y;
}
}

// A padded CUDA-graph row carries `out_loc == 0`, the reserved dummy slot, and must publish
// nothing: at ratio 1 the compressed slot is the FULL slot, so there is no other marker to read.
if (out_loc <= 0) return;
const auto kv_row = Paged::row(params.kvcache, out_loc);

if constexpr (kLayout == KVLayout::V41) {
// fp8 with one ue8m0 scale per 32 elements over the whole row, RoPE included.
return deepseek_v4::v41::store_row<kLayout>(kv_row.data, kv_row.scale, tx, data);
}

const auto value_ptr = kv_row.data;

if (tx >= kNopeLanes) {
bf16_vec_t rope_out;
#pragma unroll
for (uint32_t j = 0; j < kVecSize / 2; ++j) {
rope_out[j] = cast<bf16x2_t>(fp32x2_t{data[j * 2 + 0], data[j * 2 + 1]});
}
rope_out.store(value_ptr + (kHeadDim - kRopeDim), tx - kNopeLanes);
} else {
// fp8 e4m3 with one ue8m0 scale per 64 elements.
float abs_max = fabsf(data[0]);
#pragma unroll
for (uint32_t i = 1; i < kVecSize; ++i) {
abs_max = fmaxf(abs_max, fabsf(data[i]));
}
abs_max = warp::reduce_max<kFp8Lanes>(abs_max);
const auto scale_ue8m0 = cast_to_ue8m0(fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX);
const auto inv_scale = inv_scale_ue8m0(scale_ue8m0);
fp8_vec_t nope_out;
#pragma unroll
for (uint32_t i = 0; i < kVecSize / 2; ++i) {
nope_out[i] = pack_fp8(data[i * 2 + 0] * inv_scale, data[i * 2 + 1] * inv_scale);
}
nope_out.store(value_ptr, tx);
kv_row.scale[tx / kFp8Lanes] = scale_ue8m0;
}
}

/// \brief Host side of `flash_c1_decode_kernel`.
template <int64_t kHeadDim, int64_t kRopeDim, uint32_t kPageSize, deepseek_v4::KVLayout kLayout, bool kUsePDL>
struct FlashCompress1Kernel {
static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1;
static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes<kLayout>(kPageSize);
static constexpr uint32_t kBlockSize = kHeadDim / kC1VecSize;

static_assert(std::has_single_bit(kPageSize), "the page/slot split needs a power-of-two page");
static_assert(kBlockSize % device::kWarpThreads == 0 && kBlockSize <= 1024);
static_assert(kLayout != deepseek_v4::KVLayout::V4 || kPageBytes == host::div_ceil(584ll * kPageSize, 576) * 576);

template <typename PosT, typename LocT>
static constexpr auto kernel = flash_c1_decode_kernel<kHeadDim, kRopeDim, kPageBits, PosT, LocT, kLayout, kUsePDL>;

/// \brief The (`positions`, `out_loc`) dtype pair, resolved at run time.
static auto select(const bool pos_i32, const bool loc_i32) {
if (pos_i32) return loc_i32 ? kernel<int32_t, int32_t> : kernel<int32_t, int64_t>;
return loc_i32 ? kernel<int64_t, int32_t> : kernel<int64_t, int64_t>;
}

/// \brief RMSNorm + RoPE + fp4 fake-quant + the FlashMLA store, one launch.
///
/// \param kv_input `[num_tokens, kHeadDim]` bf16, the `wkv` projection.
/// \param kv_output `[num_tokens, kHeadDim]` bf16, the pre-RoPE latent.
/// \param norm_weight `[kHeadDim]` bf16.
/// \param freqs_cis `[max_pos, kRopeDim]` fp32, real/imag interleaved.
/// \param positions `[num_tokens]` int32 or int64, indexed as-is.
/// \param out_loc `[num_tokens]` int32 or int64, the compressed slot; `0` is a padded row.
/// \param kvcache `[npages, kPageBytes]` uint8, or the pool's fp8 view of it.
static void run_decode_fusion(
const tvm::ffi::TensorView kv_input,
const tvm::ffi::TensorView kv_output,
const tvm::ffi::TensorView norm_weight,
const tvm::ffi::TensorView freqs_cis,
const tvm::ffi::TensorView positions,
const tvm::ffi::TensorView out_loc,
const tvm::ffi::TensorView kvcache,
const float eps) {
using namespace host;

auto N = SymbolicSize{"num_tokens"};
auto device_ = SymbolicDevice{};
device_.set_options<kDLCUDA>();

TensorMatcher({N, kHeadDim}) //
.with_dtype<bf16_t>()
.with_device(device_)
.verify(kv_input)
.verify(kv_output);
TensorMatcher({kHeadDim}).with_dtype<bf16_t>().with_device(device_).verify(norm_weight);
// Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2.
TensorMatcher({-1, kRopeDim}).with_dtype<fp32_t>().with_device(device_).verify(freqs_cis);
// The scheduler's `out_cache_loc` (which `c1_out_loc` aliases at ratio 1)
// is int64; the unit tests hand int32. Both are indexed as-is.
auto pos_dtype = SymbolicDType{};
auto loc_dtype = SymbolicDType{};
TensorMatcher({N}).with_dtype<int32_t, int64_t>(pos_dtype).with_device(device_).verify(positions);
TensorMatcher({N}).with_dtype<int32_t, int64_t>(loc_dtype).with_device(device_).verify(out_loc);
// The pool allocates the buffer as uint8 and hands it out viewed as its fp8
// dtype (`get_extra_key_buffer`); both are one byte per element.
TensorMatcher({-1, kPageBytes}).with_dtype<uint8_t, fp8_e4m3_t>().with_device(device_).verify(kvcache);

const auto num_tokens = static_cast<uint32_t>(N.unwrap());
if (num_tokens == 0) return;

const auto params = Compress1DecodeParams{
.kv_input = static_cast<const bf16_t*>(kv_input.data_ptr()),
.kv_output = static_cast<bf16_t*>(kv_output.data_ptr()),
.norm_weight = static_cast<const bf16_t*>(norm_weight.data_ptr()),
.freqs_cis = static_cast<const float*>(freqs_cis.data_ptr()),
.positions = positions.data_ptr(),
.out_loc = out_loc.data_ptr(),
.kvcache = static_cast<uint8_t*>(kvcache.data_ptr()),
.eps = eps,
};
const auto k = select(pos_dtype.is_type<int32_t>(), loc_dtype.is_type<int32_t>());
LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) //
.enable_pdl(kUsePDL)(k, params);
}
};

// The JIT module names and wrappers spell the layouts as bare enumerators.
using enum deepseek_v4::KVLayout;

} // namespace sglang
Loading
Loading