LogitsProcessor : transformations des logits avant échantillonnage
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/LogitBiasretournentFalse(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_statereçoit unBatchUpdate(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 viamin_p_count)(builtin.py:54-100). - Double tenseur CPU↔GPU :
MinPLogitsProcessordétientmin_p_cpu(pinned) +min_p_device(builtin.py:30-44) ; la mise à jour se fait côté CPU, et une copie H2D a lieu avantapply, pour éviter une synchro inter-périphérique à chaqueapply. - MinP passe par softmax, pas par les logits :
applycalcule d'abordsoftmax(logits)pour obtenir la probabilité max, puis filtre selon le seuilmax_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.applyfaitlogits[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 :
MinTokensLogitsProcessormasque lesstop_token_idstant que la génération n'a pas atteintmin_tokens(builtin.py:165-186) pour forcer une génération suffisamment longue. - Spec decode désactive les processors personnalisés :
build_logitsprocsne conserve queMinTokensLogitsProcessorquandspeculative_configest non vide(__init__.py:201-209) ; MinP et LogitBias sont désactivés car incompatibles avec le rejection sampling.
Fichiers clés
LogitsProcessor ABC:60— classe abstraite à quatre méthodes :__init__,apply,is_argmax_invariant,update_state.BatchUpdate:36— frozen dataclass, porte les trois séquences removed/added/moved.LogitsProcessors:148— conteneur, sépare en deux voies argmax_invariant.LogitsProcessors.all:162-165— chaîne les deux voies en un seul itérateur.MinPLogitsProcessor:23— filtrage par seuil dans l'espace des probabilités, argmax-invariant.LogitBiasLogitsProcessor:119— addition clairsemée de biais, non argmax-invariant.MinTokensLogitsProcessor:165— masque les stop tokens jusqu'à atteindre la longueur minimale, non argmax-invariant.BUILTIN_LOGITS_PROCESSORS:49-53— ordre par défaut : MinTokens → LogitBias → MinP.build_logitsprocs:184— point d'entrée usine, gère les court-circuits pooling / spec-decode / personnalisés._load_logitsprocs_by_fqcns:86-155— chargement paresseux de chaînes<module>:<Qualname>+ validationissubclass.ThinkingBudgetStateHolder:33— processor de budget de raisonnement,apply_to_logitsest appelé par le Sampler au forward.
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 :
# 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 :
# 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 logitsSampler.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_logitsprocsretourne unLogitsProcessorsvide quandis_pooling_model=True, et lèveSTR_POOLING_REJECTS_LOGITSPROCSsi 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 queMinTokensLogitsProcessor, car MinP / LogitBias sont incompatibles avec le rejection sampling. - TPU ne supporte pas les processors personnalisés :
_load_custom_logitsprocsretourne une liste vide quandis_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_pluginslève directementraise RuntimeErrorsi 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_fqcnsdécoupe<module>:<Qualname>(__init__.py:128) et lève unValueErrorsi 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_idsde chaque tuple dansBatchUpdate.addedest 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