Skip to content

Sampler : la dernière étape, des logits aux tokens

源码版本v0.25.1

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_samplertorch.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_temperature met les positions temp < eps à 1.0(sampler.py:233-237), puis un unique torch.where(temp < eps, greedy, random) fusionne à la fin, pour éviter les branches.
  • Privilégier in-place : apply_temperature utilise logits.div_, et MinP.apply utilise masked_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 vaut processed_logits / processed_logprobs, on force forward_native.
  • Collecte des logprobs séparée de l'échantillonnage : gather_logprobs(sampler.py:308-356) tourne après l'échantillonnage, et calcule log_softmax sur 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_logprobs utilise torch._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 mode predict_bonus_token, concatène spec_token_ids à la suite de output_token_ids(sampler.py:384-393) pour que la pénalité voie les draft tokens.

Fichiers clés

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 :

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_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_greedy et all_random mutuellement exclusifs : en tête de sample, 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_logprobs s'exécute avant apply_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_ids passe par un gather dédié : gather_specific_token_logprobs(sampler.py:151-225) est utilisé par l'API generative_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_logitsprocs lève ValueError si speculative_config est 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