Sampler : la dernière étape, des logits aux tokens
Responsabilités
Une fois le forward du modèle terminé, GPUModelRunner dispose de logits [num_tokens, vocab_size]. Il reste à transformer ces logits en token ids : appliquer temperature, top_k, top_p, penalties, faire l'échantillonnage greedy ou multinomial, et au passage collecter les top-k logprobs. Sampler(sampler.py:20) est le nn.Module qui s'en charge ; chaque GPUModelRunner en détient une instance(gpu_model_runner.py:532).
Sampler.forward est un pipeline fixe de neuf étapes(sampler.py:22-59) : calculer au besoin raw_logprobs / cloner raw_logits, convertir les logits en float32, appliquer la liste blanche allowed_token_ids, la liste noire bad_words, exécuter les logits processors non argmax-invariant (MinTokens, LogitBias), appliquer les pénalités repetition / frequency / presence, puis entrer dans sample() pour l'échantillonnage greedy ou random. sample(sampler.py:243) se découpe à son tour : greedy via argmax, random via apply_temperature → processor argmax-invariant (MinP par défaut) → topk_topp_sampler → torch.where(temp < eps, greedy, random).
L'implémentation concrète de top-k/top-p ne vit pas dans Sampler mais dans TopKTopPSampler(topk_topp_sampler.py:70) : sur CUDA il passe par FlashInfer (top_p_sampling_from_probs / top_k_sampling_from_probs)(topk_topp_sampler.py:471), sur CPU/XPU/ROCm par leurs implémentations natives respectives, et sinon se replie sur forward_native.
Motivation de conception
- Classification par invariance argmax :
LogitsProcessor.is_argmax_invariant(interface.py:84-92) distingue ceux qui changent le résultat greedy. MinP, MinTokens et LogitBias tombent de chaque côté ; les requêtes greedy sautent toute la suite du chemin random après les processors argmax-invariant(sampler.py:257-271). - Température nulle vue comme greedy :
apply_temperaturemet les positionstemp < epsà 1.0(sampler.py:233-237), puis un uniquetorch.where(temp < eps, greedy, random)fusionne à la fin, pour éviter les branches. - Privilégier in-place :
apply_temperatureutiliselogits.div_, etMinP.applyutilisemasked_fill_(builtin.py:115) pour réduire le pic de mémoire GPU. - Le sampler FlashInfer est rapide mais limité : il ne retourne pas les logits filtrés(
topk_topp_sampler.py:86-95), donc quand logprobs_mode vautprocessed_logits/processed_logprobs, on forceforward_native. - Collecte des logprobs séparée de l'échantillonnage :
gather_logprobs(sampler.py:308-356) tourne après l'échantillonnage, et calculelog_softmaxsur les logits bruts (non mis à l'échelle par temperature), pour que les pénalités n'affectent pas les valeurs de logprob renvoyées à l'utilisateur. - Protection contre la spécialisation dynamique sur la dimension de batch :
gather_logprobsutilisetorch._dynamo.decorators.mark_unbacked(sampler.py:345-346) pour que dynamo ne recompile pas quand batch_size passe de 1 à ≥2. - Compatibilité spec decode :
apply_logits_processors, en modepredict_bonus_token, concatènespec_token_idsà la suite deoutput_token_ids(sampler.py:384-393) pour que la pénalité voie les draft tokens.
Fichiers clés
Sampler 类:20—nn.Module, détientTopKTopPSampler, implémente le pipeline en 9 étapes.forward docstring:22-59— spécification du pipeline 9 étapes, indique ce que fait chaque étape.Sampler.forward:72-149— entrée principale, enchaîne apply_logits_processors + sample + gather_logprobs.Sampler.sample:243— branchement greedy / random, application de temperature + top-k/top-p.apply_temperature:227-237—div_in-place, met la température 0 à 1.0 pour éviter la division par zéro.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— délègue à l'opérationapply_all_penalties.SamplingMetadata:14— porte temperature/top_p/top_k/generators/penalties, les métadonnées d'échantillonnage du batch.TopKTopPSampler:70— dispatch par plateforme : CUDA→FlashInfer, CPU→native, XPU→kernel, ROCm→aiter.flashinfer_sample:471-508— appelle FlashInfertop_p_renorm_probs/top_k_sampling_from_probs, statistiquement équivalent mais plus rapide que le rejection.GPUModelRunner 持有 Sampler:532— construitSampler(logprobs_mode, use_fp64_gumbel)à l'initialisation.
Flux de données
Chaque étape prend en entrée des logits [num_tokens, vocab_size], où num_tokens est le total des tokens à calculer pour le batch à l'étape courante (prefill : plusieurs prompt tokens ; decode : 1 par requête). Voici le branchement greedy / random dans la méthode sample :
# 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 sont des flags de batch précalculés dans SamplingMetadata qui permettent à tout le batch d'emprunter un fast path. En batch mixte, les deux chemins sont calculés puis fusionnés via torch.where en fonction de la température de chaque requête — pour éviter que le GPU ne ralentisse à cause de branches.
Limites et échecs
all_greedyetall_randommutuellement exclusifs : en tête desample,assert not (all_greedy and all_random)(sampler.py:256) ; configurer les deux lève une AssertionError.- Token à température 0 hors chemin random :
torch.where(temp < eps, greedy, random)(sampler.py:296-301) garantit que les requêtes greedy obtiennent toujours l'argmax, sans pollution par le chemin random. - logprobs calculés sur logits bruts :
compute_logprobss'exécute avantapply_temperature(sampler.py:85-93) ; penalty et temperature n'affectent pas les valeurs de logprob renvoyées à l'utilisateur, c'est une différence par rapport à v0. logprob_token_idspasse par un gather dédié :gather_specific_token_logprobs(sampler.py:151-225) est utilisé par l'APIgenerative_scoring, saute le top-k et gather directement les tokens spécifiés, avec padding et valid_mask.- Logits processors désactivés en spec decode :
build_logitsprocslèveValueErrorsispeculative_configest détecté et qu'un processor personnalisé est passé(__init__.py:201-209) ; seul MinTokens est conservé. - FlashInfer retourne du int32 : à la fin de
forward,sampled = sampled.long()(sampler.py:109) unifie en int64, car le sampler FlashInfer renvoie du int32 alors que argmax / topk de PyTorch renvoient du int64 ; les opérations d'indexation suivantes doivent rester compatibles.
Résumé
Sampler est l'unique point d'entrée de l'échantillonnage vLLM : pipeline en neuf étapes + branchement greedy/random + accélération FlashInfer. Il s'appuie sur la couche LogitsProcessor décrite dans /sampling/logits pour les transformations paramétrées, et ne s'occupe lui-même que de « comment transformer les logits en tokens ». Le SamplerOutput remonte vers /worker/gpu-model-runner, puis est réassemblé en RequestOutput par l'OutputProcessor. La manière dont le cœur du moteur orchestre ces étapes est décrite dans /engine/engine-core.
Voir la documentation officielle : Documentation vLLM · README