[DSA] Fix top-k v2 emitting invalid indices under tie overflow / inf scores (IMA in FA3 sparse decode) (#30645)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
DarkSharpness
2026-07-09 14:28:21 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent 8d0fd34150
commit bda1dc0d95
@@ -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. // Returns -inf for bin 0 (everything qualifies) and +inf for bins past the top.
template <uint32_t kBits> template <uint32_t kBits>
SGL_DEVICE float coarse_bin_lower_bound(uint32_t bin) { 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; constexpr uint32_t kShift = 16 - kBits;
const uint32_t key = bin << kShift; // ordered16 key at the low edge of `bin` 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) // ordered16 -> fp16 value (inverse of the transform in extract_coarse_bin);
const auto to_val = [](uint32_t okey) -> float { // finite keys only.
const auto to_finite_val = [](uint32_t okey) -> float {
const uint16_t ob = static_cast<uint16_t>(okey); 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); 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)); return cast<float>(*reinterpret_cast<const fp16_t*>(&hb));
}; };
// fp16 rounds to nearest, so the fp32 boundary is the midpoint between the fp16 // Fast path, hoisted above the per-key special cases so both keys are
// value at this key and the next-lower fp16 value (ordered key - 1). // 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)); 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 kBlockSize = 1024;
static constexpr uint32_t kOccupancy = 2; static constexpr uint32_t kOccupancy = 2;
static constexpr uint32_t kNumWarps = kBlockSize / kWarpSize; 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 kRadixSize = 1 << 8;
static constexpr uint32_t kTopKItems = (kMaxTopK + kBlockSize - 1) / kBlockSize; 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 TieHandleSmem {
struct alignas(16) MatchBin { struct alignas(16) MatchBin {
@@ -208,15 +244,11 @@ struct TopKConfig {
static_assert(kNumWarps == kWarpSize); static_assert(kNumWarps == kWarpSize);
if (num_ties <= topk) { if (num_ties <= topk) {
if (tx < num_ties) problem.emit(base + tx, tie_buffer[tx].idx); for (uint32_t t = tx; t < num_ties; t += kBlockSize) {
// Fewer tie candidates than remaining slots (ties beyond kMaxNumTie are problem.emit(base + t, tie_buffer[t].idx);
// 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 = num_ties + tx; t < topk; t += kBlockSize) { 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) { } else if (num_ties <= kWarpSize) {
if (lane_id >= num_ties || warp_id >= num_ties) return; // some threads are idle if (lane_id >= num_ties || warp_id >= num_ties) return; // some threads are idle
@@ -273,14 +305,45 @@ struct TopKConfig {
} }
if (lane_id == 0 && rank < topk) problem.emit(base + rank, target[i].idx); 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 { } else {
// Each thread loads one element (or becomes inactive) // Rare overflow case (kBlockSize < num_ties <= kMaxNumTie), kept out of
bool active = tx < num_ties; // the common path so it alone pays the multi-item register cost.
const auto tie = active ? tie_buffer[tx] : TieValue::invalid(); radix_tie_select<kTieItems>(tie_buffer, problem, base, num_ties, topk, smem);
const uint32_t key = extract_exact_bin(tie.value); }
const uint32_t idx = tie.idx; }
/// 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 (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; uint32_t topk_remain = topk;
uint32_t write_pos = topk;
if (tx < kRadixSize) smem->histogram[0][tx] = 0; if (tx < kRadixSize) smem->histogram[0][tx] = 0;
if (tx == kRadixSize) smem->counter = smem->counter_final = 0; if (tx == kRadixSize) smem->counter = smem->counter_final = 0;
__syncthreads(); __syncthreads();
@@ -289,12 +352,12 @@ struct TopKConfig {
#pragma unroll #pragma unroll
for (int round = 0; round < 4; round++) { for (int round = 0; round < 4; round++) {
const uint32_t shift = 24 - round * 8; const uint32_t shift = 24 - round * 8;
const uint32_t bin = (key >> shift) & 0xFFu;
const auto hist_idx = round % 2; const auto hist_idx = round % 2;
const auto histogram = smem->histogram[hist_idx]; const auto histogram = smem->histogram[hist_idx];
if (active) { #pragma unroll
atomicAdd(&histogram[bin], 1); for (uint32_t i = 0; i < kItems; ++i) {
if (active[i]) atomicAdd(&histogram[(key[i] >> shift) & 0xFFu], 1);
} }
if (round < 3 && tx < kRadixSize) { if (round < 3 && tx < kRadixSize) {
smem->histogram[hist_idx ^ 1][tx] = 0; smem->histogram[hist_idx ^ 1][tx] = 0;
@@ -325,14 +388,17 @@ struct TopKConfig {
topk_remain -= above_count; topk_remain -= above_count;
// 4. Scatter // 4. Scatter
if (active) { #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) { if (bin > threshold_bin) {
write_pos = atomicAdd(&smem->counter, 1); write_pos[i] = atomicAdd(&smem->counter, 1);
active = false; active[i] = false;
} else if (bin < threshold_bin) { } else if (bin < threshold_bin) {
active = false; active[i] = false;
} else if (round == 3) { } else if (round == 3) {
write_pos = topk - topk_remain + atomicAdd(&smem->counter_final, 1); write_pos[i] = topk - topk_remain + atomicAdd(&smem->counter_final, 1);
} }
// my_bin == thr && round < 3: stay active for next round // my_bin == thr && round < 3: stay active for next round
} }
@@ -340,7 +406,9 @@ struct TopKConfig {
if (round == 3 || topk_remain == 0) break; if (round == 3 || topk_remain == 0) break;
} }
if (write_pos < topk) problem.emit(base + write_pos, idx); #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; alignas(128) uint32_t count_gt;
uint32_t threshold_bin; uint32_t threshold_bin;
uint32_t warp_sum[kNumWarps]; 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 { union {
TieHandleSmem tie_handle_smem;
uint32_t histogram[kHistSize]; uint32_t histogram[kHistSize];
kHistVec hist_vecs[kBlockSize]; kHistVec hist_vecs[kBlockSize];
struct {
TieHandleSmem handle;
TieValue values[kMaxNumTie];
} tie;
}; };
TieValue tie_values[kMaxNumTie];
}; };
protected: protected:
@@ -514,7 +591,7 @@ struct TopKRegister : TopKRadixBase<12> {
} else if (val >= v_lo) { } else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1); const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]] if (count_eq < kMaxNumTie) [[likely]]
smem->tie_values[count_eq] = {val, idx}; smem->tie.values[count_eq] = {val, idx};
} }
}; };
#pragma unroll #pragma unroll
@@ -537,7 +614,7 @@ struct TopKRegister : TopKRadixBase<12> {
const auto equal_count = smem->count_eq; const auto equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0; const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie); 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) { } else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1); const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]] { 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 equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0; const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie); 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) { } else if (val >= v_lo) {
const auto count_eq = atomicAdd(&smem->count_eq, 1); const auto count_eq = atomicAdd(&smem->count_eq, 1);
if (count_eq < kMaxNumTie) [[likely]] { 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(); __syncthreads();
const auto start_gt_local = smem->start_gt_local; const auto start_gt_local = smem->start_gt_local;
const auto start_eq_local = smem->start_eq_local; const auto start_eq_local = smem->start_eq_local;
if (tx < local_equal_count && start_eq_local + tx < kMaxNumTie) { #pragma unroll
smem_0->tie_values[start_eq_local + tx] = smem->tie_values[tx]; 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; start_write = start_gt_local;
num_write = local_above_count; num_write = local_above_count;
@@ -753,7 +834,7 @@ struct TopKCluster : TopKRadixBase<10> {
const auto equal_count = smem->count_eq; const auto equal_count = smem->count_eq;
const auto remain_topk = above_count < topk ? topk - above_count : 0; const auto remain_topk = above_count < topk ? topk - above_count : 0;
const auto tie_count = min(equal_count, kMaxNumTie); 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);
} }
} }
}; };