Sampler:logits から token への最後の 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_sampler → torch.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_temperatureはtemp < epsの位置を 1.0 に書き換え(sampler.py:233-237)、最後にtorch.where(temp < eps, greedy, random)で一度にマージして分岐を避けます。 - in-place 優先:
apply_temperatureはlogits.div_、MinP.applyはmasked_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_logprobsはtorch._dynamo.decorators.mark_unbacked(sampler.py:345-346)を使い、batch_size=1 → ≥2 の再コンパイルを抑制します。 - spec decode 互換:
apply_logits_processorsはpredict_bonus_token時にspec_token_idsをoutput_token_idsの後ろに結合し(sampler.py:384-393)、penalty が draft token を見られるようにします。
主要ファイル
Sampler クラス:20—nn.Module、TopKTopPSamplerを持ち、9 ステップパイプラインを実装します。forward docstring:22-59— 9 ステップパイプラインの仕様、各ステップで何をすべきかを示します。Sampler.forward:72-149— 主入口、apply_logits_processors + sample + gather_logprobs を繋ぎます。Sampler.sample:243— greedy / random の分岐、temperature 適用 + top-k/top-p。apply_temperature:227-237— in-place なdiv_、0 温度を 1.0 にしてゼロ除算を防ぎます。greedy_sample:239-241—logits.argmax(dim=-1).view(-1)。apply_logits_processors:371-420— allowed_token_ids_mask + bad_words + non_argmax_invariant processors + penalties。apply_penalties:422-439—apply_all_penaltiesop に委譲します。SamplingMetadata:14— temperature/top_p/top_k/generators/penalties などバッチ毎のサンプリングメタデータを持ちます。TopKTopPSampler:70— プラットフォームディスパッチ:CUDA→FlashInfer、CPU→native、XPU→kernel、ROCm→aiter。flashinfer_sample:471-508— FlashInfer のtop_p_renorm_probs/top_k_sampling_from_probsを呼び、rejection より高速に統計等価な結果を得ます。GPUModelRunner が Sampler を保持:532— コンストラクタでSampler(logprobs_mode, use_fp64_gumbel)を呼びます。
データフロー
各ステップの入力はどれも [num_tokens, vocab_size] の logits です。num_tokens は現在のステップでバッチ全体が計算すべき token 総数です(prefill では複数 prompt token を計算、decode ではリクエスト毎に 1 つ)。以下は sample メソッド内の greedy / random の分岐です:
# 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_logprobsall_greedy / all_random は SamplingMetadata で事前計算されたバッチフラグで、バッチ全体がそのまま fast path に入れます。混合バッチでは両方のパスを計算し、最後に torch.where でリクエスト毎の temperature に従って選びます——これにより GPU が分岐で遅くなるのを防ぎます。
境界と失敗
all_greedyとall_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_logprobsはapply_temperatureの前に計算され(sampler.py:85-93)、penalty / temperature がユーザーへ返す logprob の数値に影響しません。これは v0 との差異です。 logprob_token_idsは専用 gather:gather_specific_token_logprobs(sampler.py:151-225)はgenerative_scoringAPI で使い、top-k を飛ばして指定 token を直接 gather し、padding と valid_mask を付けます。- spec decode 時は logits processor を無効化:
build_logitsprocsはspeculative_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/logits の LogitsProcessor 抽象に任せ、本体は「logits を token に変える」ことだけを担います。返した SamplerOutput は上流の /worker/gpu-model-runner へ渡り、最終的に OutputProcessor が RequestOutput に組み直します。エンジンコアがこれらのステップをどうスケジュールするかは /engine/engine-core を参照してください。