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,不走這層抽象。