Skip to content

Sampler:從 logits 到 token 的最後一公里

源码版本v0.25.1

職責

模型前向算完,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_samplertorch.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_temperaturetemp < eps 的位置改成 1.0(sampler.py:233-237),最後 torch.where(temp < eps, greedy, random) 一次合併,避免分支。
  • in-place 優先:apply_temperaturelogits.div_MinP.applymasked_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_logprobstorch._dynamo.decorators.mark_unbacked(sampler.py:345-346),讓 dynamo 不在 batch_size=1 → ≥2 時重編譯。
  • spec decode 兼容:apply_logits_processorspredict_bonus_token 時會把 spec_token_ids 拼到 output_token_ids 後(sampler.py:384-393),讓 penalty 看到 draft token。

關鍵檔案

資料流

每一步的輸入都是 [num_tokens, vocab_size] 的 logits,num_tokens 是當前 step 整批請求要算的 token 總數(prefill 算多個 prompt token,decode 每請求算 1 個)。下面是 sample 方法裡 greedy 和 random 的分流:

python
# 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_logprobs

all_greedy / all_randomSamplingMetadata 裡預算好的批次標誌,讓整批直接走 fast path。混合 batch 時兩條路徑都算一遍,最後用 torch.where 按 per-request temperature 選擇——這樣 GPU 不會因分支拖慢。

邊界與失敗

  • all_greedyall_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_logprobsapply_temperature 之前算(sampler.py:85-93),penalty / temperature 不會影響返回給使用者的 logprob 數值,這是和 v0 的差異。
  • logprob_token_ids 走專用 gather:gather_specific_token_logprobs(sampler.py:151-225)在 generative_scoring API 用,跳過 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

對照官方資料:vLLM 文件 · README