diff --git a/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/topk_impl.cuh b/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/topk_impl.cuh index ffdf0f916..ec28a8f0a 100644 --- a/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/topk_impl.cuh +++ b/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/topk_impl.cuh @@ -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 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(okey); const uint16_t hb = (ob & 0x8000) ? static_cast(ob ^ 0x8000) : static_cast(~ob); return cast(*reinterpret_cast(&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::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)); } @@ -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 { @@ -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 @@ -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(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 + 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 (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 uint32_t bin = (key >> shift) & 0xFFu; - const auto hist_idx = round % 2; - const auto histogram = smem->histogram[hist_idx]; + 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]; - if (active) { - atomicAdd(&histogram[bin], 1); - } - if (round < 3 && tx < kRadixSize) { - smem->histogram[hist_idx ^ 1][tx] = 0; - } - __syncthreads(); +#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; + 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 (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 +#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); } - __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 - } - - if (round == 3 || topk_remain == 0) break; + // 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]); } } }; @@ -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: @@ -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 @@ -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); } }; @@ -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}; } } }); @@ -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); } }; @@ -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}; } } }); @@ -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; @@ -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); } } };