From 1064ce756b56150afe182a89edcac93dcefa1d7a Mon Sep 17 00:00:00 2001 From: muh-bot Date: Thu, 6 Aug 2026 04:13:03 +0000 Subject: [PATCH] [base/sampler] CCCL dispatch_topk alias_temporaries: fix _sampler_cache bug + pre-allocate temp storage MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Source: CCCL dispatch_topk.cuh alias_temporaries() pattern - Pre-allocate counter + histogram + double-buffer into single blob - No per-kernel-launch malloc in the hot path - BI-V100 16 SMs: every unnecessary CUDA malloc stalls all SMs Changes to vllm/model_executor/layers/sampler.py: 1. Fix _sampler_cache global declaration bug: - Old: 'if "_sampler_cache" not in dir()' — dir() returns local scope names in function context, not globals. The cache was being recreated on every call, defeating the purpose of caching entirely. - New: module-level _sampler_temp_storage dict, declared once at import. 2. Apply CCCL alias_temporaries pattern: - _sampler_temp_storage is a module-level dict that maps (shape_key -> pre-allocated CUDA tensor). - bin_counts tensor (vocab=152064, int64) = 1.2MB per sequence, allocated ONCE and .zero_() reused on each decode step. - Eliminates cudaMalloc/cudaFree cycle per decode step in _apply_penalties -> _get_bin_counts_and_mask path. CCCL reference read: cccl_upstream/cub/cub/device/dispatch/dispatch_topk.cuh - 460 lines, multi-pass radix select with DoubleBuffer - alias_temporaries packs 6 allocations into 1 cudaMalloc - Grid sizing: min(MaxSmOccupancy * num_sms, num_tiles) - Key insight: BI-V100 with 16 SMs has very small grids, so per-launch overhead (malloc, memset) dominates more than on 148-SM GPUs where kernel compute time dominates --- vllm/model_executor/layers/sampler.py | 33 +++++++++++++++------------ 1 file changed, 18 insertions(+), 15 deletions(-) diff --git a/vllm/model_executor/layers/sampler.py b/vllm/model_executor/layers/sampler.py index 3129da68..95fa3b72 100644 --- a/vllm/model_executor/layers/sampler.py +++ b/vllm/model_executor/layers/sampler.py @@ -30,6 +30,16 @@ if envs.VLLM_USE_FLASHINFER_SAMPLER and find_spec("flashinfer"): else: flashinfer_top_k_top_p_sampling = None +# CCCL alias_temporaries pattern (dispatch_topk.cuh line ~340): +# All temporary buffers (counter, histogram, candidate double-buffer) are +# pre-allocated into a single d_temp_storage blob at launch time, then +# aliased via pointer arithmetic. No per-kernel-launch malloc. +# +# Python equivalent: module-level dict mapping (shape_key → pre-allocated tensor). +# Survives across decode steps within the same process lifetime. +# Keys: ("bin_counts", vocab_size, num_seqs, device) etc. +_sampler_temp_storage: Dict = {} + # (num_token_ids, num_parent_ids) per sequence group. SampleResultType = List[Tuple[List[int], List[int]]] @@ -332,23 +342,16 @@ def _get_bin_counts_and_mask( # Compute the bin counts for the tokens. # vocab_size + 1 for padding. # - # CCCL bit_packed_counter pattern (catch2_test_memcpy_bitpacked_counter.cu): - # Pack counters using minimum bits needed. Original code uses int64 - # (8 bytes per counter), but token repetition counts in a single - # generation never exceed a few hundred. We keep int64 for scatter_add_ - # compatibility but pre-allocate once to avoid per-step CUDA malloc. + # CCCL alias_temporaries pattern (dispatch_reduce.cuh, dispatch_topk.cuh): + # Pre-allocate all temporary buffers once, reuse across kernel launches. + # In dispatch_topk.cuh this is done via alias_temporaries() which packs + # counter + histogram + double-buffer into a single d_temp_storage blob. # - # CCCL dispatch_reduce.cuh alias_temporaries: pre-allocate, reuse. - # For Qwen3.6 (vocab=152064, batch=8 decode): - # bin_counts = 8 × 152065 × 8 = 9.7 MB, allocated ONCE, reused. + # For Qwen3.6 (vocab=152064, max_num_seqs=1 decode): + # bin_counts = 1 × 152065 × 8 = 1.2 MB, allocated ONCE, reused. # scatter_add_ requires int64 on CUDA, so dtype cannot change. - # - # Future: if scatter_add_ supports int16/int32, switch to reduce 4x. _cache_key = ("bin_counts", vocab_size, num_seqs, tokens.device) - global _sampler_cache - if '_sampler_cache' not in dir(): - _sampler_cache = {} - cached = _sampler_cache.get(_cache_key) + cached = _sampler_temp_storage.get(_cache_key) if cached is not None and cached.shape == (num_seqs, vocab_size + 1): bin_counts = cached bin_counts.zero_() @@ -356,7 +359,7 @@ def _get_bin_counts_and_mask( bin_counts = torch.zeros((num_seqs, vocab_size + 1), dtype=torch.long, device=tokens.device) - _sampler_cache[_cache_key] = bin_counts + _sampler_temp_storage[_cache_key] = bin_counts bin_counts.scatter_add_(1, tokens, torch.ones_like(tokens)) bin_counts = bin_counts[:, :vocab_size]