Skip to content

LogitsProcessor: logits-Transformationen vor dem Sampling

源码版本v0.25.1

Verantwortung

Nachdem Sampler die rohen logits erhalten hat, müssen noch eine Reihe von Transformationen erfolgen, bevor das Sampling beginnen kann: MinP filtert niedrigwahrscheinliche Token, LogitBias fügt bestimmten Token einen Bias hinzu, MinTokens verbietet EOS, bevor genug generiert wurde, repetition/frequency/presence penalty bestraft Wiederholungen, structured output blendet Token außerhalb der Grammatik aus, thinking budget fügt Reasoning-Modellen ein Budget hinzu. Die einheitliche Abstraktion für diese Transformationen ist LogitsProcessor (interface.py:60); jeder Processor hält seinen eigenen Batch-Zustand und wird vor jedem forward von Sampler.apply_logits_processors (sampler.py:371) nacheinander aufgerufen.

Das LogitsProcessor-Interface ist sehr schlank und besteht aus nur vier Methoden: __init__, apply(logits) -> logits, is_argmax_invariant() -> bool (interface.py:84) und update_state(batch_update) (interface.py:94). is_argmax_invariant entscheidet, ob der Processor in LogitsProcessors.argmax_invariant oder non_argmax_invariant eingeordnet wird (state.py:148-160); erstere werden nur beim random Sampling ausgeführt (MinP), letztere auch beim greedy Sampling (MinTokens, LogitBias). update_state wird bei jeder Änderung der Batch-Zusammensetzung aufgerufen; der Processor nutzt die Gelegenheit, seine gehaltenen pro-Request-Zustandstensoren zu aktualisieren.

Die eingebauten Prozessoren liegen in BUILTIN_LOGITS_PROCESSORS (__init__.py:49-53); die Reihenfolge ist [MinTokensLogitsProcessor, LogitBiasLogitsProcessor, MinPLogitsProcessor]. Darüber hinaus werden Entry-Point-Plugins (Gruppe vllm.logits_processors) und FQCN-Strings unterstützt (__init__.py:86-155); build_logitsprocs (__init__.py:184) fasst diese drei Klassen zusammen und instanziiert sie.

Entwurfsmotivation

  • Klassifikation der argmax-Invarianz: MinPLogitsProcessor.is_argmax_invariant() = True (builtin.py:47-49), weil es nur extrem niedrigwahrscheinliche Token filtert und argmax nicht beeinflusst; MinTokens / LogitBias sind False (builtin.py:183-186), weil sie das greedy-Ergebnis verändern können (nach EOS-Masking oder Bias-Zugabe kann sich argmax ändern).
  • Inkrementelle Batch-Zustandsaktualisierung: update_state erhält eine BatchUpdate (removed/added/moved) (interface.py:36-57); der Processor berechnet seine Tensoren nur neu, wenn sich die Batch-Zusammensetzung tatsächlich geändert hat (MinP nutzt min_p_count, um leere Batches zu überspringen) (builtin.py:54-100).
  • Doppelte CPU↔GPU-Tensoren: MinPLogitsProcessor hält min_p_cpu (pinned) + min_p_device (builtin.py:30-44); beim Update werden die Werte auf der CPU geändert und vor apply einmal H2D kopiert, damit nicht bei jedem apply eine geräteübergreifende Synchronisierung erfolgt.
  • MinP läuft über softmax, nicht über logits: apply berechnet zunächst softmax(logits), um die maximale Wahrscheinlichkeit zu finden, und filtert dann nach dem Schwellenwert max_prob * min_p (builtin.py:102-116), weil min_p ein Schwellenwert im Wahrscheinlichkeitsraum ist, nicht im logit-Raum.
  • LogitBias mit Sparse-Index: LogitBiasLogitsProcessor.apply nutzt logits[req_idx, tok_id] += bias (builtin.py:159-162) und verändert nur Positionen mit Bias, ohne eine Addition über das gesamte Vocab auszuführen.
  • MinTokens maskiert EOS: MinTokensLogitsProcessor maskiert stop_token_ids, solange die Anzahl generierter Token kleiner als min_tokens ist (builtin.py:165-186), damit das Modell ausreichend lange generieren muss.
  • Spec-Decode deaktiviert benutzerdefinierte Prozessoren: build_logitsprocs behält bei nicht-leerem speculative_config nur MinTokensLogitsProcessor (__init__.py:201-209); MinP und LogitBias werden abgeschaltet, weil sie mit der Semantik des Rejection-Sampling inkompatibel sind.

