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