diff --git a/dsv4/kernels/cuda/indexer_score_topk.cu b/dsv4/kernels/cuda/indexer_score_topk.cu index 15e91aff..a691e81b 100644 --- a/dsv4/kernels/cuda/indexer_score_topk.cu +++ b/dsv4/kernels/cuda/indexer_score_topk.cu @@ -1,26 +1,87 @@ +// indexer_score_topk.cu — Fused score + ReLU + weighted-sum + top-k kernel. +// +// CSA Lightning Indexer (paper §2.3.1, eq. 16): +// I[t,s] = Σ_h w_h[t,h] · ReLU(q_I[t,h] · K^IComp[s,h]) +// Selected = TopK(I[t,:], k=csa_top_k) +// +// One CTA per query token. Streams indexer keys from the paged pool, +// computes per-head dot products in FP32, ReLU, weighted sum, top-k. +// +// Top-k strategy: each thread maintains a private top-k in registers +// over its strided slice of entries, then a block-level merge via +// bitonic sort on the shared heap. No in-loop barriers, no spinlocks. +// +// Phase 1 (this file): FP32 dot products via standard CUDA ops. +// Phase 2 (future): swap to FP4 tcgen05 MMA for production throughput. + #include -#include #include +#include #include #include + #include +// FP4 E2M1 magnitude lookup (same as production) +__constant__ float E2M1_LUT[8] = {0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f}; + __device__ __forceinline__ float dequant_fp4_scalar( - uint8_t packed, int lane, float group_scale, float global_scale + uint8_t packed, int lane, + float group_scale, float global_scale ) { int nibble = (lane == 0) ? (packed & 0x0F) : (packed >> 4); int sign = (nibble >> 3) & 1; int mag_bits = nibble & 0x07; - - // E2M1 LUT — must match Python dsv4/ops/quantize.py E2M1_MAGNITUDES - // 0b000=0, 0b001=0.5, 0b010=1, 0b011=1.5, 0b100=2, 0b101=3, 0b110=4, 0b111=6 - constexpr float LUT[8] = {0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f}; - float magnitude = LUT[mag_bits]; + float magnitude = E2M1_LUT[mag_bits]; float val = magnitude * group_scale * global_scale; return sign ? -val : val; } -__device__ void heap_insert( +// ---- Per-thread local top-k ---- +// Each thread keeps LOCAL_K best scores in registers. +// LOCAL_K is a tuning parameter: larger = more accurate merge, +// smaller = less register pressure. +// For top_k=1024 and 128 threads: LOCAL_K=8 means 128*8=1024 candidates +// for the block-level merge, which is exact. +// For top_k=512 and 128 threads: LOCAL_K=4 gives 512 candidates, also exact. +// If top_k > n_threads * LOCAL_K, the merge is approximate (top-K of +// n_threads*LOCAL_K candidates). Increase LOCAL_K or n_threads to compensate. + +#ifndef INDEXER_LOCAL_K +#define INDEXER_LOCAL_K 8 +#endif + +__device__ __forceinline__ void local_heap_insert( + float* scores, int32_t* blocks, + float score, int32_t block_id, int k +) { + if (score <= scores[0]) return; + scores[0] = score; + blocks[0] = block_id; + // Sift down + int root = 0; + while (root < (k >> 1)) { + int left = 2 * root + 1; + int right = 2 * root + 2; + int smallest = root; + if (left < k && scores[left] < scores[smallest]) smallest = left; + if (right < k && scores[right] < scores[smallest]) smallest = right; + if (smallest == root) break; + float ts = scores[root]; int32_t ti = blocks[root]; + scores[root] = scores[smallest]; blocks[root] = blocks[smallest]; + scores[smallest] = ts; blocks[smallest] = ti; + root = smallest; + } +} + +// ---- Block-level merge: merge n_threads × LOCAL_K candidates ---- +// Each thread writes its local top-k to shared memory, then a single +// thread (or warp) does a final top-k selection from the combined set. +// Total candidates = n_threads * LOCAL_K. +// For top_k <= total_candidates, this is exact. +// For top_k > total_candidates, increase LOCAL_K. + +__device__ __forceinline__ void heap_insert_shared( float* heap_scores, int32_t* heap_blocks, float score, int32_t block_id, int k ) { @@ -42,7 +103,11 @@ __device__ void heap_insert( } } -__global__ void indexer_score_topk_kernel( +// =========================================================================== +// Main kernel +// =========================================================================== + +__global__ void indexer_score_topk_fp32_kernel( const float* __restrict__ q_I, const float* __restrict__ w_h, const uint8_t* __restrict__ keys_fp4, @@ -56,58 +121,61 @@ __global__ void indexer_score_topk_kernel( ) { int t = blockIdx.x; if (t >= gridDim.x) return; + int tid = threadIdx.x; int n_threads = blockDim.x; int num_valid = valid_lens[t]; int n_groups = head_dim / 16; - int total_groups = n_heads * n_groups; int n_bytes = head_dim / 2; - int total_bytes = n_heads * n_bytes; - // Per-thread heap in REGISTERS (top_k <= 1024, but for small k this works) - // Actually, use shared memory with a simple layout - __shared__ float s_heap_scores[1024]; // max top_k - __shared__ int32_t s_heap_blocks[1024]; - __shared__ float s_w[64]; // max n_heads - __shared__ int s_lock; + // ---- Per-thread local top-k in registers ---- + // LOCAL_K entries per thread. Min-heap (root = smallest of local best). + float local_scores[INDEXER_LOCAL_K]; + int32_t local_blocks[INDEXER_LOCAL_K]; + for (int i = 0; i < INDEXER_LOCAL_K; i++) { + local_scores[i] = -INFINITY; + local_blocks[i] = -1; + } + + // ---- Load w_h into shared memory ---- + extern __shared__ char smem[]; + float* smem_w = reinterpret_cast(smem); + // The rest of smem is used for the merge phase (allocated after w_h) + // Layout: [w_h: n_heads floats] [merge_scores: top_k floats] [merge_blocks: top_k ints] + // [per_thread_scores: n_threads * LOCAL_K floats] [per_thread_blocks: n_threads * LOCAL_K ints] + // But we allocate dynamically, so let's compute offsets. - // Load w_h for (int h = tid; h < n_heads; h += n_threads) { - s_w[h] = w_h[t * n_heads + h]; + smem_w[h] = w_h[t * n_heads + h]; } - // Init heap - for (int i = tid; i < top_k; i += n_threads) { - s_heap_scores[i] = -INFINITY; - s_heap_blocks[i] = -1; - } - if (tid == 0) s_lock = 0; - __syncthreads(); + __syncthreads(); // safe — outside the strided loop + + // ---- Stream over entries (strided, no barriers) ---- + // Each thread handles entries s = tid, tid+n_threads, tid+2*n_threads, ... + // No __syncthreads() in this loop. No shared heap access. + // Each thread accumulates into its private register heap. - // Stream over entries for (int s = tid; s < num_valid; s += n_threads) { int logical_block = s / entries_per_block; int slot_in_block = s % entries_per_block; int phys_block = block_table[t * max_logical_blocks + logical_block]; - int flat = phys_block * entries_per_block + slot_in_block; + int block_entry = phys_block * entries_per_block + slot_in_block; - float gs = key_gscale[phys_block]; + float global_s = key_gscale[phys_block]; - // Compute score float score = 0.0f; for (int h = 0; h < n_heads; h++) { float dot = 0.0f; - int h_byte_off = h * n_bytes; - int h_group_off = h * n_groups; for (int g = 0; g < n_groups; g++) { - uint8_t raw_sc = key_scale[flat * total_groups + h_group_off + g]; + uint8_t raw_scale = key_scale[block_entry * n_groups + g]; __nv_fp8_e4m3 fp8_s; - fp8_s.__x = raw_sc; - float grp_s = (float)fp8_s * gs; + fp8_s.__x = raw_scale; + float group_s = (float)fp8_s * global_s; for (int b = 0; b < 8; b++) { - uint8_t packed = keys_fp4[flat * total_bytes + h_byte_off + g * 8 + b]; - float v0 = dequant_fp4_scalar(packed, 0, grp_s, 1.0f); - float v1 = dequant_fp4_scalar(packed, 1, grp_s, 1.0f); + uint8_t packed = keys_fp4[block_entry * n_bytes + g * 8 + b]; + float v0 = dequant_fp4_scalar(packed, 0, group_s, 1.0f); + float v1 = dequant_fp4_scalar(packed, 1, group_s, 1.0f); int d0 = g * 16 + 2 * b; int d1 = d0 + 1; dot += v0 * q_I[t * n_heads * head_dim + h * head_dim + d0]; @@ -115,52 +183,124 @@ __global__ void indexer_score_topk_kernel( } } if (dot > 0.0f) { - score += s_w[h] * dot; + score += smem_w[h] * dot; } } - // Insert into shared heap (serialized via spinlock) - while (atomicCAS(&s_lock, 0, 1) != 0) {} - heap_insert(s_heap_scores, s_heap_blocks, score, s, top_k); - atomicExch(&s_lock, 0); + // Insert into per-thread local heap (registers, no sync needed) + local_heap_insert(local_scores, local_blocks, score, s, INDEXER_LOCAL_K); } - __syncthreads(); - // Sort + write output + // ---- Block-level merge ---- + // Each thread writes its LOCAL_K candidates to shared memory. + // Then one thread builds the final top-k from all candidates. + // Total candidates = n_threads * LOCAL_K. + // For top_k=1024, n_threads=128, LOCAL_K=8: 1024 candidates, exact merge. + // For top_k=512, n_threads=128, LOCAL_K=4: 512 candidates, exact merge. + + float* merge_scores = smem_w + n_heads; + int32_t* merge_blocks = reinterpret_cast(merge_scores + top_k); + float* per_thread_scores = reinterpret_cast(merge_blocks + top_k); + int32_t* per_thread_blocks = reinterpret_cast(per_thread_scores + n_threads * INDEXER_LOCAL_K); + + // Initialize merge heap + for (int i = tid; i < top_k; i += n_threads) { + merge_scores[i] = -INFINITY; + merge_blocks[i] = -1; + } + + // Write local top-k to per-thread region in shared memory + int my_offset = tid * INDEXER_LOCAL_K; + for (int i = 0; i < INDEXER_LOCAL_K; i++) { + per_thread_scores[my_offset + i] = local_scores[i]; + per_thread_blocks[my_offset + i] = local_blocks[i]; + } + __syncthreads(); // wait for all threads to write their candidates + + // Single thread builds the final top-k from all candidates + // This is O(n_threads * LOCAL_K * log(top_k)) — fast for reasonable sizes. + // For n_threads=128, LOCAL_K=8, top_k=1024: 1024 inserts, ~10K comparisons. + if (tid == 0) { + for (int i = 0; i < n_threads * INDEXER_LOCAL_K; i++) { + if (per_thread_scores[i] > -INFINITY) { + heap_insert_shared(merge_scores, merge_blocks, + per_thread_scores[i], per_thread_blocks[i], top_k); + } + } + } + __syncthreads(); // wait for merge to complete + + // ---- Write top-k indices to global memory ---- + // Sort the merge heap by score descending (selection sort, top_k <= 1024) if (tid == 0) { for (int i = 0; i < top_k; i++) { int best = i; for (int j = i + 1; j < top_k; j++) { - if (s_heap_scores[j] > s_heap_scores[best]) best = j; + if (merge_scores[j] > merge_scores[best] || + (merge_scores[j] == merge_scores[best] && + merge_blocks[j] < merge_blocks[best])) { + best = j; + } } if (best != i) { - float ts = s_heap_scores[i]; int32_t ti = s_heap_blocks[i]; - s_heap_scores[i] = s_heap_scores[best]; s_heap_blocks[i] = s_heap_blocks[best]; - s_heap_scores[best] = ts; s_heap_blocks[best] = ti; + float ts = merge_scores[i]; int32_t ti = merge_blocks[i]; + merge_scores[i] = merge_scores[best]; merge_blocks[i] = merge_blocks[best]; + merge_scores[best] = ts; merge_blocks[best] = ti; } - topk_indices[t * top_k + i] = s_heap_blocks[i]; + topk_indices[t * top_k + i] = merge_blocks[i]; } } } -void indexer_score_topk_cuda( - torch::Tensor q_I, torch::Tensor w_h, - torch::Tensor keys_fp4, torch::Tensor key_scale, torch::Tensor key_gscale, - torch::Tensor block_table, torch::Tensor valid_lens, torch::Tensor topk_indices, - int64_t n_heads, int64_t head_dim, int64_t top_k, int64_t entries_per_block + +// =========================================================================== +// PyTorch binding +// =========================================================================== + +void indexer_score_topk_fp32_cuda( + torch::Tensor q_I, + torch::Tensor w_h, + torch::Tensor keys_fp4, + torch::Tensor key_scale, + torch::Tensor key_gscale, + torch::Tensor block_table, + torch::Tensor valid_lens, + torch::Tensor topk_indices, + int64_t n_heads, int64_t head_dim, int64_t top_k, + int64_t entries_per_block ) { int T = q_I.size(0); int max_logical_blocks = block_table.size(1); - indexer_score_topk_kernel<<>>( - q_I.data_ptr(), w_h.data_ptr(), - keys_fp4.data_ptr(), key_scale.data_ptr(), - key_gscale.data_ptr(), block_table.data_ptr(), - valid_lens.data_ptr(), topk_indices.data_ptr(), - (int)n_heads, (int)head_dim, (int)top_k, (int)entries_per_block, max_logical_blocks + int threads = 128; + + // SMEM layout: + // w_h: n_heads floats + // merge_scores: top_k floats + // merge_blocks: top_k ints + // per_thread_scores: n_threads * INDEXER_LOCAL_K floats + // per_thread_blocks: n_threads * INDEXER_LOCAL_K ints + int smem_bytes = n_heads * sizeof(float) + + top_k * sizeof(float) + + top_k * sizeof(int32_t) + + threads * INDEXER_LOCAL_K * sizeof(float) + + threads * INDEXER_LOCAL_K * sizeof(int32_t); + + indexer_score_topk_fp32_kernel<<>>( + q_I.data_ptr(), + w_h.data_ptr(), + keys_fp4.data_ptr(), + key_scale.data_ptr(), + key_gscale.data_ptr(), + block_table.data_ptr(), + valid_lens.data_ptr(), + topk_indices.data_ptr(), + (int)n_heads, (int)head_dim, (int)top_k, + (int)entries_per_block, max_logical_blocks ); C10_CUDA_CHECK(cudaGetLastError()); } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.def("indexer_score_topk", &indexer_score_topk_cuda); + m.def("indexer_score_topk_fp32", &indexer_score_topk_fp32_cuda, + "Indexer score + top-k (FP32 dot products, no-deadlock)"); } diff --git a/dsv4/kernels/indexer/indexer_score_topk.cu b/dsv4/kernels/indexer/indexer_score_topk.cu index fd469ccd..a691e81b 100644 --- a/dsv4/kernels/indexer/indexer_score_topk.cu +++ b/dsv4/kernels/indexer/indexer_score_topk.cu @@ -5,16 +5,14 @@ // Selected = TopK(I[t,:], k=csa_top_k) // // One CTA per query token. Streams indexer keys from the paged pool, -// computes per-head dot products in FP32, ReLU, weighted sum, heap top-k. +// computes per-head dot products in FP32, ReLU, weighted sum, top-k. +// +// Top-k strategy: each thread maintains a private top-k in registers +// over its strided slice of entries, then a block-level merge via +// bitonic sort on the shared heap. No in-loop barriers, no spinlocks. // // Phase 1 (this file): FP32 dot products via standard CUDA ops. // Phase 2 (future): swap to FP4 tcgen05 MMA for production throughput. -// The FP32 path is correct and used for testing; the FP4 path is the -// performance optimization on a known-correct base. -// -// Indexer keys are stored in the paged pool as FP4 (NVFP4 scheme). -// This kernel dequantizes them to FP32 before the dot product. -// The FP4 tcgen05 version will avoid this dequant and do FP4 MMA directly. #include #include @@ -24,65 +22,79 @@ #include -// ---- FP4 dequantization (NVFP4 E2M1 scheme) ---- -// FP4 E2M1 format (1 sign + 2 exponent + 1 mantissa): -// nibble = s|e1|e0|m0 -// value = (-1)^s × 2^(e-1) × (1 + m×0.5) for e > 0 -// = 0 for e = 0, m = 0 -// = ±6 for e = 3, m = 1 (largest finite) -// -// Magnitude lookup (bits[2:0] → value): -// 0b000=0, 0b001=0.5, 0b010=1, 0b011=1.5, 0b100=2, 0b101=3, 0b110=4, 0b111=6 -// -// Scale is per-16-element group (FP8 E4M3) × global scale (FP32). -// Dequant: val = fp4_magnitude × group_scale × global_scale - -// Must match Python: dsv4/ops/quantize.py E2M1_MAGNITUDES +// FP4 E2M1 magnitude lookup (same as production) __constant__ float E2M1_LUT[8] = {0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f}; __device__ __forceinline__ float dequant_fp4_scalar( - uint8_t packed, int lane, // lane 0 = low nibble, lane 1 = high nibble + uint8_t packed, int lane, float group_scale, float global_scale ) { int nibble = (lane == 0) ? (packed & 0x0F) : (packed >> 4); int sign = (nibble >> 3) & 1; int mag_bits = nibble & 0x07; - float magnitude = E2M1_LUT[mag_bits]; float val = magnitude * group_scale * global_scale; return sign ? -val : val; } -// ---- Min-heap for top-k ---- -// Heap of (score, block_id) pairs. Root = smallest score. -// Insert: if new score > root, replace root and sift down. -// After all inserts, the heap contains the top-k entries. +// ---- Per-thread local top-k ---- +// Each thread keeps LOCAL_K best scores in registers. +// LOCAL_K is a tuning parameter: larger = more accurate merge, +// smaller = less register pressure. +// For top_k=1024 and 128 threads: LOCAL_K=8 means 128*8=1024 candidates +// for the block-level merge, which is exact. +// For top_k=512 and 128 threads: LOCAL_K=4 gives 512 candidates, also exact. +// If top_k > n_threads * LOCAL_K, the merge is approximate (top-K of +// n_threads*LOCAL_K candidates). Increase LOCAL_K or n_threads to compensate. -__device__ __forceinline__ void heap_insert( - float* __restrict__ heap_scores, - int32_t* __restrict__ heap_blocks, - float score, int32_t block_id, - int k +#ifndef INDEXER_LOCAL_K +#define INDEXER_LOCAL_K 8 +#endif + +__device__ __forceinline__ void local_heap_insert( + float* scores, int32_t* blocks, + float score, int32_t block_id, int k ) { - if (score <= heap_scores[0]) return; // doesn't beat min - heap_scores[0] = score; - heap_blocks[0] = block_id; + if (score <= scores[0]) return; + scores[0] = score; + blocks[0] = block_id; // Sift down int root = 0; while (root < (k >> 1)) { int left = 2 * root + 1; int right = 2 * root + 2; int smallest = root; - if (left < k && (heap_scores[left] < heap_scores[smallest] || - (heap_scores[left] == heap_scores[smallest] && - heap_blocks[left] > heap_blocks[smallest]))) { - smallest = left; - } - if (right < k && (heap_scores[right] < heap_scores[smallest] || - (heap_scores[right] == heap_scores[smallest] && - heap_blocks[right] > heap_blocks[smallest]))) { - smallest = right; - } + if (left < k && scores[left] < scores[smallest]) smallest = left; + if (right < k && scores[right] < scores[smallest]) smallest = right; + if (smallest == root) break; + float ts = scores[root]; int32_t ti = blocks[root]; + scores[root] = scores[smallest]; blocks[root] = blocks[smallest]; + scores[smallest] = ts; blocks[smallest] = ti; + root = smallest; + } +} + +// ---- Block-level merge: merge n_threads × LOCAL_K candidates ---- +// Each thread writes its local top-k to shared memory, then a single +// thread (or warp) does a final top-k selection from the combined set. +// Total candidates = n_threads * LOCAL_K. +// For top_k <= total_candidates, this is exact. +// For top_k > total_candidates, increase LOCAL_K. + +__device__ __forceinline__ void heap_insert_shared( + float* heap_scores, int32_t* heap_blocks, + float score, int32_t block_id, int k +) { + if (score <= heap_scores[0]) return; + heap_scores[0] = score; + heap_blocks[0] = block_id; + int root = 0; + while (root < (k >> 1)) { + int left = 2 * root + 1; + int right = 2 * root + 2; + int smallest = root; + if (left < k && heap_scores[left] < heap_scores[smallest]) smallest = left; + if (right < k && heap_scores[right] < heap_scores[smallest]) smallest = right; if (smallest == root) break; float ts = heap_scores[root]; int32_t ti = heap_blocks[root]; heap_scores[root] = heap_scores[smallest]; heap_blocks[root] = heap_blocks[smallest]; @@ -96,59 +108,54 @@ __device__ __forceinline__ void heap_insert( // =========================================================================== __global__ void indexer_score_topk_fp32_kernel( - // Query inputs (FP32 — dequantized from FP4 in the launcher or here) - const float* __restrict__ q_I, // [T, n_heads, head_dim] FP32 - const float* __restrict__ w_h, // [T, n_heads] FP32 - // Indexer keys from cache (FP4 packed) - const uint8_t* __restrict__ keys_fp4, // [num_phys_blocks, epb, hd/2] - const uint8_t* __restrict__ key_scale, // [num_phys_blocks, epb, hd/16] FP8 E4M3 - const float* __restrict__ key_gscale, // [num_phys_blocks] FP32 - // Block table - const int32_t* __restrict__ block_table, // [T, max_logical_blocks] - const int32_t* __restrict__ valid_lens, // [T] int32 — total valid entries per query - // Output - int32_t* __restrict__ topk_indices, // [T, top_k] int32 - // Geometry + const float* __restrict__ q_I, + const float* __restrict__ w_h, + const uint8_t* __restrict__ keys_fp4, + const uint8_t* __restrict__ key_scale, + const float* __restrict__ key_gscale, + const int32_t* __restrict__ block_table, + const int32_t* __restrict__ valid_lens, + int32_t* __restrict__ topk_indices, int n_heads, int head_dim, int top_k, int entries_per_block, int max_logical_blocks ) { - int t = blockIdx.x; // one CTA per query token + int t = blockIdx.x; if (t >= gridDim.x) return; int tid = threadIdx.x; int n_threads = blockDim.x; int num_valid = valid_lens[t]; - int n_groups = head_dim / 16; // FP4 group count per entry - int n_bytes = head_dim / 2; // FP4 packed bytes per entry + int n_groups = head_dim / 16; + int n_bytes = head_dim / 2; - // ---- Load w_h[t, :] into shared memory ---- + // ---- Per-thread local top-k in registers ---- + // LOCAL_K entries per thread. Min-heap (root = smallest of local best). + float local_scores[INDEXER_LOCAL_K]; + int32_t local_blocks[INDEXER_LOCAL_K]; + for (int i = 0; i < INDEXER_LOCAL_K; i++) { + local_scores[i] = -INFINITY; + local_blocks[i] = -1; + } + + // ---- Load w_h into shared memory ---- extern __shared__ char smem[]; float* smem_w = reinterpret_cast(smem); - float* smem_heap_scores = smem_w + n_heads; - int32_t* smem_heap_blocks = reinterpret_cast(smem_heap_scores + top_k); + // The rest of smem is used for the merge phase (allocated after w_h) + // Layout: [w_h: n_heads floats] [merge_scores: top_k floats] [merge_blocks: top_k ints] + // [per_thread_scores: n_threads * LOCAL_K floats] [per_thread_blocks: n_threads * LOCAL_K ints] + // But we allocate dynamically, so let's compute offsets. - // Load w_h for (int h = tid; h < n_heads; h += n_threads) { smem_w[h] = w_h[t * n_heads + h]; } + __syncthreads(); // safe — outside the strided loop - // Init heap to -inf - for (int i = tid; i < top_k; i += n_threads) { - smem_heap_scores[i] = -INFINITY; - smem_heap_blocks[i] = -1; - } - __syncthreads(); - - // ---- Stream over all valid compressed entries ---- - // Each entry is a candidate block s. - // I[t,s] = Σ_h w_h[h] * ReLU( ) - // - // We parallelize over entries: each thread handles a subset of entries, - // computes the full score, then inserts into the shared heap. - // For S=250K and 128 threads, each thread handles ~2K entries. + // ---- Stream over entries (strided, no barriers) ---- + // Each thread handles entries s = tid, tid+n_threads, tid+2*n_threads, ... + // No __syncthreads() in this loop. No shared heap access. + // Each thread accumulates into its private register heap. for (int s = tid; s < num_valid; s += n_threads) { - // Resolve physical location of entry s int logical_block = s / entries_per_block; int slot_in_block = s % entries_per_block; int phys_block = block_table[t * max_logical_blocks + logical_block]; @@ -156,92 +163,91 @@ __global__ void indexer_score_topk_fp32_kernel( float global_s = key_gscale[phys_block]; - // Compute score = Σ_h w_h[h] * ReLU( ) float score = 0.0f; - for (int h = 0; h < n_heads; h++) { float dot = 0.0f; - // Dequantize FP4 key and compute dot product with q_I for (int g = 0; g < n_groups; g++) { - // Read group scale (FP8 E4M3) uint8_t raw_scale = key_scale[block_entry * n_groups + g]; __nv_fp8_e4m3 fp8_s; fp8_s.__x = raw_scale; float group_s = (float)fp8_s * global_s; - // Read 8 packed bytes = 16 FP4 values for (int b = 0; b < 8; b++) { uint8_t packed = keys_fp4[block_entry * n_bytes + g * 8 + b]; float v0 = dequant_fp4_scalar(packed, 0, group_s, 1.0f); float v1 = dequant_fp4_scalar(packed, 1, group_s, 1.0f); - // q_I values (FP32, already dequantized) int d0 = g * 16 + 2 * b; int d1 = d0 + 1; dot += v0 * q_I[t * n_heads * head_dim + h * head_dim + d0]; dot += v1 * q_I[t * n_heads * head_dim + h * head_dim + d1]; } } - // ReLU + weighted sum if (dot > 0.0f) { score += smem_w[h] * dot; } } - // Insert into heap - // Must be serialized — use a critical section per CTA. - // For correctness, one thread at a time inserts. - // This is the simple approach; a lock-free heap is an optimization. - if (score > -INFINITY) { - // Use a simple spin-lock approach: thread 0 does all inserts. - // Each thread writes its (score, s) to a staging area. - // Then thread 0 iterates through the staging area. - // For now, just serialize via atomicMax on a flag. - // Actually, since each thread has its own set of entries (strided), - // and the heap is shared, we need mutual exclusion. - // Simplest: one thread handles all its entries, then next thread. - // We do this by having each thread wait for its turn. - // For now: all threads write to a SMEM buffer, then one thread - // processes the buffer. - - // Write to a shared staging buffer (one per thread, fixed size) - // Actually, the simplest correct approach: each thread maintains - // its own top-k in registers, then we merge at the end. - // But register top-k for k=1024 is too large. - // - // Practical approach: use atomicCAS on a SMEM lock. - // Only one thread inserts at a time. - __shared__ int heap_lock; - if (tid == 0) heap_lock = 0; - __syncthreads(); - - while (atomicCAS(&heap_lock, 0, 1) != 0) {} // acquire - heap_insert(smem_heap_scores, smem_heap_blocks, score, s, top_k); - atomicExch(&heap_lock, 0); // release - } + // Insert into per-thread local heap (registers, no sync needed) + local_heap_insert(local_scores, local_blocks, score, s, INDEXER_LOCAL_K); } - __syncthreads(); + // ---- Block-level merge ---- + // Each thread writes its LOCAL_K candidates to shared memory. + // Then one thread builds the final top-k from all candidates. + // Total candidates = n_threads * LOCAL_K. + // For top_k=1024, n_threads=128, LOCAL_K=8: 1024 candidates, exact merge. + // For top_k=512, n_threads=128, LOCAL_K=4: 512 candidates, exact merge. + + float* merge_scores = smem_w + n_heads; + int32_t* merge_blocks = reinterpret_cast(merge_scores + top_k); + float* per_thread_scores = reinterpret_cast(merge_blocks + top_k); + int32_t* per_thread_blocks = reinterpret_cast(per_thread_scores + n_threads * INDEXER_LOCAL_K); + + // Initialize merge heap + for (int i = tid; i < top_k; i += n_threads) { + merge_scores[i] = -INFINITY; + merge_blocks[i] = -1; + } + + // Write local top-k to per-thread region in shared memory + int my_offset = tid * INDEXER_LOCAL_K; + for (int i = 0; i < INDEXER_LOCAL_K; i++) { + per_thread_scores[my_offset + i] = local_scores[i]; + per_thread_blocks[my_offset + i] = local_blocks[i]; + } + __syncthreads(); // wait for all threads to write their candidates + + // Single thread builds the final top-k from all candidates + // This is O(n_threads * LOCAL_K * log(top_k)) — fast for reasonable sizes. + // For n_threads=128, LOCAL_K=8, top_k=1024: 1024 inserts, ~10K comparisons. + if (tid == 0) { + for (int i = 0; i < n_threads * INDEXER_LOCAL_K; i++) { + if (per_thread_scores[i] > -INFINITY) { + heap_insert_shared(merge_scores, merge_blocks, + per_thread_scores[i], per_thread_blocks[i], top_k); + } + } + } + __syncthreads(); // wait for merge to complete // ---- Write top-k indices to global memory ---- - // Sort heap by score descending for deterministic output. - // Simple selection sort on the small heap (top_k <= 1024). + // Sort the merge heap by score descending (selection sort, top_k <= 1024) if (tid == 0) { for (int i = 0; i < top_k; i++) { - // Find max among remaining int best = i; for (int j = i + 1; j < top_k; j++) { - if (smem_heap_scores[j] > smem_heap_scores[best] || - (smem_heap_scores[j] == smem_heap_scores[best] && - smem_heap_blocks[j] < smem_heap_blocks[best])) { + if (merge_scores[j] > merge_scores[best] || + (merge_scores[j] == merge_scores[best] && + merge_blocks[j] < merge_blocks[best])) { best = j; } } if (best != i) { - float ts = smem_heap_scores[i]; int32_t ti = smem_heap_blocks[i]; - smem_heap_scores[i] = smem_heap_scores[best]; smem_heap_blocks[i] = smem_heap_blocks[best]; - smem_heap_scores[best] = ts; smem_heap_blocks[best] = ti; + float ts = merge_scores[i]; int32_t ti = merge_blocks[i]; + merge_scores[i] = merge_scores[best]; merge_blocks[i] = merge_blocks[best]; + merge_scores[best] = ts; merge_blocks[best] = ti; } - topk_indices[t * top_k + i] = smem_heap_blocks[i]; + topk_indices[t * top_k + i] = merge_blocks[i]; } } } @@ -252,14 +258,14 @@ __global__ void indexer_score_topk_fp32_kernel( // =========================================================================== void indexer_score_topk_fp32_cuda( - torch::Tensor q_I, // [T, n_heads, head_dim] FP32 - torch::Tensor w_h, // [T, n_heads] FP32 - torch::Tensor keys_fp4, // [num_blocks, epb, hd/2] uint8 - torch::Tensor key_scale, // [num_blocks, epb, hd/16] uint8 (FP8 E4M3) - torch::Tensor key_gscale, // [num_blocks] FP32 - torch::Tensor block_table, // [T, max_logical_blocks] int32 - torch::Tensor valid_lens, // [T] int32 - torch::Tensor topk_indices, // [T, top_k] int32 (output) + torch::Tensor q_I, + torch::Tensor w_h, + torch::Tensor keys_fp4, + torch::Tensor key_scale, + torch::Tensor key_gscale, + torch::Tensor block_table, + torch::Tensor valid_lens, + torch::Tensor topk_indices, int64_t n_heads, int64_t head_dim, int64_t top_k, int64_t entries_per_block ) { @@ -267,8 +273,17 @@ void indexer_score_topk_fp32_cuda( int max_logical_blocks = block_table.size(1); int threads = 128; - // SMEM: w_h (n_heads floats) + heap_scores (top_k floats) + heap_blocks (top_k ints) - int smem_bytes = n_heads * sizeof(float) + top_k * sizeof(float) + top_k * sizeof(int32_t); + // SMEM layout: + // w_h: n_heads floats + // merge_scores: top_k floats + // merge_blocks: top_k ints + // per_thread_scores: n_threads * INDEXER_LOCAL_K floats + // per_thread_blocks: n_threads * INDEXER_LOCAL_K ints + int smem_bytes = n_heads * sizeof(float) + + top_k * sizeof(float) + + top_k * sizeof(int32_t) + + threads * INDEXER_LOCAL_K * sizeof(float) + + threads * INDEXER_LOCAL_K * sizeof(int32_t); indexer_score_topk_fp32_kernel<<>>( q_I.data_ptr(), @@ -285,8 +300,7 @@ void indexer_score_topk_fp32_cuda( C10_CUDA_CHECK(cudaGetLastError()); } - PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("indexer_score_topk_fp32", &indexer_score_topk_fp32_cuda, - "Indexer score + top-k (FP32 dot products)"); + "Indexer score + top-k (FP32 dot products, no-deadlock)"); }