[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:
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.
|
||||
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));
|
||||
}
|
||||
|
||||
@@ -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<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 (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);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user