LogitsProcessor: transformación de logits antes del muestreo
Responsabilidades
Después de que Sampler obtiene los logits crudos, aún hay que aplicar una pila de transformaciones antes de poder muestrear: MinP filtra tokens de baja probabilidad, LogitBias suma un sesgo a tokens concretos, MinTokens prohíbe EOS hasta que se generen suficientes, repetition/frequency/presence penalty penalizan repetición, structured output enmascara los tokens fuera de la gramática, y thinking budget añade un budget a modelos de razonamiento. La abstracción unificada de estas transformaciones es LogitsProcessor (interface.py:60); cada processor mantiene su estado de lote, y antes de cada forward Sampler.apply_logits_processors (sampler.py:371) los va llamando en orden.
La interfaz de LogitsProcessor es ligera, solo cuatro métodos: __init__, apply(logits) -> logits, is_argmax_invariant() -> bool (interface.py:84) y update_state(batch_update) (interface.py:94). is_argmax_invariant decide si el processor se asigna a LogitsProcessors.argmax_invariant o a non_argmax_invariant (state.py:148-160): los primeros solo corren en muestreo random (MinP), los segundos también corren en muestreo greedy (MinTokens, LogitBias). update_state se llama cada vez que cambia la composición del batch, y el processor aprovecha para refrescar los tensores por-request que mantiene.
Los processors internos están en BUILTIN_LOGITS_PROCESSORS (__init__.py:49-53) en el orden [MinTokensLogitsProcessor, LogitBiasLogitsProcessor, MinPLogitsProcessor]. Además se soportan plugins por entry-point (grupo vllm.logits_processors) y cadenas FQCN (__init__.py:86-155); build_logitsprocs (__init__.py:184) combina las tres clases y las instancia.
Motivación de diseño
- Clasificación por invarianza a argmax:
MinPLogitsProcessor.is_argmax_invariant() = True(builtin.py:47-49), porque solo filtra tokens de probabilidad muy baja y no afecta el argmax;MinTokens/LogitBiassonFalse(builtin.py:183-186), porque sí pueden cambiar el resultado greedy (enmascarar EOS o sumar un sesgo puede mover el argmax). - Actualización incremental del estado de batch:
update_staterecibe unBatchUpdate(removed/added/moved) (interface.py:36-57); el processor solo recalcula sus tensores cuando la composición del batch cambia de verdad (MinP salta conmin_p_countsi el lote está vacío) (builtin.py:54-100). - Tensores dobles CPU↔GPU:
MinPLogitsProcessormantienemin_p_cpu(pinned) +min_p_device(builtin.py:30-44); en update modifica el valor en CPU, y antes deapplyhace una copia H2D, evitando una sincronización entre dispositivos en cadaapply. - MinP va por softmax, no por logits:
applyprimero calculasoftmax(logits)para encontrar la max prob, y luego filtra por el umbralmax_prob * min_p(builtin.py:102-116), porque min_p es un umbral en el espacio de probabilidad, no en el espacio de logits. - LogitBias usa indexación dispersa:
LogitBiasLogitsProcessor.applyhacelogits[req_idx, tok_id] += bias(builtin.py:159-162), solo modifica las posiciones con bias, sin sumar sobre todo el vocabulario. - MinTokens enmascara EOS:
MinTokensLogitsProcessorenmascara losstop_token_idsmientras la cantidad generada no llegue amin_tokens(builtin.py:165-186), forzando al modelo a generar lo suficientemente largo. - Spec decode desactiva processors custom: cuando
speculative_configno es vacío,build_logitsprocssolo conservaMinTokensLogitsProcessor(__init__.py:201-209); MinP / LogitBias se apagan, porque son incompatibles con la semántica de rejection sampling.
Archivos clave
LogitsProcessor ABC:60— clase abstracta con cuatro métodos:__init__,apply,is_argmax_invariant,update_state.BatchUpdate:36— frozen dataclass con tres secuencias: removed/added/moved.LogitsProcessors:148— contenedor que parte en dos rutas por argmax_invariant.LogitsProcessors.all:162-165— concatena las dos rutas en un iterator.MinPLogitsProcessor:23— filtrado por umbral en espacio de probabilidad, argmax-invariant.LogitBiasLogitsProcessor:119— suma de sesgo disperso, non-argmax-invariant.MinTokensLogitsProcessor:165— enmascara stop tokens hasta generar lo suficiente, non-argmax-invariant.BUILTIN_LOGITS_PROCESSORS:49-53— orden por defecto de las tres piezas: MinTokens → LogitBias → MinP.build_logitsprocs:184— entrada de fábrica, con tres cortocircuitos para pooling / spec-decode / custom._load_logitsprocs_by_fqcns:86-155— carga perezosa de cadenas<module>:<Qualname>+ validación conissubclass.ThinkingBudgetStateHolder:33— processor de budget de razonamiento;apply_to_logitses invocado por el Sampler durante el forward.
Flujo de datos
Antes de cada forward, el Scheduler o GPUModelRunner construye un BatchUpdate para decirle al processor "este lote añadió estas requests, quitó estas, movió estas". El processor aprovecha para refrescar sus tensores por-request. El update_state de MinP es un ejemplo típico:
# 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 es un contador de escape: si en todo el lote ninguna request tiene min_p activado, apply hace return directo sin ningún cómputo. El cuerpo de apply también es corto:
# 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 llama los processors non-argmax-invariant y argmax-invariant en dos tramos: los primeros corren también en el camino greedy, los segundos solo en el camino random (sampler.py:403-405). Las penalties van por una op independiente apply_all_penalties, no pasan por la capa LogitsProcessor.
Límites y fallos
- Los modelos de pooling no soportan processors:
build_logitsprocsdevuelve unLogitsProcessorsvacío cuandois_pooling_model=True; si se pasó algo custom, lanzaSTR_POOLING_REJECTS_LOGITSPROCS(__init__.py:191-198). - Spec decode desactiva processors custom:
STR_SPEC_DEC_REJECTS_LOGITSPROCS(__init__.py:201-209) deja soloMinTokensLogitsProcessor, porque MinP / LogitBias son incompatibles con rejection sampling. - TPU no soporta processors custom:
_load_custom_logitsprocsdevuelve directamente una lista vacía cuandois_tpu()(__init__.py:176-179); v1 aún no cablea processors custom en TPU. - Fallos de carga de plugins se relanzan: si un entry point falla al cargar,
_load_logitsprocs_pluginslanza directamenteraise RuntimeError(__init__.py:78-82), sin saltárselo en silencio, para que no se descubra en runtime que un processor no aplicó. - FQCN debe ser subclase de LogitsProcessor:
_load_logitsprocs_by_fqcnsparte la cadena como<module>:<Qualname>(__init__.py:128) y si no es subclase lanzaValueError, evitando pasar por error una función común. - Semántica de referencia en batch_update: el
output_tok_idsen cada tupla deBatchUpdate.addedes una referencia a la lista de tokens en running de la request (interface.py:44-50); por esa referencia el processor ve el último token generado, sin necesidad de actualizar manualmente.
Resumen
LogitsProcessor es la capa de transformaciones enchufables, sobre Sampler y antes del muestreo: argmax-invariant decide si corre en el camino greedy o en el random, y update_state le permite seguir cambios incrementales del batch. Las tres piezas internas MinP / LogitBias / MinTokens, más plugins y cadenas FQCN, permiten extensión externa. El proceso de muestreo concreto está en /sampling/sampler; las penalties y bad_words se invocan directamente como apply_all_penalties / apply_bad_words dentro de Sampler, sin pasar por esta abstracción.