Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
104 changes: 39 additions & 65 deletions python/sglang/jit_kernel/csrc/deepseek_v4/c128_v2.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -89,10 +89,10 @@ struct C128Trait {
static_assert(kHeadDim % kTileDim == 0);
};

template <typename Trait, bool kUsePDL, typename BufFloat, typename InFloat, typename OutFloat>
template <typename Trait, bool kUsePDL, typename InFloat, typename OutFloat>
SGL_DEVICE void c128_forward(
const BufFloat* kv_buf, // [128n, 128n + 127]
const InFloat* kv_src, // ragged pointer at position = 128n + 127
const InFloat* kv_buf, // [128n, 128n + 127]
const InFloat* kv_src, // ragged pointer at position = 128n + 127
OutFloat* kv_out,
const InFloat* score_bias,
const int32_t buffer_len) {
Expand All @@ -101,15 +101,11 @@ SGL_DEVICE void c128_forward(
const auto warp_id = threadIdx.x / kWarpThreads;
const auto lane_id = threadIdx.x % kWarpThreads;

/// NOTE: part 1: load kv + score. kv_score_buffer (fp32, runtime state pool)
/// keeps its own BufFloat dtype; input/ape share InFloat (ape is cast to bf16
/// at load). Every value is converted to fp32 right after load.
using StorageBuf = AlignedVector<BufFloat, kTileElements>;
/// NOTE: part 1: load kv + score
using StorageIn = AlignedVector<InFloat, kTileElements>;
const auto gmem_buf = tile::Memory<StorageBuf>{lane_id, kWarpThreads};
const auto gmem_in = tile::Memory<StorageIn>{lane_id, kWarpThreads};
float kv[kElementsPerWarp][kTileElements];
float score[kElementsPerWarp][kTileElements];
StorageIn kv[kElementsPerWarp];
StorageIn score[kElementsPerWarp];
StorageIn bias[kElementsPerWarp];
const int32_t warp_offset = warp_id * kElementsPerWarp;

Expand All @@ -125,23 +121,9 @@ SGL_DEVICE void c128_forward(
for (int32_t i = 0; i < kElementsPerWarp; ++i) {
const int32_t j = i + warp_offset;
__builtin_assume(j < 128);
if (j < buffer_len) {
const auto k = gmem_buf.load(kv_buf + j * Trait::kElementSize);
const auto s = gmem_buf.load(kv_buf + j * Trait::kElementSize + Trait::kScoreOffset);
#pragma unroll
for (int32_t t = 0; t < kTileElements; ++t) {
kv[i][t] = cast<float>(k[t]);
score[i][t] = cast<float>(s[t]);
}
} else {
const auto k = gmem_in.load(kv_start + j * Trait::kElementSize);
const auto s = gmem_in.load(kv_start + j * Trait::kElementSize + Trait::kScoreOffset);
#pragma unroll
for (int32_t t = 0; t < kTileElements; ++t) {
kv[i][t] = cast<float>(k[t]);
score[i][t] = cast<float>(s[t]);
}
}
const auto src = j < buffer_len ? kv_buf : kv_start;
kv[i] = gmem_in.load(src + j * Trait::kElementSize);
score[i] = gmem_in.load(src + j * Trait::kElementSize + Trait::kScoreOffset);
}

/// NOTE: part 2: safe online softmax + weighted sum
Expand All @@ -156,11 +138,11 @@ SGL_DEVICE void c128_forward(

float score_fp32[kTileElements][kElementsPerWarp];

// kv/score already fp32 (converted at load); just add the bias
// convert to fp32 and apply bias first
#pragma unroll
for (int32_t i = 0; i < kTileElements; ++i) {
for (int32_t j = 0; j < kElementsPerWarp; ++j) {
score_fp32[i][j] = score[j][i] + cast<float>(bias[j][i]);
score_fp32[i][j] = cast<float>(score[j][i]) + cast<float>(bias[j][i]);
}
}

Expand All @@ -181,7 +163,7 @@ SGL_DEVICE void c128_forward(
for (int32_t j = 0; j < 8; ++j) {
const auto fp32_score = score[j];
const auto exp_score = expf(fp32_score - max_value);
sum_product += kv[j][i] * exp_score;
sum_product += cast<float>(kv[j][i]) * exp_score;
sum_exp_value += exp_score;
}

Expand Down Expand Up @@ -233,27 +215,25 @@ SGL_DEVICE void c128_forward(
}
}

template <typename Trait, typename BufFloat, typename InFloat>
SGL_DEVICE void c128_write_decode(BufFloat* kv_buf, const InFloat* kv_src) {
template <typename Trait, typename InFloat>
SGL_DEVICE void c128_write_decode(InFloat* kv_buf, const InFloat* kv_src) {
using namespace device;

using StorageIn = AlignedVector<InFloat, kTileElements>;
using StorageBuf = AlignedVector<BufFloat, kTileElements>;
const auto gmem_in = tile::Memory<StorageIn>::warp();
const auto gmem_buf = tile::Memory<StorageBuf>::warp();
using Storage = AlignedVector<InFloat, kTileElements>;
const auto gmem = tile::Memory<Storage>::warp();

Storage data[2];
#pragma unroll
for (int32_t i = 0; i < 2; ++i) {
const auto d = gmem_in.load(kv_src + Trait::kHeadDim * i);
StorageBuf o;
data[i] = gmem.load(kv_src + Trait::kHeadDim * i);
}
#pragma unroll
for (int32_t t = 0; t < kTileElements; ++t)
o[t] = cast<BufFloat>(d[t]);
gmem_buf.store(kv_buf + Trait::kHeadDim * i, o);
for (int32_t i = 0; i < 2; ++i) {
gmem.store(kv_buf + Trait::kHeadDim * i, data[i]);
}
}

template <int64_t kHeadDim, typename BufFloat, typename InFloat, typename OutFloat, bool kUsePDL>
template <int64_t kHeadDim, typename InFloat, typename OutFloat, bool kUsePDL>
C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) {
using namespace device;
using Trait = C128Trait<kHeadDim>;
Expand All @@ -267,7 +247,7 @@ C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodePara
const auto plan = params.plan_d[global_bid];
const auto kv_input = static_cast<const InFloat*>(params.kv_input) + split_offset;
const auto kv_output = static_cast<OutFloat*>(params.kv_output) + split_offset;
const auto kv_buffer = static_cast<BufFloat*>(params.kv_buffer) + split_offset;
const auto kv_buffer = static_cast<InFloat*>(params.kv_buffer) + split_offset;
const auto score_bias = static_cast<const InFloat*>(params.score_bias) + split_offset;

const auto kv_src = kv_input + global_bid * Trait::kElementSize;
Expand All @@ -278,15 +258,15 @@ C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodePara
PDLWaitPrimary<kUsePDL>();
// the write warp must match the load warp in the following `c128_forward`
if (warp_id == kNumWarps - 1) {
c128_write_decode<Trait, BufFloat, InFloat>(kv_dst, kv_src);
c128_write_decode<Trait>(kv_dst, kv_src);
}
if (plan.write_loc % 128 == 127) {
c128_forward<Trait, kUsePDL, BufFloat, InFloat, OutFloat>(kv_buf, kv_src, kv_out, score_bias, 128);
c128_forward<Trait, kUsePDL>(kv_buf, kv_src, kv_out, score_bias, 128);
}
}

// compress kernel
template <int64_t kHeadDim, typename BufFloat, typename InFloat, typename OutFloat, bool kUsePDL>
template <int64_t kHeadDim, typename InFloat, typename OutFloat, bool kUsePDL>
C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillParams params) {
using namespace device;
using Trait = C128Trait<kHeadDim>;
Expand All @@ -299,7 +279,7 @@ C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillPa
const auto plan = params.plan_c[global_pid];
const auto kv_input = static_cast<const InFloat*>(params.kv_input) + split_offset;
const auto kv_output = static_cast<OutFloat*>(params.kv_output) + split_offset;
const auto kv_buffer = static_cast<BufFloat*>(params.kv_buffer) + split_offset;
const auto kv_buffer = static_cast<InFloat*>(params.kv_buffer) + split_offset;
const auto score_bias = static_cast<const InFloat*>(params.score_bias) + split_offset;
if (plan.is_invalid()) return;

Expand All @@ -308,15 +288,14 @@ C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillPa
const auto kv_out = kv_output + global_pid * Trait::kHeadDim;
const auto kv_buf = kv_buffer + plan.read_page_1 * Trait::kPageElementSize;
PDLWaitPrimary<kUsePDL>();
c128_forward<Trait, kUsePDL, BufFloat, InFloat, OutFloat>(kv_buf, kv_src, kv_out, score_bias, plan.buffer_len);
c128_forward<Trait, kUsePDL>(kv_buf, kv_src, kv_out, score_bias, plan.buffer_len);
}

template <int64_t kHeadDim, typename BufFloat, typename InFloat, typename OutFloat, bool kUsePDL>
template <int64_t kHeadDim, typename InFloat, typename OutFloat, bool kUsePDL>
WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillParams params) {
using namespace device;
using Trait = C128Trait<kHeadDim>;
using StorageIn = AlignedVector<InFloat, kTileElements>;
using StorageBuf = AlignedVector<BufFloat, kTileElements>;

const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x;
const uint32_t global_wid = global_tid / kWarpThreads; // warp id
Expand All @@ -329,37 +308,32 @@ WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillP

const auto plan = params.plan_w[global_pid];
const auto kv_input = static_cast<const InFloat*>(params.kv_input) + split_offset;
const auto kv_buffer = static_cast<BufFloat*>(params.kv_buffer) + split_offset;
const auto kv_buffer = static_cast<InFloat*>(params.kv_buffer) + split_offset;
if (plan.is_invalid()) return;

// each warp will handle a contiguous region
const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize;
const auto kv_buf = kv_buffer + plan.write_loc * Trait::kElementSize;
const auto gmem_in = tile::Memory<StorageIn>::warp();
const auto gmem_buf = tile::Memory<StorageBuf>::warp();
const auto gmem = tile::Memory<StorageIn>::warp();

PDLWaitPrimary<kUsePDL>();
StorageIn data[2];
#pragma unroll
for (int32_t i = 0; i < 2; ++i) {
data[i] = gmem_in.load(kv_src, i);
data[i] = gmem.load(kv_src, i);
}
PDLTriggerSecondary<kUsePDL>();
#pragma unroll
for (int32_t i = 0; i < 2; ++i) {
StorageBuf o;
#pragma unroll
for (int32_t t = 0; t < kTileElements; ++t)
o[t] = cast<BufFloat>(data[i][t]);
gmem_buf.store(kv_buf, o, i);
gmem.store(kv_buf, data[i], i);
}
}

template <int64_t kHeadDim, typename BufFloat, typename InFloat, typename OutFloat, bool kUsePDL>
template <int64_t kHeadDim, typename InFloat, typename OutFloat, bool kUsePDL>
struct FlashCompress128Kernel {
static constexpr auto decode_kernel = flash_c128_decode<kHeadDim, BufFloat, InFloat, OutFloat, kUsePDL>;
static constexpr auto prefill_c_kernel = flash_c128_prefill<kHeadDim, BufFloat, InFloat, OutFloat, kUsePDL>;
static constexpr auto prefill_w_kernel = write_c128_prefill<kHeadDim, BufFloat, InFloat, OutFloat, kUsePDL>;
static constexpr auto decode_kernel = flash_c128_decode<kHeadDim, InFloat, OutFloat, kUsePDL>;
static constexpr auto prefill_c_kernel = flash_c128_prefill<kHeadDim, InFloat, OutFloat, kUsePDL>;
static constexpr auto prefill_w_kernel = write_c128_prefill<kHeadDim, InFloat, OutFloat, kUsePDL>;
static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64
static constexpr uint32_t kNumSplit = kHeadDim / kTileDim;
using Trait = C128Trait<kHeadDim>;
Expand All @@ -377,7 +351,7 @@ struct FlashCompress128Kernel {
device_.set_options<kDLGPU>();

TensorMatcher({-1, 128, Trait::kElementSize}) // kv score
.with_dtype<BufFloat>()
.with_dtype<InFloat>()
.with_device(device_)
.verify(kv_buffer);
TensorMatcher({N, Trait::kElementSize}) // kv score input
Expand Down Expand Up @@ -424,7 +398,7 @@ struct FlashCompress128Kernel {
device_.set_options<kDLGPU>();

TensorMatcher({-1, 128, Trait::kElementSize}) // kv score
.with_dtype<BufFloat>()
.with_dtype<InFloat>()
.with_device(device_)
.verify(kv_buffer);
TensorMatcher({N, Trait::kElementSize}) // kv score input (ragged)
Expand Down
Loading
Loading