Skip to content

vllm.v1.worker.gpu.sample.logits_processor

Custom logits processors for the V2 model runner.

Kept import-light: the frontend process imports this package to validate per-request params, so nothing here may pull in model-runner side modules (torch, triton, worker state) at import time.

Modules:

  • interface –

    V2 logits processor interface.

  • loader –

    Loading of custom logits processor classes for the V2 model runner.

Classes:

Functions:

LogitsContext dataclass

The current step's batch layout, passed to every apply() call.

A row is a logits row, not a request: rows are reordered every step, and under speculative decoding a request owns one row per draft token. Committed tokens live in req_states.all_token_ids (valid up to total_len); this step's draft tokens are only in input_ids.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
@dataclass(frozen=True)
class LogitsContext:
    """The current step's batch layout, passed to every ``apply()`` call.

    A row is a logits row, not a request: rows are reordered every step, and
    under speculative decoding a request owns one row per draft token.
    Committed tokens live in ``req_states.all_token_ids`` (valid up to
    ``total_len``); this step's draft tokens are only in ``input_ids``.
    """

    # [num_logits_rows] row -> persistent request slot.
    expanded_idx_mapping: torch.Tensor
    # [num_reqs] batch position -> persistent request slot.
    idx_mapping: torch.Tensor
    # [num_reqs] batch position -> persistent request slot, on the host, for
    # skipping work without a device sync.
    idx_mapping_np: np.ndarray
    # [num_logits_rows] row -> its offset among the rows of its own request.
    expanded_local_pos: torch.Tensor
    # [num_logits_rows] token fed to the model at each row's input position.
    input_ids: torch.Tensor
    # [num_logits_rows] position of each row within its sequence.
    pos: torch.Tensor
    # [num_reqs] batch position -> upper bound of tokens visible to the token
    # being sampled this step, on the host. Exact when spec decoding isn't in use.
    # Exact per-row lengths on device are `pos + 1`.
    seq_lens_upper_bound_np: np.ndarray

LogitsProcRequestState dataclass

State associated with active requests, shared with logits processors.

Wraps the model runner's per-slot buffers, which are mutated in place, so reads always see current values. Processors must treat every field as read-only.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
@dataclass(frozen=True)
class LogitsProcRequestState:
    """State associated with active requests, shared with logits processors.

    Wraps the model runner's per-slot buffers, which are mutated in place,
    so reads always see current values. Processors must treat every field
    as read-only.
    """

    device: torch.device
    max_num_reqs: int
    vocab_size: int

    # [max_num_reqs, max_model_len] committed token ids per request slot.
    all_token_ids: StagedWriteTensor
    # [max_num_reqs] tokens in the user-provided prompt.
    prompt_len: UvaBackedTensor
    # [max_num_reqs] tokens fed at the latest (re)fill: the prompt plus any
    # partial output on resumption after preemption.
    prefill_len: UvaBackedTensor
    # [max_num_reqs] prompt_len + output_len; grows as the request progresses.
    total_len: StagedWriteTensor

    @classmethod
    def from_request_state(cls, req_states: RequestState) -> LogitsProcRequestState:
        return cls(
            device=req_states.device,
            max_num_reqs=req_states.max_num_reqs,
            vocab_size=req_states.vocab_size,
            all_token_ids=req_states.all_token_ids,
            prompt_len=req_states.prompt_len,
            prefill_len=req_states.prefill_len,
            total_len=req_states.total_len,
        )

LogitsProcessor

Bases: ABC

Custom logits processor for Model Runner V2.

Per-request state is keyed by the request slot index; slots are recycled through a free list, so per-slot state must be fully (re)initialized in add_request().

apply() runs after the built-in bias, penalty, bad-words and grammar stages and before temperature, min_p and top-k/top-p, so it sees unscaled logits and must not re-inflate grammar-masked tokens. Thinking-budget forcing runs after apply() and wins over its edits for requests whose budget is exhausted.

State that is constant for a request belongs in __init__() or add_request(); apply() receives only what changes per step.

