Skip to content

Sampler:logits から token への最後の 1 マイル

源码版本v0.25.1

役割

モデルの forward が終わると、GPUModelRunner[num_tokens, vocab_size] の logits を受け取ります。残りの仕事は logits を token id に変えることです。temperature、top_k、top_p、penalties を適用し、greedy または multinomial でサンプリングし、ついでに top-k logprobs も集めます。Sampler(sampler.py:20)はこの一連を行う nn.Module で、GPUModelRunner はインスタンスを 1 つ持ちます(gpu_model_runner.py:532)。

Sampler.forward は固定 9 ステップのパイプラインです(sampler.py:22-59):必要に応じて raw_logprobs を算出 / raw_logits を clone し、logits を float32 に変換し、allowed_token_ids のホワイトリスト、bad_words のブラックリストを適用し、non-argmax-invariant な logits processor(MinTokens、LogitBias)を走らせ、repetition / frequency / presence penalty を適用してから sample() で greedy または random サンプリングを行います。sample(sampler.py:243)は内部でさらに分岐し、greedy は argmax、random は apply_temperature → argmax-invariant processor(デフォルトは MinP)→ topk_topp_samplertorch.where(temp < eps, greedy, random) と進みます。

具体的な top-k/top-p の実装は Sampler ではなく TopKTopPSampler(topk_topp_sampler.py:70)にあります。CUDA では FlashInfer の top_p_sampling_from_probs / top_k_sampling_from_probs(topk_topp_sampler.py:471)を呼び、CPU/XPU/ROCm ではそれぞれのネイティブ実装を使い、いずれでもなければ forward_native にフォールバックします。

設計動機

  • argmax 不変性の分類:LogitsProcessor.is_argmax_invariant(interface.py:84-92)は greedy の結果を変えるかどうかを区別し、MinP、MinTokens、LogitBias はそれぞれどちらかに分類されます。greedy リクエストは argmax-invariant 以降の random パイプライン全体をスキップします(sampler.py:257-271)。
  • 温度 0 を greedy として扱う:apply_temperaturetemp < eps の位置を 1.0 に書き換え(sampler.py:233-237)、最後に torch.where(temp < eps, greedy, random) で一度にマージして分岐を避けます。
  • in-place 優先:apply_temperaturelogits.div_MinP.applymasked_fill_(builtin.py:115)を使い、VRAM のピークを下げます。
  • FlashInfer sampler は速いが機能制限あり:フィルタ後の logits を返しません(topk_topp_sampler.py:86-95)。そのため logprobs_mode が processed_logits / processed_logprobs のときは forward_native に強制します。
  • logprobs 収集とサンプリングの分離:gather_logprobs(sampler.py:308-356)はサンプリング後に走り、元の(temperature scale されていない)logits で log_softmax を計算します。penalty がユーザーへ返す logprob 値に影響しないようにするためです。
  • バッチ次元の動的特化を防ぐ:gather_logprobstorch._dynamo.decorators.mark_unbacked(sampler.py:345-346)を使い、batch_size=1 → ≥2 の再コンパイルを抑制します。
  • spec decode 互換:apply_logits_processorspredict_bonus_token 時に spec_token_idsoutput_token_ids の後ろに結合し(sampler.py:384-393)、penalty が draft token を見られるようにします。

主要ファイル

データフロー

各ステップの入力はどれも [num_tokens, vocab_size] の logits です。num_tokens は現在のステップでバッチ全体が計算すべき token 総数です(prefill では複数 prompt token を計算、decode ではリクエスト毎に 1 つ)。以下は sample メソッド内の greedy / random の分岐です:

python
# vllm/v1/sample/sampler.py L243-L302
def sample(
    self,
    logits: torch.Tensor,
    sampling_metadata: SamplingMetadata,
    logprobs_mode_override: LogprobsMode | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    logprobs_mode = logprobs_mode_override or self.logprobs_mode
    assert not (sampling_metadata.all_greedy and sampling_metadata.all_random)
    if sampling_metadata.all_random:
        greedy_sampled = None
    else:
        greedy_sampled = self.greedy_sample(logits)
        if sampling_metadata.all_greedy:
            processed_logprobs = None
            if (
                sampling_metadata.max_num_logprobs is not None
                or sampling_metadata.logprob_token_ids
            ):
                if logprobs_mode == "processed_logits":
                    processed_logprobs = logits
                elif logprobs_mode == "processed_logprobs":
                    processed_logprobs = self.compute_logprobs(logits)
            return greedy_sampled, processed_logprobs

    assert sampling_metadata.temperature is not None

    # Apply temperature.
    logits = self.apply_temperature(
        logits, sampling_metadata.temperature, sampling_metadata.all_random
    )

    # Apply logits processors that only apply to random sampling
    # (argmax invariant)
    for processor in sampling_metadata.logitsprocs.argmax_invariant:
        logits = processor.apply(logits)

    # Apply top_k and/or top_p.
    random_sampled, processed_logprobs = self.topk_topp_sampler(
        logits,
        sampling_metadata.generators,
        sampling_metadata.top_k,
        sampling_metadata.top_p,
    )

    if greedy_sampled is None:
        return random_sampled, processed_logprobs

    sampled = torch.where(
        sampling_metadata.temperature < _SAMPLING_EPS,
        greedy_sampled,
        random_sampled,
        out=greedy_sampled,  # Reuse tensor
    )
    return sampled, processed_logprobs

all_greedy / all_randomSamplingMetadata で事前計算されたバッチフラグで、バッチ全体がそのまま fast path に入れます。混合バッチでは両方のパスを計算し、最後に torch.where でリクエスト毎の temperature に従って選びます——これにより GPU が分岐で遅くなるのを防ぎます。

境界と失敗

  • all_greedyall_random は排他:sample の冒頭で assert not (all_greedy and all_random)(sampler.py:256)があり、両方を同時に有効にすると AssertionError になります。
  • 温度 0 の token は random を通らない:torch.where(temp < eps, greedy, random)(sampler.py:296-301)により、greedy リクエストは常に argmax を取得し、random に汚染されません。
  • logprobs は元の logits を使用:compute_logprobsapply_temperature の前に計算され(sampler.py:85-93)、penalty / temperature がユーザーへ返す logprob の数値に影響しません。これは v0 との差異です。
  • logprob_token_ids は専用 gather:gather_specific_token_logprobs(sampler.py:151-225)は generative_scoring API で使い、top-k を飛ばして指定 token を直接 gather し、padding と valid_mask を付けます。
  • spec decode 時は logits processor を無効化:build_logitsprocsspeculative_config を検出してカスタム processor が渡されると ValueError を投げ(__init__.py:201-209)、MinTokens のみ残します。
  • FlashInfer は int32 を返す:forward の末尾で sampled = sampled.long()(sampler.py:109)で int64 に揃えます。FlashInfer sampler は int32、PyTorch argmax / topk は int64 を返し、後続の index 操作との互換性のためです。

まとめ

Sampler は vLLM のサンプリングの統一入口です:9 ステップパイプライン + greedy/random 分岐 + FlashInfer 加速。パラメータ化変換は /sampling/logitsLogitsProcessor 抽象に任せ、本体は「logits を token に変える」ことだけを担います。返した SamplerOutput は上流の /worker/gpu-model-runner へ渡り、最終的に OutputProcessorRequestOutput に組み直します。エンジンコアがこれらのステップをどうスケジュールするかは /engine/engine-core を参照してください。

公式資料: vLLM 文档 · README