How to Implement Custom Logit Processors and Sampling Parameters in vLLM

You can implement custom logit processors in vLLM by subclassing the LogitsProcessor base class, implementing the required life-cycle hooks, and registering the processor via FQCN, entry points, or direct class injection when constructing an LLM or AsyncLLM engine.

vLLM is a high-throughput inference engine that processes raw model logits through a pipeline of logits processors before final token selection. Understanding how to implement custom logit processors and sampling parameters in vLLM allows you to inject arbitrary token-level logic—such as dynamic masking, bias injection, or constrained decoding—while maintaining compatibility with tensor-parallel generation and continuous batching.

Core Architecture of vLLM Logits Processing

SamplingParams and Per-Request Configuration

The SamplingParams class in vllm/sampling_params.py holds per-request generation options including temperature, top_p, logit_bias, allowed_token_ids, and extra_args. When a request enters the engine, vLLM inspects these fields to build a list of logits processors specific to that request. Custom arguments can be passed via extra_args and accessed within your processor implementation.

LogitsProcessor Base Class

All built-in and custom processors inherit from LogitsProcessor, defined in vllm/model_executor/layers/logits_processor.py. This abstract base class defines the contract between your custom logic and the vLLM engine. The engine calls the processor's apply() method on every generation step, passing a (num_requests) × (vocab_size) tensor of logits that your code can inspect and modify.

Engine Wiring and Registration

When an LLM or AsyncLLM instance is created, vLLM scans the logits_processors constructor argument to discover processor classes. You can provide these as fully-qualified class names (FQCN), entry-point references, or direct class objects. The engine instantiates one copy of each processor class and stores it internally, invoking it across all applicable requests during the sampling loop.

Implementing a Custom Logit Processor

Required Hooks and Life-Cycle

A valid LogitsProcessor subclass must implement five key hooks that the vLLM engine calls in sequence:

  1. validate_params(cls, sampling_params) – Called once during request creation. Raise ValueError if extra_args or other fields are malformed.
  2. __init__(self, vllm_config, device, is_pin_memory) – Called once at engine startup. Use this to allocate buffers or initialize per-request metadata dictionaries.
  3. is_argmax_invariant(self) – Return True if your processor never changes the argmax token. When True, vLLM skips the processor during greedy decoding to save computation.
  4. update_state(self, batch_update) – Called whenever the persistent batch changes (requests added, removed, or moved). Update internal mappings from request indices to custom data.
  5. apply(self, logits) – Called every generation step. Modify the logits tensor in-place or return a new tensor.

Complete Example: Target Token Masking

The following implementation, derived from the official vLLM documentation, creates a processor that masks all logits except for a user-specified target token:

import torch
from vllm.v1.sample.logits_processor import (
    BatchUpdate,
    LogitsProcessor,
    MoveDirectionality,
)
from vllm.config import VllmConfig
from vllm.sampling_params import SamplingParams


class DummyLogitsProcessor(LogitsProcessor):
    """Keeps only a user-specified target token, masks everything else."""

    @classmethod
    def validate_params(cls, params: SamplingParams):
        # The processor expects an integer `target_token` in extra_args.

        target = params.extra_args and params.extra_args.get("target_token")
        if target is not None and not isinstance(target, int):
            raise ValueError(f"target_token must be int, got {type(target)}")

    def __init__(self, vllm_config: VllmConfig,
                 device: torch.device, is_pin_memory: bool):
        # Mapping request_index → target token id

        self.req_info: dict[int, int] = {}

    def is_argmax_invariant(self) -> bool:
        # This processor can change which token has the highest logit.

        return False

    def update_state(self, batch_update: BatchUpdate | None):
        if not batch_update:
            return

        # 1) New requests – read the custom argument.

        for idx, params, _, _ in batch_update.added:
            self.validate_params(params)
            if params.extra_args and (tok := params.extra_args.get("target_token")):
                self.req_info[idx] = tok
            else:
                self.req_info.pop(idx, None)

        # 2) Removed requests – clean up our dict.

        for idx in batch_update.removed:
            self.req_info.pop(idx, None)

        # 3) Moves – keep indices in sync.

        for src, dst, direction in batch_update.moved:
            val = self.req_info.pop(src, None)
            if val is not None:
                self.req_info[dst] = val
            # Swaps (direction == MoveDirectionality.SWAP) are handled

            # automatically because we pop both sides.

    def apply(self, logits: torch.Tensor) -> torch.Tensor:
        # If no request uses the processor, return unchanged.

        if not self.req_info:
            return logits

        # Gather row/col indices for the target tokens.

        rows = torch.tensor(list(self.req_info.keys()),
                            dtype=torch.long, device=logits.device)
        cols = torch.tensor(list(self.req_info.values()),
                            dtype=torch.long, device=logits.device)

        # Preserve the original values for the target tokens.

        saved = logits[rows, cols].clone()

        # Mask everything else.

        logits[rows] = float("-inf")
        logits[rows, cols] = saved

        return logits

