LogitsProcessor:采样前的 logits 变换
职责
Sampler 拿到原始 logits 之后,还要做一堆变换才能开始采样:MinP 过滤低概率 token、LogitBias 给特定 token 加偏置、MinTokens 在生成够数前禁止 EOS、repetition/frequency/presence penalty 惩罚重复、structured output 把不在语法内的 token 屏蔽掉、thinking budget 给 reasoning 模型加 budget。这些变换的统一抽象就是 LogitsProcessor(interface.py:60),每个 processor 持自己的 batch 状态,在每个 forward 之前被 Sampler.apply_logits_processors(sampler.py:371)依次调用。
LogitsProcessor 接口很轻,就四个方法:__init__、apply(logits) -> logits、is_argmax_invariant() -> bool(interface.py:84)、update_state(batch_update)(interface.py:94)。is_argmax_invariant 决定它会被分到 LogitsProcessors.argmax_invariant 还是 non_argmax_invariant(state.py:148-160),前者只在 random 采样时跑(MinP),后者在贪心采样也跑(MinTokens、LogitBias)。update_state 在每次 batch 成分变化时被调用,processor 借机更新自己持有的 per-request 状态张量。
内置 processor 在 BUILTIN_LOGITS_PROCESSORS(__init__.py:49-53),按顺序是 [MinTokensLogitsProcessor, LogitBiasLogitsProcessor, MinPLogitsProcessor]。除此之外还支持 entry-point 插件(vllm.logits_processors 组)和 FQCN 字符串(__init__.py:86-155),build_logitsprocs(__init__.py:184)把这三类合到一起实例化。
设计动机
- argmax 不变性分类:
MinPLogitsProcessor.is_argmax_invariant() = True(builtin.py:47-49),因为它只过滤概率极低的 token,不影响 argmax;MinTokens/LogitBias是False(builtin.py:183-186),因为它们能改 greedy 的结果(屏蔽 EOS、加偏置后 argmax 可能变)。 - 批量状态增量更新:
update_state接收BatchUpdate(removed/added/moved)(interface.py:36-57),processor 只在 batch 成分真的变了时才重算自己的张量(MinP 用min_p_count跳过空批次)(builtin.py:54-100)。 - CPU↔GPU 双张量:
MinPLogitsProcessor持min_p_cpu(pinned)+min_p_device(builtin.py:30-44),update 时在 CPU 改值,apply之前 H2D 拷一次,避免每次apply都跨设备同步。 - MinP 走 softmax 而不是 logits:
apply先算softmax(logits)找 max prob,再按max_prob * min_p阈值过滤(builtin.py:102-116),因为 min_p 是概率空间的阈值,不是 logit 空间。 - LogitBias 用稀疏索引:
LogitBiasLogitsProcessor.apply用logits[req_idx, tok_id] += bias(builtin.py:159-162),只在有 bias 的位置改,不对整个 vocab 做加法。 - MinTokens 屏蔽 EOS:
MinTokensLogitsProcessor在生成数不到min_tokens时把stop_token_ids屏蔽(builtin.py:165-186),让模型必须生成够长。 - Spec decode 禁用自定义 processor:
build_logitsprocs在speculative_config非空时只保留MinTokensLogitsProcessor(__init__.py:201-209),MinP / LogitBias 都被关掉,因为它们和 rejection sampling 语义不兼容。
关键文件
LogitsProcessor ABC:60— 四方法抽象基类:__init__、apply、is_argmax_invariant、update_state。BatchUpdate:36— frozen dataclass,持 removed/added/moved 三个序列。LogitsProcessors:148— 装载容器,按 argmax_invariant 分两路。LogitsProcessors.all:162-165— 把两路串成一个 iterator。MinPLogitsProcessor:23— 概率空间阈值过滤,argmax-invariant。LogitBiasLogitsProcessor:119— 稀疏偏置加法,non-argmax-invariant。MinTokensLogitsProcessor:165— 屏蔽 stop tokens 直到生成够数,non-argmax-invariant。BUILTIN_LOGITS_PROCESSORS:49-53— 默认三件套顺序:MinTokens → LogitBias → MinP。build_logitsprocs:184— 工厂入口,处理 pooling / spec-decode / 自定义三类短路。_load_logitsprocs_by_fqcns:86-155—<module>:<Qualname>字符串懒加载 +issubclass校验。ThinkingBudgetStateHolder:33— reasoning budget processor,apply_to_logits在 forward 时被 Sampler 调用。
数据流
每个 forward 之前,Scheduler 或 GPUModelRunner 会构造一个 BatchUpdate,告诉 processor「这一批新增了哪些请求、移除了哪些、移动了哪些」。processor 借这个机会刷新自己的 per-request 张量。MinP 的 update_state 是典型例子:
# vllm/v1/sample/logits_processor/builtin.py L54-L100
def update_state(self, batch_update: BatchUpdate | None):
if not batch_update:
return
needs_update = False
# Process added requests.
for index, params, _, _ in batch_update.added:
min_p = params.min_p
min_p_before = self.min_p_cpu[index]
if min_p_before != min_p:
needs_update = True
self.min_p_cpu[index] = min_p
if min_p and not min_p_before:
self.min_p_count += 1
elif not min_p and min_p_before:
self.min_p_count -= 1
if self.min_p_count:
# Process removed requests.
if batch_update.removed:
needs_update = True
for index in batch_update.removed:
if self.min_p_cpu[index]:
self.min_p_cpu[index] = 0
self.min_p_count -= 1
# Process moved requests, unidirectional (a->b) and swap (a<->b).
for adx, bdx, direct in batch_update.moved:
min_p_a, min_p_b = self.min_p_cpu[adx], self.min_p_cpu[bdx]
if min_p_a != min_p_b:
needs_update = True
self.min_p_cpu[bdx] = min_p_a
if direct == MoveDirectionality.SWAP:
self.min_p_cpu[adx] = min_p_b
if direct == MoveDirectionality.UNIDIRECTIONAL:
if min_p_a:
self.min_p_cpu[adx] = 0
if min_p_b:
self.min_p_count -= 1
# Update tensors if needed.
size = batch_update.batch_size
if self.min_p_count and (needs_update or self.min_p.shape[0] != size):
self.min_p = self.min_p_device[:size]
if self.use_double_tensor:
self.min_p.copy_(self.min_p_cpu_tensor[:size], non_blocking=True)
self.min_p.unsqueeze_(1)min_p_count 是个逃逸计数:只要全批没一个请求开了 min_p,apply 直接 return 不做任何运算。apply 本体也很短:
# vllm/v1/sample/logits_processor/builtin.py L102-L116
def apply(self, logits: torch.Tensor) -> torch.Tensor:
if not self.min_p_count:
return logits
# Convert logits to probability distribution
probability_values = torch.nn.functional.softmax(logits, dim=-1)
# Calculate maximum probabilities per sequence
max_probabilities = torch.amax(probability_values, dim=-1, keepdim=True)
# Adjust min_p
adjusted_min_p = max_probabilities.mul_(self.min_p)
# Identify valid tokens using threshold comparison
invalid_token_mask = probability_values < adjusted_min_p
# Apply mask using boolean indexing
logits.masked_fill_(invalid_token_mask, -float("inf"))
return logitsSampler.apply_logits_processors 把 non-argmax-invariant 和 argmax-invariant 分两段调用:前者在 greedy 路径也跑,后者只在 random 路径跑(sampler.py:403-405)。penalty 走独立的 apply_all_penalties op,不走 LogitsProcessor 这层。
边界与失败
- Pooling model 不支持 processor:
build_logitsprocs在is_pooling_model=True时直接返回空LogitsProcessors,传了自定义就抛STR_POOLING_REJECTS_LOGITSPROCS(__init__.py:191-198)。 - Spec decode 禁用自定义 processor:
STR_SPEC_DEC_REJECTS_LOGITSPROCS(__init__.py:201-209),只留MinTokensLogitsProcessor,因为 MinP / LogitBias 和 rejection sampling 不兼容。 - TPU 不支持自定义 processor:
_load_custom_logitsprocs在is_tpu()时直接返回空列表(__init__.py:176-179),v1 在 TPU 上还没接自定义 processor。 - 插件加载失败硬抛:
_load_logitsprocs_plugins在某个 entry point 加载失败时直接raise RuntimeError(__init__.py:78-82),不会静默跳过,避免运行时才发现某 processor 没生效。 - FQCN 必须是 LogitsProcessor 子类:
_load_logitsprocs_by_fqcns走<module>:<Qualname>切分(__init__.py:128),不是子类就抛ValueError,防止误传普通函数。 - batch_update 引用语义:
BatchUpdate.added里每个 tuple 的output_tok_ids是对 request 的 running tokens 列表的引用(interface.py:44-50),processor 通过这个引用看到最新生成 token,不需要每次更新。
小结
LogitsProcessor 是 Sampler 之上、sampling 之前的可插拔变换层:argmax-invariant 决定它跑在 greedy 还是 random 路径,update_state 让它跟踪 batch 增量变化。内置 MinP / LogitBias / MinTokens 三件套,加上插件和 FQCN 字符串支持外部扩展。具体采样过程见 /sampling/sampler,penalty 和 bad_words 是在 Sampler 里直接调 apply_all_penalties / apply_bad_words,不走这层抽象。