fix(tuning_radix_sort): remove invented portioned_smem_per_warp field
The previous version had a `portioned_smem_per_warp` field that doesn't exist in CCCL. The actual CCCL RadixSortOnesweepPolicy has: threads, items, store_algorithm, rank_algorithm, scan_algorithm, rank_private_partitions, radix_bits Also adds proper SMEM calculation: total = max(keys_tile, values_tile, rank_smem) + offsets with 2KB headroom for kernel stack/locals. rank_private_partitions set to 1 to minimize SMEM pressure.
This commit is contained in:
@@ -1,15 +1,28 @@
|
|||||||
// muh/include/muh/tuning/tuning_radix_sort.cuh — BI-V100
|
// muh/include/muh/tuning/tuning_radix_sort.cuh — BI-V100
|
||||||
//
|
//
|
||||||
// Mirrors: cccl_upstream/cub/cub/device/dispatch/tuning/tuning_radix_sort.cuh
|
// Mirrors: cccl_upstream/cub/cub/device/dispatch/tuning/tuning_radix_sort.cuh
|
||||||
// CCCL: 2383 lines, chained_policy architecture (not sm100_tuning structs)
|
// CCCL: 2383 lines, chained_policy architecture with ONESWEEP on SM90+.
|
||||||
// Uses ONESWEEP algorithm on SM90+, fallback multi-pass on older.
|
|
||||||
//
|
//
|
||||||
// vllm relevance: top-p (nucleus) sampling sorts the full vocab logits
|
// SMEM analysis for ONESWEEP:
|
||||||
// SMEM risk: radix sort SMEM = ONESWEEP_RADIX_BITS^2 bins × sizeof(int) per warp
|
// TempStorage_ is a union of:
|
||||||
// bits=8 → 256 bins × 4B × (threads/32 warps) = 256*4*(512/32) = 16384 (safe)
|
// keys_out[TILE_ITEMS] = threads * items * sizeof(key_type)
|
||||||
// bits=11 → 2048 bins × 4B × 16 = 131072 (OVERFLOW at default threads)
|
// values_out[TILE_ITEMS] = threads * items * sizeof(value_type)
|
||||||
|
// rank_temp_storage (from BlockRadixRank)
|
||||||
|
// PLUS global_offsets[RADIX_DIGITS] = (1 << bits) * sizeof(OffsetT)
|
||||||
//
|
//
|
||||||
// Strategy: use ONESWEEP=true with bits=8 (safe), not 11
|
// For bits=8, threads=256, items=4, key=float32:
|
||||||
|
// keys_out = 256*4*4 = 4096
|
||||||
|
// offsets = 256*8 = 2048 (OffsetT=int64)
|
||||||
|
// total ≈ 6144 (safe)
|
||||||
|
//
|
||||||
|
// For bits=11: offsets = 2048*8 = 16384. rank_temp_storage with
|
||||||
|
// MATCH_EARLY_COUNTS uses per-warp privatized bins: 2048*num_parts*4.
|
||||||
|
// At num_parts=1, threads=256: just offsets + rank already ~32KB.
|
||||||
|
// But TILE_ITEMS = 256*4*4 = 4096 in the union, so total ≈ 36KB.
|
||||||
|
// Tight but might fit. However, rank_temp_storage for MATCH_EARLY_COUNTS
|
||||||
|
// with num_parts>1 can push past 48KB. Use bits=8 to be safe.
|
||||||
|
//
|
||||||
|
// vllm relevance: top-p (nucleus) sampling sorts full vocab (152064 logits)
|
||||||
|
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
@@ -18,19 +31,27 @@
|
|||||||
|
|
||||||
namespace muh::tuning::radix_sort {
|
namespace muh::tuning::radix_sort {
|
||||||
|
|
||||||
struct RadixSortHistogramPolicy {
|
// CCCL's actual RadixSortOnesweepPolicy fields (from tuning_radix_sort.cuh):
|
||||||
int threads_per_block;
|
// threads_per_block, items_per_thread, store_algorithm, rank_algorithm,
|
||||||
int items_per_thread;
|
// scan_algorithm, rank_private_partitions, radix_bits
|
||||||
int num_parts;
|
|
||||||
};
|
enum class RadixSortStoreAlgo { DIRECT, ALIGNED };
|
||||||
|
enum class RadixRankAlgo { MATCH, MATCH_EARLY_COUNTS_ANY, MATCH_EARLY_COUNTS_ATOMIC_OR };
|
||||||
|
|
||||||
struct RadixSortOnesweepPolicy {
|
struct RadixSortOnesweepPolicy {
|
||||||
int threads_per_block;
|
int threads_per_block;
|
||||||
int items_per_thread;
|
int items_per_thread;
|
||||||
|
RadixSortStoreAlgo store_algorithm;
|
||||||
|
RadixRankAlgo rank_algorithm;
|
||||||
|
BlockScanAlgorithm scan_algorithm;
|
||||||
|
int rank_private_partitions;
|
||||||
int radix_bits;
|
int radix_bits;
|
||||||
int rank_algorithm; // 0=MATCH, 1=MATCH_EARLY_COUNTS_ANY
|
};
|
||||||
BlockStoreAlgorithm store_algorithm;
|
|
||||||
int portioned_smem_per_warp;
|
struct RadixSortHistogramPolicy {
|
||||||
|
int threads_per_block;
|
||||||
|
int items_per_thread;
|
||||||
|
int num_parts;
|
||||||
};
|
};
|
||||||
|
|
||||||
struct RadixSortExclusiveSumPolicy {
|
struct RadixSortExclusiveSumPolicy {
|
||||||
@@ -64,38 +85,61 @@ struct policy_selector {
|
|||||||
bool keys_only;
|
bool keys_only;
|
||||||
|
|
||||||
constexpr RadixSortPolicy operator()(const hardware_capability& hw) const {
|
constexpr RadixSortPolicy operator()(const hardware_capability& hw) const {
|
||||||
bool is_onesweep = true; // SM90+ equivalent for BI-V100
|
// ONESWEEP with bits=8 is the safe choice for BI-V100.
|
||||||
|
// bits=11 risks SMEM overflow in rank_temp_storage with multiple partitions.
|
||||||
// SMEM-safe radix bits: 8 for all key sizes on BI-V100
|
constexpr int onesweep_bits = 8;
|
||||||
// bits=11 would need 2048 bins × warps × 4B → overflow
|
|
||||||
int onesweep_bits = 8;
|
|
||||||
int primary_bits = (key_size > 1) ? 7 : 5;
|
int primary_bits = (key_size > 1) ? 7 : 5;
|
||||||
int single_tile_bits = (key_size > 1) ? 6 : 5;
|
int single_tile_bits = (key_size > 1) ? 6 : 5;
|
||||||
int segmented_bits = (key_size > 1) ? 6 : 5;
|
int segmented_bits = (key_size > 1) ? 6 : 5;
|
||||||
|
|
||||||
int hist_items = (4 * 4) / key_size;
|
// items: 16 bytes per thread / key_size
|
||||||
if (hist_items < 1) hist_items = 1;
|
int items = 16 / key_size;
|
||||||
|
if (items < 1) items = 1;
|
||||||
|
|
||||||
int sweep_items = hist_items;
|
// SMEM check for onesweep: max(keys_tile, values_tile) + offsets
|
||||||
|
// keys_tile = threads * items * key_size
|
||||||
|
// offsets = (1 << bits) * 8 (OffsetT = int64)
|
||||||
|
int threads = 256;
|
||||||
|
int keys_tile = threads * items * key_size;
|
||||||
|
int offsets = (1 << onesweep_bits) * 8;
|
||||||
|
// rank_temp_storage: approximately radix_digits * sizeof(int) * num_parts
|
||||||
|
int rank_smem = (1 << onesweep_bits) * 4 * 1; // num_parts=1
|
||||||
|
int total_smem = keys_tile + offsets + rank_smem; // union: max(keys,values) not sum
|
||||||
|
// Actually it is a union, so: max(keys_tile, values_tile, rank_smem) + offsets
|
||||||
|
int val_tile = keys_only ? 0 : threads * items * value_size;
|
||||||
|
int main_union = keys_tile;
|
||||||
|
if (val_tile > main_union) main_union = val_tile;
|
||||||
|
if (rank_smem > main_union) main_union = rank_smem;
|
||||||
|
total_smem = main_union + offsets;
|
||||||
|
|
||||||
// Downsweep (fallback for non-onesweep path)
|
while (total_smem > hw.max_shared_memory_per_block - 2048 && items > 1) {
|
||||||
int ds_items = (4 * 4) / key_size;
|
// Leave 2KB headroom for kernel stack/locals
|
||||||
if (ds_items < 1) ds_items = 1;
|
items--;
|
||||||
|
keys_tile = threads * items * key_size;
|
||||||
|
val_tile = keys_only ? 0 : threads * items * value_size;
|
||||||
|
main_union = keys_tile > val_tile ? keys_tile : val_tile;
|
||||||
|
if (rank_smem > main_union) main_union = rank_smem;
|
||||||
|
total_smem = main_union + offsets;
|
||||||
|
}
|
||||||
|
|
||||||
return {
|
return {
|
||||||
is_onesweep,
|
true, // onesweep
|
||||||
primary_bits,
|
primary_bits,
|
||||||
single_tile_bits,
|
single_tile_bits,
|
||||||
segmented_bits,
|
segmented_bits,
|
||||||
// histogram
|
// histogram: same threads/items as onesweep
|
||||||
{256, hist_items, 1},
|
{threads, items, 1},
|
||||||
// exclusive_sum
|
// exclusive_sum
|
||||||
{256, onesweep_bits},
|
{256, onesweep_bits},
|
||||||
// onesweep
|
// onesweep
|
||||||
{256, sweep_items, onesweep_bits, 1, BLOCK_STORE_DIRECT,
|
{threads, items,
|
||||||
hw.max_shared_memory_per_block / (256 / hw.warp_size)},
|
RadixSortStoreAlgo::DIRECT,
|
||||||
// downsweep
|
RadixRankAlgo::MATCH_EARLY_COUNTS_ANY,
|
||||||
{256, ds_items, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
|
BLOCK_SCAN_WARP_SCANS,
|
||||||
|
1, // rank_private_partitions: 1 to minimize SMEM
|
||||||
|
onesweep_bits},
|
||||||
|
// downsweep (fallback for non-onesweep)
|
||||||
|
{256, items, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
|
||||||
primary_bits, BLOCK_SCAN_WARP_SCANS},
|
primary_bits, BLOCK_SCAN_WARP_SCANS},
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user