Sampler:從 logits 到 token 的最後一公里
職責
模型前向算完,GPUModelRunner 拿到 [num_tokens, vocab_size] 的 logits。剩下的事就是把 logits 變成 token id:應用 temperature、top_k、top_p、penalties,做 greedy 或 multinomial 採樣,再順手把 top-k logprobs 收集好。Sampler(sampler.py:20)就是幹這件事的 nn.Module,每個 GPUModelRunner 持有一個實例(gpu_model_runner.py:532)。
Sampler.forward 是一個固定九步流水線(sampler.py:22-59):先按需算 raw_logprobs / clone raw_logits,把 logits 轉 float32,應用 allowed_token_ids 白名單、bad_words 黑名單,跑 non-argmax-invariant 的 logits processor(MinTokens、LogitBias),應用 repetition / frequency / presence penalty,然後進 sample() 裡做 greedy 或 random 採樣。sample(sampler.py:243)內部再分:greedy 走 argmax,random 走 apply_temperature → argmax-invariant processor(預設 MinP)→ topk_topp_sampler → torch.where(temp < eps, greedy, random)。
具體的 top-k/top-p 實現不在 Sampler 裡,而是 TopKTopPSampler(topk_topp_sampler.py:70),它在 CUDA 上走 FlashInfer 的 top_p_sampling_from_probs / top_k_sampling_from_probs(topk_topp_sampler.py:471),在 CPU/XPU/ROCm 走各自的 native 實現,都不是的話退到 forward_native。
設計動機
- argmax 不變性分類:
LogitsProcessor.is_argmax_invariant(interface.py:84-92)區分會不會改 greedy 結果,MinP、MinTokens、LogitBias 分別落在兩邊;greedy 請求跳過 argmax-invariant 後面的整套 random 流程(sampler.py:257-271)。 - 溫度為 0 視為 greedy:
apply_temperature把temp < eps的位置改成 1.0(sampler.py:233-237),最後torch.where(temp < eps, greedy, random)一次合併,避免分支。 - in-place 優先:
apply_temperature用logits.div_、MinP.apply用masked_fill_(builtin.py:115),降低顯存峰值。 - FlashInfer sampler 快但功能受限:它不返回過濾後的 logits(
topk_topp_sampler.py:86-95),所以 logprobs_mode 為processed_logits/processed_logprobs時強制走forward_native。 - logprobs 收集和採樣分開:
gather_logprobs(sampler.py:308-356)在採樣之後跑,拿原始(未 temperature scale)logits 算 log_softmax,避免 penalty 影響返回給使用者的 logprob 值。 - 批維度動態特化防護:
gather_logprobs用torch._dynamo.decorators.mark_unbacked(sampler.py:345-346),讓 dynamo 不在 batch_size=1 → ≥2 時重編譯。 - spec decode 兼容:
apply_logits_processors在predict_bonus_token時會把spec_token_ids拼到output_token_ids後(sampler.py:384-393),讓 penalty 看到 draft token。
關鍵檔案
Sampler 类:20—nn.Module,持TopKTopPSampler,實現 9 步流水線。forward docstring:22-59— 9 步流水線規範,標註每步要做什麼。Sampler.forward:72-149— 主入口,串起 apply_logits_processors + sample + gather_logprobs。Sampler.sample:243— greedy / random 分流,temperature 應用 + top-k/top-p。apply_temperature:227-237— in-placediv_,把 0 溫度改 1.0 防除零。greedy_sample:239-241—logits.argmax(dim=-1).view(-1)。apply_logits_processors:371-420— allowed_token_ids_mask + bad_words + non_argmax_invariant processors + penalties。apply_penalties:422-439— 委託給apply_all_penaltiesop。SamplingMetadata:14— 持 temperature/top_p/top_k/generators/penalties 等每批採樣元資料。TopKTopPSampler:70— 平臺分發:CUDA→FlashInfer、CPU→native、XPU→kernel、ROCm→aiter。flashinfer_sample:471-508— 調 FlashInfertop_p_renorm_probs/top_k_sampling_from_probs,統計等價但比 rejection 快。GPUModelRunner 持有 Sampler:532— 在構造時Sampler(logprobs_mode, use_fp64_gumbel)。
資料流
每一步的輸入都是 [num_tokens, vocab_size] 的 logits,num_tokens 是當前 step 整批請求要算的 token 總數(prefill 算多個 prompt token,decode 每請求算 1 個)。下面是 sample 方法裡 greedy 和 random 的分流:
# vllm/v1/sample/sampler.py L243-L302
def sample(
self,
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
logprobs_mode_override: LogprobsMode | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
logprobs_mode = logprobs_mode_override or self.logprobs_mode
assert not (sampling_metadata.all_greedy and sampling_metadata.all_random)
if sampling_metadata.all_random:
greedy_sampled = None
else:
greedy_sampled = self.greedy_sample(logits)
if sampling_metadata.all_greedy:
processed_logprobs = None
if (
sampling_metadata.max_num_logprobs is not None
or sampling_metadata.logprob_token_ids
):
if logprobs_mode == "processed_logits":
processed_logprobs = logits
elif logprobs_mode == "processed_logprobs":
processed_logprobs = self.compute_logprobs(logits)
return greedy_sampled, processed_logprobs
assert sampling_metadata.temperature is not None
# Apply temperature.
logits = self.apply_temperature(
logits, sampling_metadata.temperature, sampling_metadata.all_random
)
# Apply logits processors that only apply to random sampling
# (argmax invariant)
for processor in sampling_metadata.logitsprocs.argmax_invariant:
logits = processor.apply(logits)
# Apply top_k and/or top_p.
random_sampled, processed_logprobs = self.topk_topp_sampler(
logits,
sampling_metadata.generators,
sampling_metadata.top_k,
sampling_metadata.top_p,
)
if greedy_sampled is None:
return random_sampled, processed_logprobs
sampled = torch.where(
sampling_metadata.temperature < _SAMPLING_EPS,
greedy_sampled,
random_sampled,
out=greedy_sampled, # Reuse tensor
)
return sampled, processed_logprobsall_greedy / all_random 是 SamplingMetadata 裡預算好的批次標誌,讓整批直接走 fast path。混合 batch 時兩條路徑都算一遍,最後用 torch.where 按 per-request temperature 選擇——這樣 GPU 不會因分支拖慢。
邊界與失敗
all_greedy和all_random互斥:sample開頭assert not (all_greedy and all_random)(sampler.py:256),配置同時開兩個會直接 AssertionError。- 溫度為 0 的 token 不走 random:
torch.where(temp < eps, greedy, random)(sampler.py:296-301)保證 greedy 請求永遠拿 argmax,不被 random 汙染。 - logprobs 用原始 logits:
compute_logprobs在apply_temperature之前算(sampler.py:85-93),penalty / temperature 不會影響返回給使用者的 logprob 數值,這是和 v0 的差異。 logprob_token_ids走專用 gather:gather_specific_token_logprobs(sampler.py:151-225)在generative_scoringAPI 用,跳過 top-k 直接 gather 指定 token,帶 padding 和 valid_mask。- spec decode 時禁用 logits processor:
build_logitsprocs檢測到speculative_config且傳了自定義 processor 就拋ValueError(__init__.py:201-209),只留 MinTokens。 - FlashInfer 輸出 int32:
forward末尾sampled = sampled.long()(sampler.py:109)統一到 int64,因為 FlashInfer sampler 返回 int32 而 PyTorch argmax / topk 返回 int64,後續 index 操作要兼容。
小結
Sampler 是 vLLM 採樣的統一入口:九步流水線 + greedy/random 分流 + FlashInfer 加速。它依賴 /sampling/logits 這層 LogitsProcessor 抽象做參數化變換,本身只關心「怎麼把 logits 變 token」這一件事。返回的 SamplerOutput 再往上傳到 /worker/gpu-model-runner,最終經 OutputProcessor 拼回 RequestOutput。引擎核心怎麼排程這些 step,見 /engine/engine-core。