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