Methods:

  • __init__ –

    Capture what stays constant for the processor's lifetime.

  • add_request –

    Initialize per-slot state for a request entering the batch.

  • apply –

    Modify logits in place or return a new tensor. In-place modification

  • apply_staged_writes –

    Flush any host-side writes staged by add_request() to the device.

  • validate_params –

    Raise ValueError for invalid per-request arguments.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
class LogitsProcessor(ABC):
    """Custom logits processor for Model Runner V2.

    Per-request state is keyed by the request slot index; slots are recycled
    through a free list, so per-slot state must be fully (re)initialized in
    ``add_request()``.

    ``apply()`` runs after the built-in bias, penalty, bad-words and grammar
    stages and before temperature, min_p and top-k/top-p, so it sees unscaled
    logits and must not re-inflate grammar-masked tokens. Thinking-budget
    forcing runs after ``apply()`` and wins over its edits for requests whose
    budget is exhausted.

    State that is constant for a request belongs in ``__init__()`` or
    ``add_request()``; ``apply()`` receives only what changes per step.
    """

    def __init__(  # noqa: B027
        self, vllm_config: VllmConfig, req_states: LogitsProcRequestState
    ):
        """Capture what stays constant for the processor's lifetime.

        ``req_states`` exposes the on-device token history and batch constants a
        processor may read. Treat it as read-only.
        """

    @classmethod  # noqa: B027
    def validate_params(cls, sampling_params: SamplingParams) -> None:
        """Raise ``ValueError`` for invalid per-request arguments.

        Runs at request admission, so invalid arguments fail the request
        with an error instead of reaching the sampler.
        """

    def add_request(self, req_idx: int, sampling_params: SamplingParams) -> bool:
        """Initialize per-slot state for a request entering the batch.

        The slot may hold a previous occupant's state; overwrite or
        neutralize all of it here.

        Returns whether this processor modifies logits for the request.
        """
        return True

    def apply_staged_writes(self) -> None:  # noqa: B027
        """Flush any host-side writes staged by ``add_request()`` to the device.

        Called once per step before the forward pass, after the model runner
        has flushed ``req_states``, so a processor that stages writes here can
        read the request's tokens back on device.
        """

    @abstractmethod
    def apply(self, logits: torch.Tensor, ctx: LogitsContext) -> torch.Tensor:
        """Modify logits in place or return a new tensor. In-place modification
        is preferred for efficiency.

        ``apply()`` is called once for the whole batch, including rows of
        requests this processor declined in ``add_request()``, so filter rows
        via ``ctx.expanded_idx_mapping``.

        Args:
            logits: [num_logits_rows, vocab_size] float32 tensor.
            ctx: this step's batch layout.

        """
        raise NotImplementedError

__init__(vllm_config, req_states)

Capture what stays constant for the processor's lifetime.

req_states exposes the on-device token history and batch constants a processor may read. Treat it as read-only.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
def __init__(  # noqa: B027
    self, vllm_config: VllmConfig, req_states: LogitsProcRequestState
):
    """Capture what stays constant for the processor's lifetime.

    ``req_states`` exposes the on-device token history and batch constants a
    processor may read. Treat it as read-only.
    """

add_request(req_idx, sampling_params)

Initialize per-slot state for a request entering the batch.

The slot may hold a previous occupant's state; overwrite or neutralize all of it here.

Returns whether this processor modifies logits for the request.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
def add_request(self, req_idx: int, sampling_params: SamplingParams) -> bool:
    """Initialize per-slot state for a request entering the batch.

    The slot may hold a previous occupant's state; overwrite or
    neutralize all of it here.

    Returns whether this processor modifies logits for the request.
    """
    return True

apply(logits, ctx) abstractmethod

Modify logits in place or return a new tensor. In-place modification is preferred for efficiency.

apply() is called once for the whole batch, including rows of requests this processor declined in add_request(), so filter rows via ctx.expanded_idx_mapping.

Parameters:

  • logits

    (Tensor) –

    [num_logits_rows, vocab_size] float32 tensor.

  • ctx

    (LogitsContext) –

    this step's batch layout.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
