fix(tuning_batched_topk): force bits=8, fix SMEM overflow
Previous version used base topk policy's bits (11 for key>=2B), causing SMEM overflow: 512*4*key_size + 2048*4*batches > 49152. Fix: force bits=8 (same as radix_sort decision for BI-V100). SMEM: 512*4*key_size + 256*4*batches = manageable. Also adds while-loop SMEM check on max_batches. Detected by test_smem_safety.py: 3 overflows at key_size=2,4,8.
This commit is contained in:
@@ -4,7 +4,14 @@
|
|||||||
// CCCL: extends topk with batch-level parallelism
|
// CCCL: extends topk with batch-level parallelism
|
||||||
//
|
//
|
||||||
// vllm relevance: multi-sequence parallel decode top-k sampling
|
// vllm relevance: multi-sequence parallel decode top-k sampling
|
||||||
// SMEM risk: same as topk (handled by bits_per_pass)
|
// SMEM risk: histogram SMEM = (1 << bits) * sizeof(int) * max_batches
|
||||||
|
// bits=11: 2048*4*batches — even at batches=1, keys_tile + 8192 can overflow
|
||||||
|
// bits=8: 256*4*batches — much safer
|
||||||
|
// Decision: use bits=8 (same as radix_sort) for all key sizes on BI-V100.
|
||||||
|
//
|
||||||
|
// SMEM layout: keys_tile (union with values_tile) + histogram per batch
|
||||||
|
// total = threads * items * max(key_size, value_size) + (1<<bits) * 4 * batches
|
||||||
|
// Must be ≤ 49152
|
||||||
|
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
@@ -26,15 +33,33 @@ struct policy_selector {
|
|||||||
constexpr BatchedTopkPolicy operator()(const hardware_capability& hw) const {
|
constexpr BatchedTopkPolicy operator()(const hardware_capability& hw) const {
|
||||||
auto base = topk::policy_selector{key_size}(hw);
|
auto base = topk::policy_selector{key_size}(hw);
|
||||||
|
|
||||||
// Batches per block: limited by SMEM
|
// Force bits=8 for BI-V100 (same decision as radix_sort)
|
||||||
// Each batch needs: bits_per_pass buckets * sizeof(int) for histogram
|
// bits=11 → 2048 buckets → histogram SMEM explodes with batching
|
||||||
int buckets = 1 << base.bits_per_pass;
|
int bits = 8;
|
||||||
int hist_smem = buckets * 4; // sizeof(int)
|
int buckets = 1 << bits; // 256
|
||||||
int max_batches = hw.max_shared_memory_per_block / hist_smem;
|
int hist_smem_per_batch = buckets * 4; // 1024 bytes per batch
|
||||||
|
|
||||||
|
// keys_tile for base policy
|
||||||
|
int keys_tile = base.threads_per_block * base.items_per_thread * key_size;
|
||||||
|
|
||||||
|
// Max batches: (SMEM - keys_tile) / hist_per_batch
|
||||||
|
int remaining_smem = hw.max_shared_memory_per_block - keys_tile;
|
||||||
|
int max_batches = remaining_smem > 0 ? remaining_smem / hist_smem_per_batch : 1;
|
||||||
if (max_batches < 1) max_batches = 1;
|
if (max_batches < 1) max_batches = 1;
|
||||||
if (max_batches > 32) max_batches = 32; // cap for occupancy
|
if (max_batches > 32) max_batches = 32; // cap for occupancy
|
||||||
|
|
||||||
return {base, max_batches};
|
// Verify total SMEM
|
||||||
|
int total_smem = keys_tile + hist_smem_per_batch * max_batches;
|
||||||
|
while (total_smem > hw.max_shared_memory_per_block && max_batches > 1) {
|
||||||
|
max_batches--;
|
||||||
|
total_smem = keys_tile + hist_smem_per_batch * max_batches;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Override base bits_per_pass to 8
|
||||||
|
topk::TopkPolicy adjusted_base = base;
|
||||||
|
adjusted_base.bits_per_pass = bits;
|
||||||
|
|
||||||
|
return {adjusted_base, max_batches};
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user