# How to Implement Custom Logit Processors and Sampling Parameters in vLLM

> Learn to implement custom logit processors and sampling parameters in vLLM. Extend vLLM's capabilities by subclassing LogitsProcessor and registering your custom logic for advanced text generation control.

- Repository: [vLLM/vllm](https://github.com/vllm-project/vllm)
- Tags: how-to-guide
- Published: 2026-03-03

---

**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`](https://github.com/vllm-project/vllm/blob/main/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`](https://github.com/vllm-project/vllm/blob/main/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:

```python
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/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:

```python
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`](https://github.com/vllm-project/vllm/blob/main/pyproject.toml) for automatic discovery:

```toml
[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:

```python
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/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`](https://github.com/vllm-project/vllm/blob/main/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:

```python

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

```python
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`](https://github.com/vllm-project/vllm/blob/main/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`](https://github.com/vllm-project/vllm/blob/main/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.