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 モデルに予算を与える。これらの変換を統一する抽象が LogitsProcessor(interface.py:60)です。各 processor は自身のバッチ状態を持ち、forward 毎に Sampler.apply_logits_processors(sampler.py:371)から順に呼ばれます。

LogitsProcessor のインターフェースは軽く、メソッドは 4 つだけ:__init__apply(logits) -> logitsis_argmax_invariant() -> bool(interface.py:84)、update_state(batch_update)(interface.py:94)。is_argmax_invariantLogitsProcessors.argmax_invariantnon_argmax_invariant のどちらに分類されるかを決め(state.py:148-160)、前者は random サンプリング時のみ走り(MinP)、後者は貪心サンプリングでも走ります(MinTokens、LogitBias)。update_state はバッチ構成が変化する毎に呼ばれ、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)がこの 3 種をまとめてインスタンス化します。

設計動機

  • argmax 不変性の分類:MinPLogitsProcessor.is_argmax_invariant() = True(builtin.py:47-49)は、確率が極端に低い token をフィルタするだけで argmax に影響しないためです。MinTokens / LogitBiasFalse(builtin.py:183-186)で、greedy の結果を変え得ます(EOS のマスク、バイアス追加で argmax が変わりうる)。
  • バッチ状態の差分更新:update_stateBatchUpdate(removed/added/moved)を受け取り(interface.py:36-57)、processor はバッチ構成が本当に変わった時だけ自身のテンソルを再計算します(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 は logits ではなく softmax を使う:apply は先に softmax(logits) を計算して max prob を探し、max_prob * min_p の閾値でフィルタします(builtin.py:102-116)。min_p は確率空間の閾値であり、logit 空間ではないためです。
  • LogitBias は sparse index を使用: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 の前に、Scheduler または GPUModelRunnerBatchUpdate を構築し、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 はエスケープカウンタです:バッチ全体で 1 件も 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 を 2 段に分けて呼びます:前者は greedy パスでも走り、後者は random パスのみで走ります(sampler.py:403-405)。penalty は独立した apply_all_penalties op を経由し、LogitsProcessor 層は通りません。

境界と失敗

  • Pooling model は processor 非サポート:build_logitsprocsis_pooling_model=True のとき空の LogitsProcessors を返し、カスタム processor が渡されると 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 の上、サンプリング前のプラグ可能な変換レイヤです:argmax-invariant が greedy か random のどちらのパスで走るかを決め、update_state がバッチの差分変化を追跡します。組み込みで MinP / LogitBias / MinTokens の 3 点セットを持ち、プラグインと FQCN 文字列で外部拡張をサポートします。具体的なサンプリング過程は /sampling/sampler を、penalty と bad_words は Sampler 内で直接 apply_all_penalties / apply_bad_words を呼び、この抽象層は通りません。

公式資料: vLLM 文档 · README