Skip to content

LogitsProcessor: transformación de logits antes del muestreo

源码版本v0.25.1

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 / LogitBias son False (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_state recibe un BatchUpdate (removed/added/moved) (interface.py:36-57); el processor solo recalcula sus tensores cuando la composición del batch cambia de verdad (MinP salta con min_p_count si el lote está vacío) (builtin.py:54-100).
  • Tensores dobles CPU↔GPU: MinPLogitsProcessor mantiene min_p_cpu (pinned) + min_p_device (builtin.py:30-44); en update modifica el valor en CPU, y antes de apply hace una copia H2D, evitando una sincronización entre dispositivos en cada apply.
  • MinP va por softmax, no por logits: apply primero calcula softmax(logits) para encontrar la max prob, y luego filtra por el umbral max_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.apply hace logits[req_idx, tok_id] += bias (builtin.py:159-162), solo modifica las posiciones con bias, sin sumar sobre todo el vocabulario.
  • MinTokens enmascara EOS: MinTokensLogitsProcessor enmascara los stop_token_ids mientras la cantidad generada no llegue a min_tokens (builtin.py:165-186), forzando al modelo a generar lo suficientemente largo.
  • Spec decode desactiva processors custom: cuando speculative_config no es vacío, build_logitsprocs solo conserva MinTokensLogitsProcessor (__init__.py:201-209); MinP / LogitBias se apagan, porque son incompatibles con la semántica de rejection sampling.

Archivos clave

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:

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 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:

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 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_logitsprocs devuelve un LogitsProcessors vacío cuando is_pooling_model=True; si se pasó algo custom, lanza STR_POOLING_REJECTS_LOGITSPROCS (__init__.py:191-198).
  • Spec decode desactiva processors custom: STR_SPEC_DEC_REJECTS_LOGITSPROCS (__init__.py:201-209) deja solo MinTokensLogitsProcessor, porque MinP / LogitBias son incompatibles con rejection sampling.
  • TPU no soporta processors custom: _load_custom_logitsprocs devuelve directamente una lista vacía cuando is_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_plugins lanza directamente raise 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_fqcns parte la cadena como <module>:<Qualname> (__init__.py:128) y si no es subclase lanza ValueError, evitando pasar por error una función común.
  • Semántica de referencia en batch_update: el output_tok_ids en cada tupla de BatchUpdate.added es 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.

Véase la documentación oficial: vLLM 文档 · README.