Skip to content
Merged
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
255 changes: 168 additions & 87 deletions python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/topk_impl.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -91,18 +91,46 @@ SGL_DEVICE uint32_t extract_coarse_bin(float x) {
// Returns -inf for bin 0 (everything qualifies) and +inf for bins past the top.
template <uint32_t kBits>
SGL_DEVICE float coarse_bin_lower_bound(uint32_t bin) {
if (bin == 0) return -FLT_MAX;
if (bin >= (1u << kBits)) return FLT_MAX;
constexpr uint32_t kShift = 16 - kBits;
const uint32_t key = bin << kShift; // ordered16 key at the low edge of `bin`
// ordered16 -> fp16 value (inverse of the transform in extract_coarse_bin)
const auto to_val = [](uint32_t okey) -> float {
// ordered16 -> fp16 value (inverse of the transform in extract_coarse_bin);
// finite keys only.
const auto to_finite_val = [](uint32_t okey) -> float {
const uint16_t ob = static_cast<uint16_t>(okey);
const uint16_t hb = (ob & 0x8000) ? static_cast<uint16_t>(ob ^ 0x8000) : static_cast<uint16_t>(~ob);
return cast<float>(*reinterpret_cast<const fp16_t*>(&hb));
};
// fp16 rounds to nearest, so the fp32 boundary is the midpoint between the fp16
// value at this key and the next-lower fp16 value (ordered key - 1).
// Fast path, hoisted above the per-key special cases so both keys are
// range-checked at once: `key` and `key - 1` both land in the finite band
// [0x0401, 0xFBFF] -- every boundary a finite-score threshold produces.
// fp16 rounds to nearest, so the fp32 boundary is the midpoint between the
// fp16 values at `key` and `key - 1`. (Verified bit-exact against the slow
// path for every bin of kBits 10 and 12, and measured faster than either
// per-key dispatch or an ordered-bit decrement trick -- the two conversions
// are independent and issue in parallel.)
if (key - 0x0401u <= 0xFBFFu - 0x0401u && bin < (1u << kBits)) {
return 0.5f * (to_finite_val(key) + to_finite_val(key - 1));
}
// Slow path: an edge of `bin` touches the +/-inf keys or NaN key space.
// The ordered-key line is: [0, 0x03FF) negative-NaN space, 0x03FF = -inf,
// [0x0400, 0xFC00) finite, 0xFC00 = +inf, (0xFC00, 0xFFFF] positive-NaN
// space. Treat the +/-inf keys as +/-65536 (one ideal step past fp16 max,
// so the midpoint lands exactly on +/-65520 -- the fp32->fp16
// round-to-nearest overflow threshold) and saturate NaN-space keys, keeping
// the returned boundaries finite-or-inf and monotone. Otherwise a threshold
// bin at/next to the inf bin gets NaN boundaries, the collect pass matches
// nothing, and rows whose scores contain >= topk (+/-)inf or >65504 values
// come back short -- the padded slots then illegal-address downstream.
if (bin == 0) return -FLT_MAX;
if (bin >= (1u << kBits)) return FLT_MAX;
const auto to_val = [&](uint32_t okey) -> float {
constexpr float kInf = std::numeric_limits<float>::infinity();
if (okey < 0x03FFu) return -kInf;
if (okey == 0x03FFu) return -65536.0f;
if (okey == 0xFC00u) return 65536.0f;
if (okey > 0xFC00u) return FLT_MAX;
return to_finite_val(okey);
};
return 0.5f * (to_val(key) + to_val(key - 1));
}

Expand Down Expand Up @@ -170,10 +198,18 @@ struct TopKConfig {
static constexpr uint32_t kBlockSize = 1024;
static constexpr uint32_t kOccupancy = 2;
static constexpr uint32_t kNumWarps = kBlockSize / kWarpSize;
static constexpr uint32_t kMaxNumTie = 1024;
// kMaxNumTie must be >= kMaxTopK: the collect pass keeps at most kMaxNumTie
// threshold-bin candidates, and up to `topk` output slots may have to be
// filled from them (above_count can be 0, e.g. heavily tied or all-inf
// scores). A smaller cap leaves slots that handle_tie can only pad, and
// padded slots inside the first min(seq_len, topk) entries are dereferenced
// by downstream sparse attention.
static constexpr uint32_t kMaxNumTie = 2048;
static constexpr uint32_t kRadixSize = 1 << 8;
static constexpr uint32_t kTopKItems = (kMaxTopK + kBlockSize - 1) / kBlockSize;
static_assert(kMaxNumTie <= kBlockSize && kBlockSize % kNumWarps == 0);
// tie candidates owned per thread in the strided handle_tie loops
static constexpr uint32_t kTieItems = kMaxNumTie / kBlockSize;
static_assert(kMaxNumTie >= kMaxTopK && kMaxNumTie % kBlockSize == 0 && kBlockSize % kNumWarps == 0);

struct TieHandleSmem {
struct alignas(16) MatchBin {
Expand Down Expand Up @@ -208,15 +244,11 @@ struct TopKConfig {
static_assert(kNumWarps == kWarpSize);

if (num_ties <= topk) {
if (tx < num_ties) problem.emit(base + tx, tie_buffer[tx].idx);
// Fewer tie candidates than remaining slots (ties beyond kMaxNumTie are
// dropped at collect): pad [num_ties, topk) with -1 ("no token"). The
// transform pass reads all `topk` output slots, and any slot left
// unwritten holds uninitialized staging memory whose page-table
// translation yields a garbage KV index (-> illegal memory access in
// the downstream sparse attention kernel).
for (uint32_t t = tx; t < num_ties; t += kBlockSize) {
problem.emit(base + t, tie_buffer[t].idx);
}
for (uint32_t t = num_ties + tx; t < topk; t += kBlockSize) {
problem.emit(base + t, -1u);
problem.emit(base + t, base + t);
}
} else if (num_ties <= kWarpSize) {
if (lane_id >= num_ties || warp_id >= num_ties) return; // some threads are idle
Expand Down Expand Up @@ -273,74 +305,110 @@ struct TopKConfig {
}
if (lane_id == 0 && rank < topk) problem.emit(base + rank, target[i].idx);
}
} else if (num_ties <= kBlockSize) {
// Common case: one candidate per thread.
radix_tie_select<1>(tie_buffer, problem, base, num_ties, topk, smem);
} else {
// Each thread loads one element (or becomes inactive)
bool active = tx < num_ties;
const auto tie = active ? tie_buffer[tx] : TieValue::invalid();
const uint32_t key = extract_exact_bin(tie.value);
const uint32_t idx = tie.idx;
uint32_t topk_remain = topk;
uint32_t write_pos = topk;
if (tx < kRadixSize) smem->histogram[0][tx] = 0;
if (tx == kRadixSize) smem->counter = smem->counter_final = 0;
__syncthreads();
uint32_t total_active = num_ties;
// Rare overflow case (kBlockSize < num_ties <= kMaxNumTie), kept out of
// the common path so it alone pays the multi-item register cost.
radix_tie_select<kTieItems>(tie_buffer, problem, base, num_ties, topk, smem);
}
}

/// Exact radix select over the tie candidates: each thread owns kItems
/// strided elements (inactive beyond num_ties). Requires
/// num_ties <= kItems * kBlockSize.
template <uint32_t kItems>
SGL_DEVICE static void radix_tie_select( //
const TieValue* tie_buffer,
const TopKProblem& problem,
const uint32_t base,
const uint32_t num_ties,
const uint32_t topk,
TieHandleSmem* smem) {
const auto tx = threadIdx.x;
const auto lane_id = tx % kWarpSize;
const auto warp_id = tx / kWarpSize;

bool active[kItems];
uint32_t key[kItems];
uint32_t idx[kItems];
uint32_t write_pos[kItems];
#pragma unroll
for (int round = 0; round < 4; round++) {
const uint32_t shift = 24 - round * 8;
const uint32_t bin = (key >> shift) & 0xFFu;
const auto hist_idx = round % 2;
const auto histogram = smem->histogram[hist_idx];

if (active) {
atomicAdd(&histogram[bin], 1);
}
if (round < 3 && tx < kRadixSize) {
smem->histogram[hist_idx ^ 1][tx] = 0;
}
__syncthreads();

uint32_t hist_val = 0;
uint32_t warp_inc = 0;
if (tx < kRadixSize) {
hist_val = histogram[tx];
warp_inc = warp_inclusive_sum(lane_id, hist_val);
if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc;
}
__syncthreads();
if (tx < kRadixSize) {
const auto inter = warp::reduce_sum(lane_id < warp_id ? smem->warp_sum[lane_id] : 0);
const auto prefix = inter + warp_inc; // inclusive prefix through this bin
const auto above = total_active - prefix; // elements in bins ABOVE this one
// 3. Find threshold bin
if (above < topk_remain && above + hist_val >= topk_remain) {
smem->match = {tx, above, hist_val};
}
}
__syncthreads();

const auto [threshold_bin, above_count, equal_count, __] = smem->match;
if (round < 3) total_active = equal_count;
topk_remain -= above_count;

// 4. Scatter
if (active) {
if (bin > threshold_bin) {
write_pos = atomicAdd(&smem->counter, 1);
active = false;
} else if (bin < threshold_bin) {
active = false;
} else if (round == 3) {
write_pos = topk - topk_remain + atomicAdd(&smem->counter_final, 1);
}
// my_bin == thr && round < 3: stay active for next round
for (uint32_t i = 0; i < kItems; ++i) {
const auto t = tx + i * kBlockSize;
active[i] = t < num_ties;
const auto tie = active[i] ? tie_buffer[t] : TieValue::invalid();
key[i] = extract_exact_bin(tie.value);
idx[i] = tie.idx;
write_pos[i] = topk;
}
uint32_t topk_remain = topk;
if (tx < kRadixSize) smem->histogram[0][tx] = 0;
if (tx == kRadixSize) smem->counter = smem->counter_final = 0;
__syncthreads();
uint32_t total_active = num_ties;

#pragma unroll
for (int round = 0; round < 4; round++) {
const uint32_t shift = 24 - round * 8;
const auto hist_idx = round % 2;
const auto histogram = smem->histogram[hist_idx];

#pragma unroll
for (uint32_t i = 0; i < kItems; ++i) {
if (active[i]) atomicAdd(&histogram[(key[i] >> shift) & 0xFFu], 1);
}
if (round < 3 && tx < kRadixSize) {
smem->histogram[hist_idx ^ 1][tx] = 0;
}
__syncthreads();

uint32_t hist_val = 0;
uint32_t warp_inc = 0;
if (tx < kRadixSize) {
hist_val = histogram[tx];
warp_inc = warp_inclusive_sum(lane_id, hist_val);
if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc;
}
__syncthreads();
if (tx < kRadixSize) {
const auto inter = warp::reduce_sum(lane_id < warp_id ? smem->warp_sum[lane_id] : 0);
const auto prefix = inter + warp_inc; // inclusive prefix through this bin
const auto above = total_active - prefix; // elements in bins ABOVE this one
// 3. Find threshold bin
if (above < topk_remain && above + hist_val >= topk_remain) {
smem->match = {tx, above, hist_val};
}
}
__syncthreads();

if (round == 3 || topk_remain == 0) break;
const auto [threshold_bin, above_count, equal_count, __] = smem->match;
if (round < 3) total_active = equal_count;
topk_remain -= above_count;

// 4. Scatter
#pragma unroll
for (uint32_t i = 0; i < kItems; ++i) {
if (!active[i]) continue;
const uint32_t bin = (key[i] >> shift) & 0xFFu;
if (bin > threshold_bin) {
write_pos[i] = atomicAdd(&smem->counter, 1);
active[i] = false;
} else if (bin < threshold_bin) {
active[i] = false;
} else if (round == 3) {
write_pos[i] = topk - topk_remain + atomicAdd(&smem->counter_final, 1);
}
// my_bin == thr && round < 3: stay active for next round
}

if (write_pos < topk) problem.emit(base + write_pos, idx);
if (round == 3 || topk_remain == 0) break;
}

#pragma unroll
for (uint32_t i = 0; i < kItems; ++i) {
if (write_pos[i] < topk) problem.emit(base + write_pos[i], idx[i]);
}
}
};
Expand All @@ -362,12 +430,21 @@ struct TopKRadixBase : TopKConfig {
alignas(128) uint32_t count_gt;
uint32_t threshold_bin;
uint32_t warp_sum[kNumWarps];
// The coarse histogram is dead once find_threshold() has published
// threshold_bin, and the tie machinery only comes alive after that: the
// collect pass fills tie.values, then handle_tie works over them with
// tie.handle as scratch. Overlaying the two phases keeps the
// kMaxNumTie-candidate buffer from growing the block's shared-memory
// footprint. tie.handle and tie.values are live TOGETHER, so they sit
// side by side inside the overlay, not in a union with each other.
union {
TieHandleSmem tie_handle_smem;
uint32_t histogram[kHistSize];
kHistVec hist_vecs[kBlockSize];
struct {
TieHandleSmem handle;
TieValue values[kMaxNumTie];
} tie;
};
TieValue tie_values[kMaxNumTie];
};

protected:
Expand Down Expand Up @@ -514,7 +591,7 @@ struct TopKRegister : TopKRadixBase<12> {
} else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]]
smem->tie_values[count_eq] = {val, idx};
smem->tie.values[count_eq] = {val, idx};
}
};
#pragma unroll
Expand All @@ -537,7 +614,7 @@ struct TopKRegister : TopKRadixBase<12> {
const auto equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie);
handle_tie(smem->tie_values, problem, above_count, tie_count, remain_topk, &smem->tie_handle_smem);
handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle);
}
};