Schlüsseldateien

Datenfluss

Vor jedem forward konstruieren Scheduler oder GPUModelRunner eine BatchUpdate, die dem Processor mitteilt: „diese Batch hat diese Requests hinzugefügt, diese entfernt und diese verschoben". Der Processor nutzt die Gelegenheit, seine pro-Request-Tensoren aufzufrischen. update_state von MinP ist ein typisches Beispiel:

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 ist ein Escape-Zähler: Wenn in der gesamten Batch kein einziges Request min_p aktiviert hat, gibt apply direkt zurück und führt keine Berechnung aus. Der Körper von apply ist ebenfalls sehr kurz:

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 ruft non-argmax-invariant und argmax-invariant in zwei getrennten Abschnitten auf: erstere laufen auch im greedy-Pfad, letztere nur im random-Pfad (sampler.py:403-405). Penalties laufen über das separate Op apply_all_penalties und nicht über die LogitsProcessor-Schicht.

Grenzen und Fehler

  • Pooling-Modelle unterstützen keine Prozessoren: build_logitsprocs gibt bei is_pooling_model=True direkt ein leeres LogitsProcessors zurück; wenn benutzerdefinierte Prozessoren übergeben wurden, wird STR_POOLING_REJECTS_LOGITSPROCS geworfen (__init__.py:191-198).
  • Spec-Decode deaktiviert benutzerdefinierte Prozessoren: STR_SPEC_DEC_REJECTS_LOGITSPROCS (__init__.py:201-209) — nur MinTokensLogitsProcessor bleibt erhalten, weil MinP und LogitBias mit Rejection-Sampling inkompatibel sind.
  • TPU unterstützt keine benutzerdefinierten Prozessoren: _load_custom_logitsprocs gibt bei is_tpu() direkt eine leere Liste zurück (__init__.py:176-179); in v1 sind auf TPU noch keine benutzerdefinierten Prozessoren angeschlossen.
  • Plugin-Ladefehler werden hart geworfen: _load_logitsprocs_plugins macht bei einem Ladefehler eines Entry-Points direkt raise RuntimeError (__init__.py:78-82) ohne stillen Sprung, damit nicht erst zur Laufzeit auffällt, dass ein Processor nicht aktiv wurde.
  • FQCN muss eine LogitsProcessor-Unterklasse sein: _load_logitsprocs_by_fqcns splittet <module>:<Qualname> (__init__.py:128); ist das Ergebnis keine Unterklasse, wird ValueError geworfen, um versehentliche Übergabe gewöhnlicher Funktionen zu verhindern.
  • Referenzsemantik von batch_update: Der output_tok_ids-Eintrag in jedem Tuple von BatchUpdate.added ist eine Referenz auf die Running-Token-Liste des Requests (interface.py:44-50); der Processor sieht über diese Referenz stets die neuesten generierten Token, ohne dass jedes Mal ein Update erfolgen muss.

Zusammenfassung

LogitsProcessor ist die steckbare Transformationsschicht über Sampler und vor dem Sampling: is_argmax_invariant entscheidet, ob der Processor im greedy- oder im random-Pfad läuft, und update_state lässt ihn die inkrementellen Batch-Änderungen verfolgen. Eingebaut sind die drei Prozessoren MinP, LogitBias und MinTokens; daneben werden Plugins und FQCN-Strings als externe Erweiterung unterstützt. Der eigentliche Sampling-Prozess siehe /sampling/sampler; Penalty und bad_words werden in Sampler direkt über apply_all_penalties / apply_bad_words aufgerufen und laufen nicht über diese Abstraktionsschicht.

Siehe offizielle Dokumentation: vLLM-Dokumentation · README.