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:
validate_params(cls, sampling_params)– Called once during request creation. RaiseValueErrorifextra_argsor other fields are malformed.__init__(self, vllm_config, device, is_pin_memory)– Called once at engine startup. Use this to allocate buffers or initialize per-request metadata dictionaries.is_argmax_invariant(self)– ReturnTrueif your processor never changes the argmax token. WhenTrue, vLLM skips the processor during greedy decoding to save computation.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.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:
- Built-in fields – Directly access
sampling_params.temp_scaleif you modified the data class. - 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
LogitsProcessorfromvllm/model_executor/layers/logits_processor.pyand implementvalidate_params,__init__,is_argmax_invariant,update_state, andapply. - Register processors via FQCN strings, entry points in
pyproject.toml, or direct class objects when constructingLLMorAsyncLLM. - Pass per-request configuration through
SamplingParams.extra_args(JSON-serializable dict) or by extending theSamplingParamsdataclass with new fields. - Manage state using the
BatchUpdateobject inupdate_stateto handle request additions, removals, and index moves in the persistent batch. - Optimize for greedy decoding by returning
Truefromis_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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →