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。