Source: [custom_logitsprocs.md](https://github.com/vllm-project/vllm/blob/main/docs/features/custom_logitsprocs.md)

Loading Methods: FQCN, Entry Points, and Direct Injection

vLLM supports three methods for registering your custom processor:

1. Fully-Qualified Class Name (FQCN)

Pass the module path and class name as a string when constructing the engine:

from vllm import LLM, SamplingParams

logits_processor_fqcn = "my_module.my_processors:DummyLogitsProcessor"

llm = LLM(
    model="facebook/opt-125m",
    logits_processors=[logits_processor_fqcn],
)

params = SamplingParams(
    temperature=0.7,
    extra_args={"target_token": 42},
)

output = llm.generate("Hello world", sampling_params=params)

2. Entry-Point Registration

Add an entry point to your pyproject.toml for automatic discovery:

[project.entry-points."vllm.logits_processors"]
my_processor = "my_pkg.path:MyProcessor"

The engine automatically loads any processors registered under the vllm.logits_processors group on startup.

3. Direct Class Object

Import the class directly and pass the object:

from my_pkg.path import MyProcessor
from vllm import LLM

llm = LLM(
    model="facebook/opt-125m",
    logits_processors=[MyProcessor],
)

Source: [custom_logitsprocs.md](https://github.com/vllm-project/vllm/blob/main/docs/features/custom_logitsprocs.md)

Extending Sampling Parameters

Adding Custom Fields to SamplingParams

While SamplingParams in vllm/sampling_params.py provides built-in fields like temperature, top_p, and logit_bias, you can pass arbitrary configuration through the extra_args dictionary. This field accepts any JSON-serializable dict and is exposed to custom processors via the SamplingParams object passed to validate_params and update_state.

For example, to add a custom temperature scaling factor distinct from the main temperature field:


# In vllm/sampling_params.py (add near other fields)

temp_scale: float = 1.0   # New custom knob

Accessing Custom Parameters in Processors

Your processor can read custom arguments from two sources:

  1. Built-in fields – Directly access sampling_params.temp_scale if you modified the data class.
  2. Extra arguments – Access sampling_params.extra_args.get("your_key") for ad-hoc configuration.

In the apply method, use the stored per-request metadata (populated during update_state) to modify logits:

def apply(self, logits: torch.Tensor) -> torch.Tensor:
    # Example: multiply logits by custom scale factor stored in self.req_info

    for req_idx, scale in self.req_info.items():
        logits[req_idx] *= scale
    return logits

Note: After modifying SamplingParams, run the test suite to ensure backward compatibility; the repository's CI verifies that unknown fields are correctly ignored by default to prevent breaking changes.

Summary

  • Subclass LogitsProcessor from vllm/model_executor/layers/logits_processor.py and implement validate_params, __init__, is_argmax_invariant, update_state, and apply.
  • Register processors via FQCN strings, entry points in pyproject.toml, or direct class objects when constructing LLM or AsyncLLM.
  • Pass per-request configuration through SamplingParams.extra_args (JSON-serializable dict) or by extending the SamplingParams dataclass with new fields.
  • Manage state using the BatchUpdate object in update_state to handle request additions, removals, and index moves in the persistent batch.
  • Optimize for greedy decoding by returning True from is_argmax_invariant() when your processor cannot change the argmax token, allowing vLLM to skip unnecessary computation.

Frequently Asked Questions

How do I access custom arguments passed from the OpenAI API in my logit processor?

When using the OpenAI-compatible API, pass custom arguments via extra_body["vllm_xargs"] or the extra_args field in the Python SDK. In your processor's validate_params and update_state methods, read these values from sampling_params.extra_args.get("your_key"). The engine automatically deserializes these JSON values and passes them through the SamplingParams object.

What is the difference between logit_bias and a custom logit processor?

logit_bias is a built-in field in SamplingParams that vLLM automatically converts into an internal LogitBiasProcessor. It accepts a dictionary mapping token IDs to bias values and is suitable for simple additive biases. A custom logit processor is a Python class you write that implements the full LogitsProcessor interface, allowing complex stateful logic, dynamic masking, or external constraint checking that goes beyond simple bias addition.

Can I use multiple logit processors simultaneously?

Yes. The logits_processors parameter accepts a list of processor classes or FQCN strings. vLLM instantiates each processor once and chains them in the order provided. Each processor's apply method receives the logits tensor modified by the previous processor in the chain. Ensure that processors earlier in the list do not produce invalid logit values (like NaN) that would break subsequent processors.

How do I handle request state that persists across generation steps?

Implement the update_state(self, batch_update) method in your processor. The batch_update object contains added, removed, and moved lists describing changes to the persistent batch. Store per-request metadata (like target tokens or custom scales) in a dictionary keyed by request index, updating it whenever update_state is called. This ensures your processor correctly handles continuous batching where requests dynamically enter, leave, or shift positions in the batch.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →