Skip to content

LogitsProcessor:采样前的 logits 变换

源码版本v0.25.1

职责

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) -> logitsis_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 / LogitBiasFalse(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 双张量:MinPLogitsProcessormin_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.applylogits[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_logitsprocsspeculative_config 非空时只保留 MinTokensLogitsProcessor(__init__.py:201-209),MinP / LogitBias 都被关掉,因为它们和 rejection sampling 语义不兼容。

关键文件

数据流

每个 forward 之前,SchedulerGPUModelRunner 会构造一个 BatchUpdate,告诉 processor「这一批新增了哪些请求、移除了哪些、移动了哪些」。processor 借这个机会刷新自己的 per-request 张量。MinP 的 update_state 是典型例子:

python
# 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 本体也很短:

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

Sampler.apply_logits_processors 把 non-argmax-invariant 和 argmax-invariant 分两段调用:前者在 greedy 路径也跑,后者只在 random 路径跑(sampler.py:403-405)。penalty 走独立的 apply_all_penalties op,不走 LogitsProcessor 这层。

边界与失败

  • Pooling model 不支持 processor:build_logitsprocsis_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_logitsprocsis_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,不需要每次更新。

小结

LogitsProcessorSampler 之上、sampling 之前的可插拔变换层:argmax-invariant 决定它跑在 greedy 还是 random 路径,update_state 让它跟踪 batch 增量变化。内置 MinP / LogitBias / MinTokens 三件套,加上插件和 FQCN 字符串支持外部扩展。具体采样过程见 /sampling/sampler,penalty 和 bad_words 是在 Sampler 里直接调 apply_all_penalties / apply_bad_words,不走这层抽象。

对照官方资料:vLLM 文档 · README