Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# Prefix-Tokenization Cache in TensorRT LLM
# Prefix-Tokenization Cache

## Overview

Expand All @@ -21,15 +21,25 @@ The output is always identical to tokenizing the whole prompt.

## Enabling the cache

The cache is **off by default**. Enable it with:
The cache is **off by default**. Enable it with the `enable_tokenization_cache`
LLM argument, either in the `trtllm-serve` configuration file:

```bash
export TLLM_PREFIX_TOKEN_CACHE=1
```yaml
enable_tokenization_cache: true
```

or from Python:

```python
from tensorrt_llm import LLM

llm = LLM(model="zai-org/GLM-5.2", enable_tokenization_cache=True)
```

The cache size can be tuned with environment variables:

| Environment variable | Default | Behavior |
|---|---:|---|
| `TLLM_PREFIX_TOKEN_CACHE` | unset | Set to exactly `1` to enable the cache. Any other value leaves it disabled. |
| `TLLM_PREFIX_TOKEN_CACHE_ENTRIES` | 512 | Maximum number of cached prompts. Eviction is least-recently-used. |
| `TLLM_PREFIX_TOKEN_CACHE_MAX_CHARS` | 67108864 | Maximum total characters of cached prompt text. Cached ids are stored as int32, about one byte per character of English text, so this bounds host memory to roughly twice this many bytes. |
| `TLLM_PREFIX_TOKEN_CACHE_MIN_CHARS` | 4096 | Prompts shorter than this are tokenized normally and never cached. |
Expand All @@ -42,9 +52,6 @@ feature you turned on.

- Each `DefaultInputProcessor` owns its own cache, so cached ids are never
shared across tokenizers.
- MiniMax-M3 has its own input processor (`MiniMaxM3VLInputProcessor`) and wires the
same cache, with the same eligibility rules, into its text-only path; requests with
images or videos take the HF processor path unchanged.
- The cache requires a fast (Rust-backed) tokenizer, because it relies on
character offsets. With a slow tokenizer the cache is disabled with a warning.
- The cache is used only when the tokenizer would be called exactly as the
Expand All @@ -68,8 +75,7 @@ feature you turned on.
## Measured effect

An A/B on GLM-5.2 (GB300, disaggregated, matched pair, 3600 s, ~29.6k requests
per arm, 0.34% error rate in both, identical configuration except the
environment variable):
per arm, identical configuration except `enable_tokenization_cache`):

| Metric | OFF | ON | Delta |
|---|---:|---:|---:|
Expand Down
2 changes: 1 addition & 1 deletion docs/source/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,6 @@ Welcome to TensorRT LLM's Documentation!
examples/dynamo_k8s_example.rst
deployment-guide/index.rst
deployment-guide/configuring-cpu-affinity.md
deployment-guide/prefix-tokenization-cache.md

.. toctree::
:maxdepth: 2
Expand Down Expand Up @@ -89,6 +88,7 @@ Welcome to TensorRT LLM's Documentation!
features/helix.md
features/kv-cache-connector.md
features/sparse-attention.md
features/prefix-tokenization-cache.md


.. toctree::
Expand Down
59 changes: 15 additions & 44 deletions tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py
Original file line number Diff line number Diff line change
Expand Up @@ -1737,15 +1737,19 @@ class MiniMaxM3VLInputProcessor:
``MINIMAX_M3_VL_VISION_END_TOKEN`` above, resolved via the tokenizer).
"""

supports_tokenization_cache = True

def __init__(
self,
model_path: str,
config: Any,
tokenizer: Any = None,
trust_remote_code: bool = True,
enable_tokenization_cache: bool = False,
**kwargs: Any,
):
from tensorrt_llm.inputs.registry import BaseMultimodalInputProcessor
from tensorrt_llm.logger import logger

BaseMultimodalInputProcessor.__init__(
self,
Expand All @@ -1770,10 +1774,15 @@ def __init__(
use_fast=self._use_fast,
trust_remote_code=trust_remote_code,
)
# Opt-in prefix-tokenization cache (TLLM_PREFIX_TOKEN_CACHE=1), shared with
# DefaultInputProcessor, which MiniMax-M3 text-only chat prompts never reach:
# this processor tokenizes them through the HF MiniMaxVLProcessor instead.
self._prefix_token_cache = self._create_prefix_token_cache()
# The HF processor tokenizes with add_special_tokens=True and the cache with False,
# so their ids match only if the tokenizer adds no special tokens.
if enable_tokenization_cache and self._processor.tokenizer.num_special_tokens_to_add() != 0:
logger.warning(
"enable_tokenization_cache is ignored: the MiniMax-M3 tokenizer adds "
"special tokens, so cached ids would differ from the HF processor's."
)
enable_tokenization_cache = False
self._init_tokenization_cache(enable_tokenization_cache, self._processor.tokenizer)
text_cfg = getattr(config, "text_config", None)
if isinstance(text_cfg, dict):
self._dtype = getattr(text_cfg, "torch_dtype", torch.bfloat16)
Expand All @@ -1800,36 +1809,6 @@ def __init__(
self._vision_start_token_id = self._resolve_token_id(MINIMAX_M3_VL_VISION_START_TOKEN)
self._vision_end_token_id = self._resolve_token_id(MINIMAX_M3_VL_VISION_END_TOKEN)

def _create_prefix_token_cache(self):
"""Prefix-tokenization cache for text-only prompts, or None when disabled or not exact.

Driven by the HF processor's own tokenizer, so cached ids come from the same tokenizer
as the processor path. The cache encodes with ``add_special_tokens=False``; it is used
only if the processor yields the same ids for a probe prompt, i.e. adds no special
tokens on this checkpoint. A probe failure disables the cache instead of failing
model loading.
"""
from tensorrt_llm.inputs.prefix_token_cache import create_prefix_token_cache
from tensorrt_llm.logger import logger

tokenizer = self._processor.tokenizer
cache = create_prefix_token_cache(tokenizer)
if cache is None:
return None
probe = "]~b]user\nhello 你好 <mm:think>x</mm:think> ]<]minimax[>[<tool_call>[e~[\n]~b]ai\n"
try:
ids = self._processor(text=[probe], return_tensors="pt")["input_ids"]
expected = tokenizer(probe, add_special_tokens=False)["input_ids"]
same = ids[0].tolist() == list(expected)
except Exception as e: # a cache problem must never fail model loading
same = False
logger.warning(f"MiniMax-M3 prefix token cache probe failed: {e!r}")
if not same:
logger.warning(
"Prefix token cache disabled for MiniMax-M3: HF processor and tokenizer ids differ"
)
return cache if same else None

def _resolve_token_id(self, token: str) -> int:
"""Look ``token`` up via the tokenizer. Raise if the tokenizer has
no ``convert_tokens_to_ids`` or maps the token to its unk id."""
Expand Down Expand Up @@ -2084,16 +2063,8 @@ def call_with_text_prompt(
templated_text = "\n".join(explicit)
else:
templated_text = text_prompt or ""

# Text-only fast path: splice cached prefix ids (opt-in, exact; see _create_prefix_token_cache).
if self._prefix_token_cache is not None and not images and not videos:
# Same eligibility as DefaultInputProcessor: no special tokens added, no truncation.
eligible = sampling_params is None or (
not sampling_params.add_special_tokens
and sampling_params.truncate_prompt_tokens is None
)
if eligible:
ids = self._prefix_token_cache.encode(self._processor.tokenizer, templated_text)
ids = self._encode_with_tokenization_cache(templated_text, sampling_params)
if ids is not None:
return ids, {"multimodal_data": {}}

# Run the HF processor. ``return_tensors='pt'`` yields tensors
Expand Down
27 changes: 11 additions & 16 deletions tensorrt_llm/inputs/prefix_token_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,10 @@
equal the cached ids over the same span. Otherwise the prompt is tokenized in
full.

Off by default; enable with ``TLLM_PREFIX_TOKEN_CACHE=1``. Each
``DefaultInputProcessor`` owns its own cache, so cached ids are never shared
across tokenizers. Environment variables, eligibility rules, and measured
effect are documented in ``docs/source/deployment-guide/prefix-tokenization-cache.md``.
Off by default; enable with the ``enable_tokenization_cache`` LLM argument.
Each ``DefaultInputProcessor`` owns its own cache, so cached ids are never
shared across tokenizers. Sizing knobs, eligibility rules, and measured effect
are documented in ``docs/source/features/prefix-tokenization-cache.md``.
"""

from __future__ import annotations
Expand All @@ -35,10 +35,8 @@
"PrefixTokenCache",
"PrefixTokenCacheConfig",
"create_prefix_token_cache",
"prefix_cache_enabled",
]

ENABLE_ENV_VAR = "TLLM_PREFIX_TOKEN_CACHE"
# Log the hit rate every this many cache-eligible requests, so an operator can
# see whether the cache is doing anything without reading its counters.
LOG_EVERY_REQUESTS = 1000
Expand All @@ -52,10 +50,6 @@
PROBE_CHARS = 256


def prefix_cache_enabled() -> bool:
return os.environ.get(ENABLE_ENV_VAR, "0") == "1"


class OffsetTokenizer(Protocol):
"""The subset of the HF fast-tokenizer call interface the cache relies on."""

Expand Down Expand Up @@ -278,17 +272,18 @@ def _remove(self, eid: int) -> None:


def create_prefix_token_cache(tokenizer: object) -> PrefixTokenCache | None:
"""Return a cache for ``tokenizer`` if the feature is enabled and usable.
"""Return a cache for ``tokenizer`` if the tokenizer can support one.

Returns None when the feature is off or the tokenizer cannot report
offsets. Raises ``ValueError`` for a malformed env override.
Returns None when there is no tokenizer or it cannot report offsets.
Raises ``ValueError`` for a malformed env override.
"""
if not prefix_cache_enabled() or tokenizer is None:
if tokenizer is None:
return None
if not getattr(tokenizer, "is_fast", False):
logger.warning(
f"{ENABLE_ENV_VAR} is set but the tokenizer is not a fast tokenizer, "
"which the prefix token cache needs for offset mappings; disabling it."
"enable_tokenization_cache is set but the tokenizer is not a fast "
"tokenizer, which the prefix token cache needs for offset mappings; "
"disabling it."
)
return None
config = PrefixTokenCacheConfig.from_env()
Expand Down
47 changes: 43 additions & 4 deletions tensorrt_llm/inputs/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,13 +86,15 @@ def __init__(self,
model_path,
config,
tokenizer,
trust_remote_code: bool = True) -> None:
trust_remote_code: bool = True,
enable_tokenization_cache: bool = False) -> None:
self.tokenizer = tokenizer
self.config = config
self.model_path = model_path
self.multimodal_hashing_supported = None
# Opt-in via TLLM_PREFIX_TOKEN_CACHE=1; None when disabled.
self._prefix_token_cache = create_prefix_token_cache(tokenizer)
# None when disabled or the tokenizer cannot support the cache.
self._prefix_token_cache = (create_prefix_token_cache(tokenizer)
if enable_tokenization_cache else None)

def __call__(
self, inputs: TextPrompt, sampling_params: SamplingParams
Expand Down Expand Up @@ -189,6 +191,10 @@ class BaseMultimodalInputProcessor(ABC):
# inputs to `call_with_token_ids` instead of detokenizing upstream.
supports_token_id_mm_expansion: ClassVar[bool] = False

# Whether the subclass takes `enable_tokenization_cache` in `__init__` and
# uses `_init_tokenization_cache` and `_encode_with_tokenization_cache`.
supports_tokenization_cache: ClassVar[bool] = False

def __init__(self,
model_path,
config,
Expand All @@ -203,6 +209,27 @@ def __init__(self,
self._trust_remote_code = trust_remote_code
self._multimodal_hashing_supported: Optional[bool] = None

def _init_tokenization_cache(self, enable_tokenization_cache: bool,
tokenizer: PreTrainedTokenizerBase) -> None:
"""Build the cache on `tokenizer` if `enable_tokenization_cache`.

Enable it only if this processor's ids for a text-only prompt equal
`tokenizer(prompt, add_special_tokens=False)`.
"""
self._prefix_token_cache_tokenizer = tokenizer
self._prefix_token_cache = (create_prefix_token_cache(tokenizer)
if enable_tokenization_cache else None)

def _encode_with_tokenization_cache(
self, prompt: str,
sampling_params: Optional[SamplingParams]) -> Optional[List[int]]:
if (self._prefix_token_cache is None or sampling_params is None
or sampling_params.add_special_tokens
or sampling_params.truncate_prompt_tokens is not None):
return None
return self._prefix_token_cache.encode(
self._prefix_token_cache_tokenizer, prompt)

def attach_multimodal_embeddings(
self,
inputs: TextPrompt,
Expand Down Expand Up @@ -921,6 +948,7 @@ def create_input_processor(
tokenizer,
checkpoint_format: Optional[str] = "HF",
trust_remote_code: bool = True,
enable_tokenization_cache: bool = False,
**kwargs,
) -> Union[InputProcessor, BaseMultimodalInputProcessor]:
"""Create an input processor for a specific model.
Expand All @@ -932,6 +960,10 @@ def create_input_processor(
config loading; any other value skips HF config loading. Default is "HF".
trust_remote_code: Whether Hugging Face config/processor loading may run
model-provided Python code.
enable_tokenization_cache: Whether the ``DefaultInputProcessor`` caches
the tokenization of recent prompts. Ignored for model-specific
(multimodal) input processors unless they set
``supports_tokenization_cache``.
**kwargs: Additional arguments passed to input processor constructors
(e.g., video_pruning_rate for multimodal models).

Expand Down Expand Up @@ -978,13 +1010,20 @@ def create_input_processor(
logger.info("Unregistered model, using DefaultInputProcessor")
input_processor_cls = None
if input_processor_cls is not None:
if (issubclass(input_processor_cls, BaseMultimodalInputProcessor)
and input_processor_cls.supports_tokenization_cache):
kwargs["enable_tokenization_cache"] = enable_tokenization_cache
return input_processor_cls(model_path_or_dir,
config,
tokenizer,
trust_remote_code=trust_remote_code,
**kwargs)

return DefaultInputProcessor(None, None, tokenizer)
return DefaultInputProcessor(
None,
None,
tokenizer,
enable_tokenization_cache=enable_tokenization_cache)


def _mm_data_to_counts(mm_data: Dict[str, Any]) -> Dict[str, int]:
Expand Down
1 change: 1 addition & 0 deletions tensorrt_llm/llmapi/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1572,6 +1572,7 @@ def _build_model(self):
self.tokenizer,
checkpoint_format,
trust_remote_code=self.args.trust_remote_code,
enable_tokenization_cache=self.args.enable_tokenization_cache,
**input_processor_kwargs)
self._tokenizer = self.input_processor.tokenizer

Expand Down
6 changes: 6 additions & 0 deletions tensorrt_llm/llmapi/llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -5267,6 +5267,12 @@ def validate_encoder_runtime_sizes(cls, v: Optional[int]) -> Optional[int]:
description="Enable iteration performance statistics.",
status="prototype")

enable_tokenization_cache: bool = Field(
default=False,
description=
"Cache the tokenization of recent prompts so that a prompt extending a cached one only tokenizes its tail, which speeds up multi-turn serving with long prompts. Requires a fast tokenizer and applies only to prompts tokenized with add_special_tokens=False and no truncation. The output is identical to tokenizing the whole prompt.",
status="prototype")

enable_iter_req_stats: bool = Field(
default=False,
description=
Expand Down
7 changes: 7 additions & 0 deletions tensorrt_llm/usage/llm_args_golden_manifest.json
Original file line number Diff line number Diff line change
Expand Up @@ -477,6 +477,13 @@
"kind": "value",
"path": "enable_speculative_beam_history_d2h"
},
{
"allowed_values": [],
"annotation": "<class 'bool'>",
"converter": "",
"kind": "value",
"path": "enable_tokenization_cache"
},
{
"allowed_values": [],
"annotation": "<class 'bool'>",
Expand Down
1 change: 0 additions & 1 deletion tests/integration/test_lists/test-db/l0_a10.yml
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,6 @@ l0_a10:
- unittest/inputs/test_multimodal.py
- unittest/inputs/test_multimodal_input_processor.py
- unittest/inputs/test_prefix_token_cache.py
- unittest/_torch/modeling/test_minimaxm3_vl_prefix_token_cache.py
- unittest/inputs/test_video_decode.py
- unittest/others/test_convert_utils.py
- unittest/others/test_lora_manager.py
Expand Down
Loading
Loading