From 2d1588d261ce010c63a41f7f4894b7838d50f87b Mon Sep 17 00:00:00 2001 From: muh-pipeline Date: Thu, 6 Aug 2026 02:38:51 +0000 Subject: [PATCH] [BASE] qwen3_6_scripts/sampler.py: CCCL dispatch_topk DoubleBuffer pattern MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Random CCCL pick: cub/cub/device/dispatch/dispatch_topk.cuh (480 lines, full read) CCCL's DeviceTopK uses DoubleBuffer to ping-pong between two pre-allocated buffers across radix passes, achieving zero allocation in the hot loop. Our sampler.py's _apply_top_k_top_p was allocating 2 new tensors (logits_sort + logits_idx, each vocab_size×4B = 600KB) on every single decode step via torch.sort(). Change: cache sort output tensors keyed on (batch, vocab, device) and reuse them via torch.sort(..., out=(cached_sort, cached_idx)). This eliminates 1.2MB of GPU allocation per decode step. For competition max_num_seqs=1, vocab=152064: Before: 2 × 152064 × 4B = 1.2MB allocated per step After: 0 bytes allocated per step (reuse cached buffers) At 395 tokens/sec target: saves 474MB/sec of allocator pressure. BI-V100 has no async CUDA allocator, so this is synchronous overhead. CCCL architecture insight used: dispatch_topk.cuh line ~430: DoubleBuffer key_bufs(alloc[3], alloc[2]) for pass: key_bufs.Current() → read, key_bufs.Alternate() → write, swap Base file modified: qwen3_6_scripts/sampler.py (deployed via patch_ops.sh) --- qwen3_6_scripts/sampler.py | 48 +++++++++++++++++++++++++++----------- 1 file changed, 35 insertions(+), 13 deletions(-) diff --git a/qwen3_6_scripts/sampler.py b/qwen3_6_scripts/sampler.py index 3129da68..88f76632 100644 --- a/qwen3_6_scripts/sampler.py +++ b/qwen3_6_scripts/sampler.py @@ -440,34 +440,56 @@ def _apply_top_k_top_p( p: torch.Tensor, k: torch.Tensor, ) -> torch.Tensor: - # CCCL insight from tuning_topk.cuh: radix select (used by torch.topk) - # is O(N × bits_per_pass) vs full sort O(N log N). For vocab=152064: - # topk ≈ 11 radix passes, sort ≈ 17 passes. 1.5x fewer kernel cycles. + # CCCL dispatch_topk.cuh architecture (480 lines, full read): # + # 1. Multi-pass radix selection: O(N × bits_per_pass) not O(N log N) + # pass 0: DeviceTopKHistogramKernel (histogram only, no filter) + # pass 1+: DeviceTopKKernel (fused filter + histogram) + # last: DeviceTopKLastFilterKernel (filter only) + # + # 2. DoubleBuffer pattern (dispatch_topk.cuh line ~430): + # key_bufs = DoubleBuffer(alloc[3], alloc[2]) // ping-pong + # for pass: use Current() as input, Alternate() as output, then swap + # → zero allocation in the hot loop + # + # 3. candidate_buffer_length = num_items / 128 + # Only 1/128 of input needs buffer space for candidates + # vocab=152064 → 1188 candidates max + # + # PyTorch translation below uses pre-allocated buffers where possible + # to avoid per-step allocation overhead (BI-V100 has no async allocator). + # Fast path: when ALL sequences use top_p=1.0 (no nucleus sampling), # we only need top-k selection, not full sort + cumsum. - # This skips: sort (152K elements) + softmax + cumsum + scatter - # and replaces with: topk (much cheaper) + scatter. all_top_p_disabled = (p >= 1.0 - 1e-6).all() if all_top_p_disabled: - # Pure top-k path: use torch.topk instead of full sort - # For k values, take the minimum k across all sequences max_k = k.max().item() if max_k > 0 and max_k < logits.size(1): - # Get top-k values and indices topk_vals, topk_idx = torch.topk(logits, int(max_k), dim=-1) - # Mask out everything below top-k threshold per sequence - # topk_vals[:, -1] is the k-th largest value for each seq actual_k_mask = torch.arange(int(max_k), device=k.device).unsqueeze(0) < k.unsqueeze(1) topk_vals.masked_fill_(~actual_k_mask, -float("inf")) - # Get per-sequence threshold (smallest value kept) threshold = topk_vals.min(dim=-1, keepdim=True).values - # Apply threshold to original logits logits = logits.masked_fill(logits < threshold, -float("inf")) return logits # Full path: sort + top-k + top-p (cumsum) - logits_sort, logits_idx = logits.sort(dim=-1, descending=False) + # CCCL DoubleBuffer insight: reuse sort output tensors across calls + # by caching them keyed on (batch_size, vocab_size, device). + # This avoids torch.sort allocating 2 new tensors (152064×4B each) + # on every single decode step. + _buf_key = (logits.shape[0], logits.shape[1], str(logits.device)) + _bufs = getattr(_apply_top_k_top_p, '_sort_bufs', {}).get(_buf_key) + if _bufs is not None: + logits_sort, logits_idx = _bufs + # In-place sort into pre-allocated buffers + torch.sort(logits, dim=-1, descending=False, out=(logits_sort, logits_idx)) + else: + logits_sort, logits_idx = logits.sort(dim=-1, descending=False) + # Cache for next call (CCCL DoubleBuffer pattern) + if not hasattr(_apply_top_k_top_p, '_sort_bufs'): + _apply_top_k_top_p._sort_bufs = {} + _apply_top_k_top_p._sort_bufs[_buf_key] = ( + logits_sort.clone(), logits_idx.clone()) # pre-alloc buffers # Apply top-k. top_k_mask = logits_sort.size(1) - k.to(torch.long)