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