Expand Down Expand Up @@ -594,7 +671,7 @@ struct TopKStreaming : TopKRegister<2> {
} else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]] {
smem->tie_values[count_eq] = {val, idx};
smem->tie.values[count_eq] = {val, idx};
}
}
});
Expand All @@ -609,7 +686,7 @@ struct TopKStreaming : TopKRegister<2> {
const auto equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie);
handle_tie(smem->tie_values, problem, above_count, tie_count, remain_topk, &smem->tie_handle_smem);
handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle);
}
};

Expand Down Expand Up @@ -710,7 +787,7 @@ struct TopKCluster : TopKRadixBase<10> {
} else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]] {
smem->tie_values[count_eq] = {val, idx};
smem->tie.values[count_eq] = {val, idx};
}
}
});
Expand All @@ -732,8 +809,12 @@ struct TopKCluster : TopKRadixBase<10> {
__syncthreads();
const auto start_gt_local = smem->start_gt_local;
const auto start_eq_local = smem->start_eq_local;
if (tx < local_equal_count && start_eq_local + tx < kMaxNumTie) {
smem_0->tie_values[start_eq_local + tx] = smem->tie_values[tx];
#pragma unroll
for (uint32_t i = 0; i < kTieItems; ++i) {
const auto t = tx + i * kBlockSize;
if (t < local_equal_count && start_eq_local + t < kMaxNumTie) {
smem_0->tie.values[start_eq_local + t] = smem->tie.values[t];
}
}
start_write = start_gt_local;
num_write = local_above_count;
Expand All @@ -753,7 +834,7 @@ struct TopKCluster : TopKRadixBase<10> {
const auto equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie);
handle_tie(smem->tie_values, problem, above_count, tie_count, remain_topk, &smem->tie_handle_smem);
handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle);
}
}
};
Expand Down
Loading