@abstractmethod
def apply(self, logits: torch.Tensor, ctx: LogitsContext) -> torch.Tensor:
    """Modify logits in place or return a new tensor. In-place modification
    is preferred for efficiency.

    ``apply()`` is called once for the whole batch, including rows of
    requests this processor declined in ``add_request()``, so filter rows
    via ``ctx.expanded_idx_mapping``.

    Args:
        logits: [num_logits_rows, vocab_size] float32 tensor.
        ctx: this step's batch layout.

    """
    raise NotImplementedError

apply_staged_writes()

Flush any host-side writes staged by add_request() to the device.

Called once per step before the forward pass, after the model runner has flushed req_states, so a processor that stages writes here can read the request's tokens back on device.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
def apply_staged_writes(self) -> None:  # noqa: B027
    """Flush any host-side writes staged by ``add_request()`` to the device.

    Called once per step before the forward pass, after the model runner
    has flushed ``req_states``, so a processor that stages writes here can
    read the request's tokens back on device.
    """

validate_params(sampling_params) classmethod

Raise ValueError for invalid per-request arguments.

Runs at request admission, so invalid arguments fail the request with an error instead of reaching the sampler.

Source code in vllm/v1/worker/gpu/sample/logits_processor/interface.py
@classmethod  # noqa: B027
def validate_params(cls, sampling_params: SamplingParams) -> None:
    """Raise ``ValueError`` for invalid per-request arguments.

    Runs at request admission, so invalid arguments fail the request
    with an error instead of reaching the sampler.
    """

build_custom_logits_processors(vllm_config, req_states, is_pooling_model, custom_logitsprocs=())

Load and instantiate custom logits processors, entrypoint plugins first.

Raises:

  • ValueError –

    if a pooling model specifies custom processors, or a loaded class does not implement the V2 interface.

  • RuntimeError –

    if an FQCN fails to import.

Source code in vllm/v1/worker/gpu/sample/logits_processor/loader.py
def build_custom_logits_processors(
    vllm_config: "VllmConfig",
    req_states: "RequestState",
    is_pooling_model: bool,
    custom_logitsprocs: Sequence[str | type] = (),
) -> list[LogitsProcessor]:
    """Load and instantiate custom logits processors, entrypoint plugins first.

    Raises:
        ValueError: if a pooling model specifies custom processors, or a
            loaded class does not implement the V2 interface.
        RuntimeError: if an FQCN fails to import.

    """
    from vllm.v1.sample.logits_processor import STR_POOLING_REJECTS_LOGITSPROCS

    if is_pooling_model:
        if custom_logitsprocs:
            raise ValueError(STR_POOLING_REJECTS_LOGITSPROCS)
        logger.debug(
            "Skipping logits processor loading because pooling models"
            " do not support logits processors."
        )
        return []

    custom_logitsprocs_classes = _load_v2_logitsprocs(custom_logitsprocs)
    lp_req_state = LogitsProcRequestState.from_request_state(req_states)
    return [ctor(vllm_config, lp_req_state) for ctor in custom_logitsprocs_classes]

build_custom_logits_processors_params_validator(custom_logitsprocs)

Load custom processor classes once and return a params validator.

Called from the frontend at startup. The returned callable runs each processor's validate_params at request admission.

Raises:

  • ValueError –

    if a loaded class does not implement the V2 interface.

  • RuntimeError –

    if an FQCN fails to import.

Source code in vllm/v1/worker/gpu/sample/logits_processor/loader.py
def build_custom_logits_processors_params_validator(
    custom_logitsprocs: Sequence[str | type] | None,
) -> "Callable[[SamplingParams], None]":
    """Load custom processor classes once and return a params validator.

    Called from the frontend at startup. The returned callable runs each
    processor's ``validate_params`` at request admission.

    Raises:
        ValueError: if a loaded class does not implement the V2 interface.
        RuntimeError: if an FQCN fails to import.

    """
    classes = _cached_load_v2_logitsprocs(tuple(custom_logitsprocs or ()))

    if not classes:
        return lambda _: None

    def validate_params(sampling_params: "SamplingParams") -> None:
        for cls in classes:
            try:
                cls.validate_params(sampling_params)
            except ValueError as e:
                raise VLLMValidationError(str(e)) from e

    return validate_params