Skip to content

LogitsProcessor : transformations des logits avant échantillonnage

源码版本v0.25.1

Responsabilités

Une fois que le Sampler dispose des logits bruts, toute une série de transformations reste à appliquer avant l'échantillonnage : MinP filtre les tokens de faible probabilité, LogitBias ajoute un biais à certains tokens, MinTokens interdit EOS tant que la génération n'est pas assez longue, les pénalités repetition/frequency/presence sanctionnent les répétitions, le structured output masque les tokens hors grammaire, et le thinking budget impose un budget aux modèles de raisonnement. L'abstraction unificatrice de ces transformations est LogitsProcessor(interface.py:60). Chaque processor conserve son état par batch et est invoqué successivement avant chaque forward par Sampler.apply_logits_processors(sampler.py:371).

L'interface LogitsProcessor est légère, quatre méthodes : __init__, apply(logits) -> logits, is_argmax_invariant() -> bool(interface.py:84) et update_state(batch_update)(interface.py:94). is_argmax_invariant détermine si elle est rangée dans LogitsProcessors.argmax_invariant ou non_argmax_invariant(state.py:148-160) : les premières ne tournent qu'en échantillonnage aléatoire (MinP), les secondes aussi en greedy (MinTokens, LogitBias). update_state est appelée chaque fois que la composition du batch change ; le processor en profite pour rafraîchir ses tenseurs d'état par requête.

Les processors intégrés sont listés dans BUILTIN_LOGITS_PROCESSORS(__init__.py:49-53) dans l'ordre [MinTokensLogitsProcessor, LogitBiasLogitsProcessor, MinPLogitsProcessor]. S'y ajoutent les plugins entry-point (groupe vllm.logits_processors) et les chaînes FQDN(__init__.py:86-155) ; build_logitsprocs(__init__.py:184) fusionne les trois sources et instancie le tout.

Motivation de conception

  • Classification par invariance argmax : MinPLogitsProcessor.is_argmax_invariant() = True(builtin.py:47-49), car il ne fait que filtrer des tokens de probabilité négligeable, sans affecter l'argmax ; MinTokens / LogitBias retournent False(builtin.py:183-186), car ils peuvent changer le résultat greedy (masquer EOS, ou modifier l'argmax après ajout d'un biais).
  • Mise à jour incrémentale de l'état de batch : update_state reçoit un BatchUpdate (removed/added/moved)(interface.py:36-57), et le processor ne recalcule ses tenseurs que lorsque la composition du batch change réellement (MinP saute les batches vides via min_p_count)(builtin.py:54-100).
  • Double tenseur CPU↔GPU : MinPLogitsProcessor détient min_p_cpu (pinned) + min_p_device(builtin.py:30-44) ; la mise à jour se fait côté CPU, et une copie H2D a lieu avant apply, pour éviter une synchro inter-périphérique à chaque apply.
  • MinP passe par softmax, pas par les logits : apply calcule d'abord softmax(logits) pour obtenir la probabilité max, puis filtre selon le seuil max_prob * min_p(builtin.py:102-116), car min_p est un seuil dans l'espace des probabilités, pas des logits.
  • LogitBias en indexation clairsemée : LogitBiasLogitsProcessor.apply fait logits[req_idx, tok_id] += bias(builtin.py:159-162), ne modifiant que les positions concernées, sans addition sur tout le vocabulaire.
  • MinTokens masque EOS : MinTokensLogitsProcessor masque les stop_token_ids tant que la génération n'a pas atteint min_tokens(builtin.py:165-186) pour forcer une génération suffisamment longue.
  • Spec decode désactive les processors personnalisés : build_logitsprocs ne conserve que MinTokensLogitsProcessor quand speculative_config est non vide(__init__.py:201-209) ; MinP et LogitBias sont désactivés car incompatibles avec le rejection sampling.

Fichiers clés

Flux de données

Avant chaque forward, le Scheduler ou le GPUModelRunner construit un BatchUpdate qui indique au processor « quelles requêtes ont été ajoutées, retirées ou déplacées dans ce batch ». Le processor en profite pour rafraîchir ses tenseurs par requête. Le update_state de MinP est un exemple typique :

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 est un compteur d'échappement : si aucune requête du batch n'a activé min_p, apply retourne immédiatement sans rien calculer. Le corps de apply est lui aussi très court :

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 appelle les processors non argmax-invariant et argmax-invariant en deux phases : les premiers sont exécutés aussi sur le chemin greedy, les seconds uniquement sur le chemin aléatoire(sampler.py:403-405). Les pénalités passent par une opération apply_all_penalties indépendante et ne traversent pas la couche LogitsProcessor.

Limites et échecs

  • Pas de processor pour les modèles pooling : build_logitsprocs retourne un LogitsProcessors vide quand is_pooling_model=True, et lève STR_POOLING_REJECTS_LOGITSPROCS si un processor personnalisé est passé(__init__.py:191-198).
  • Spec decode désactive les processors personnalisés : STR_SPEC_DEC_REJECTS_LOGITSPROCS(__init__.py:201-209) ne conserve que MinTokensLogitsProcessor, car MinP / LogitBias sont incompatibles avec le rejection sampling.
  • TPU ne supporte pas les processors personnalisés : _load_custom_logitsprocs retourne une liste vide quand is_tpu() est vrai(__init__.py:176-179) ; v1 n'a pas encore de branchement des processors personnalisés sur TPU.
  • Échec de chargement d'un plugin : erreur brute : _load_logitsprocs_plugins lève directement raise RuntimeError si un entry point échoue au chargement(__init__.py:78-82) ; pas de sauit silencieux, pour éviter de découvrir à l'exécution qu'un processor n'a pas été appliqué.
  • Le FQCN doit être une sous-classe de LogitsProcessor : _load_logitsprocs_by_fqcns découpe <module>:<Qualname>(__init__.py:128) et lève un ValueError si ce n'est pas une sous-classe, pour empêcher de passer une fonction ordinaire par erreur.
  • Sémantique de référence pour batch_update : le output_tok_ids de chaque tuple dans BatchUpdate.added est une référence vers la liste des running tokens de la requête(interface.py:44-50) ; le processor voit ainsi le dernier token généré à travers cette référence, sans mise à jour nécessaire.

Résumé

LogitsProcessor est la couche de transformations plug-in située au-dessus du Sampler, avant l'échantillonnage : argmax-invariant décide si elle s'exécute en chemin greedy ou aléatoire, et update_state lui permet de suivre les changements incrémentaux du batch. Les trois processors intégrés (MinP / LogitBias / MinTokens) sont complétés par les plugins et les chaînes FQDN pour l'extension externe. Le processus d'échantillonnage proprement dit est décrit dans /sampling/sampler ; les pénalités et bad_words sont gérées directement dans le Sampler via apply_all_penalties / apply_bad_words, sans passer par cette couche.

Voir la documentation officielle : Documentation vLLM · README