diff --git a/docs/source/deployment-guide/prefix-tokenization-cache.md b/docs/source/features/prefix-tokenization-cache.md
similarity index 87%
rename from docs/source/deployment-guide/prefix-tokenization-cache.md
rename to docs/source/features/prefix-tokenization-cache.md
index 764b7ba0725c..4855e01901f1 100644
--- a/docs/source/deployment-guide/prefix-tokenization-cache.md
+++ b/docs/source/features/prefix-tokenization-cache.md
@@ -1,4 +1,4 @@
-# Prefix-Tokenization Cache in TensorRT LLM
+# Prefix-Tokenization Cache
## Overview
@@ -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. |
@@ -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
@@ -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 |
|---|---:|---:|---:|
diff --git a/docs/source/index.rst b/docs/source/index.rst
index 252942adfbec..9820d3770d86 100644
--- a/docs/source/index.rst
+++ b/docs/source/index.rst
@@ -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
@@ -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::
diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py b/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py
index bd5fd21d6e26..024e7ae68b45 100644
--- a/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py
+++ b/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py
@@ -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,
@@ -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)
@@ -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 你好 x ]<]minimax[>[[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."""
@@ -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
diff --git a/tensorrt_llm/inputs/prefix_token_cache.py b/tensorrt_llm/inputs/prefix_token_cache.py
index 4a68261f2bf7..666e15f3b42d 100644
--- a/tensorrt_llm/inputs/prefix_token_cache.py
+++ b/tensorrt_llm/inputs/prefix_token_cache.py
@@ -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
@@ -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
@@ -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."""
@@ -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()
diff --git a/tensorrt_llm/inputs/registry.py b/tensorrt_llm/inputs/registry.py
index 365b3cf65914..de8a2ca7ac78 100644
--- a/tensorrt_llm/inputs/registry.py
+++ b/tensorrt_llm/inputs/registry.py
@@ -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
@@ -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,
@@ -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,
@@ -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.
@@ -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).
@@ -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]:
diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py
index 577717586748..5d170be95559 100644
--- a/tensorrt_llm/llmapi/llm.py
+++ b/tensorrt_llm/llmapi/llm.py
@@ -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
diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py
index 6d22c1262b26..b13b1171d6bb 100644
--- a/tensorrt_llm/llmapi/llm_args.py
+++ b/tensorrt_llm/llmapi/llm_args.py
@@ -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=
diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json
index 18b784e62495..70eb0a822f4e 100644
--- a/tensorrt_llm/usage/llm_args_golden_manifest.json
+++ b/tensorrt_llm/usage/llm_args_golden_manifest.json
@@ -477,6 +477,13 @@
"kind": "value",
"path": "enable_speculative_beam_history_d2h"
},
+ {
+ "allowed_values": [],
+ "annotation": "",
+ "converter": "",
+ "kind": "value",
+ "path": "enable_tokenization_cache"
+ },
{
"allowed_values": [],
"annotation": "",
diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml
index f70c7b5033dc..8863a599de05 100644
--- a/tests/integration/test_lists/test-db/l0_a10.yml
+++ b/tests/integration/test_lists/test-db/l0_a10.yml
@@ -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
diff --git a/tests/unittest/_torch/modeling/test_minimaxm3_vl_prefix_token_cache.py b/tests/unittest/_torch/modeling/test_minimaxm3_vl_prefix_token_cache.py
deleted file mode 100644
index b525ff644122..000000000000
--- a/tests/unittest/_torch/modeling/test_minimaxm3_vl_prefix_token_cache.py
+++ /dev/null
@@ -1,162 +0,0 @@
-# SPDX-License-Identifier: Apache-2.0
-# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
-"""MiniMax-M3 VL input processor: text-only prompts take the opt-in prefix-tokenization cache.
-
-The processor is built without a checkpoint (``__new__`` + the attributes ``call_with_text_prompt`` reads), with a
-context-free stub tokenizer and a stub HF processor that tokenizes through it, so the test covers the wiring and the
-eligibility rules rather than the tokenizer itself (``tests/unittest/inputs/test_prefix_token_cache.py`` does that).
-"""
-
-from __future__ import annotations
-
-from typing import Any
-
-import pytest
-import torch
-
-from tensorrt_llm._torch.models.modeling_minimaxm3_vl import MiniMaxM3VLInputProcessor
-from tensorrt_llm.inputs.prefix_token_cache import ENABLE_ENV_VAR
-from tensorrt_llm.sampling_params import SamplingParams
-
-_MIN_CHARS = 16
-
-
-class _Tokenizer:
- """One id per space-separated word, character offsets: the fast-tokenizer call contract the cache relies on."""
-
- is_fast = True
-
- def __call__(
- self, text: str, *, add_special_tokens: bool, return_offsets_mapping: bool = False
- ) -> dict[str, Any]:
- ids, offsets, pos = [], [], 0
- for word in text.split(" "):
- ids.append(sum(map(ord, word)) % 50_000 + 1)
- offsets.append((pos, pos + len(word)))
- pos += len(word) + 1
- out: dict[str, Any] = {"input_ids": ids}
- if return_offsets_mapping:
- out["offset_mapping"] = offsets
- return out
-
- def apply_chat_template(self, messages: list[dict[str, Any]], **kwargs: Any) -> str:
- return " ".join(str(m["content"]) for m in messages)
-
-
-class _Processor:
- """Stands in for the HF ``MiniMaxVLProcessor``: tokenizes ``text`` through the tokenizer, counts its calls."""
-
- def __init__(self, tokenizer: _Tokenizer, extra_leading_ids: tuple[int, ...] = ()) -> None:
- self.tokenizer = tokenizer
- self._extra = list(extra_leading_ids)
- self.calls = 0
-
- def __call__(
- self,
- text: list[str],
- images: Any = None,
- videos: Any = None,
- return_tensors: str | None = None,
- ) -> dict[str, Any]:
- self.calls += 1
- ids = self._extra + self.tokenizer(text[0], add_special_tokens=False)["input_ids"]
- out: dict[str, Any] = {"input_ids": torch.tensor([ids], dtype=torch.int32)}
- if images:
- out["pixel_values"] = torch.zeros(1, 4, dtype=torch.float32)
- out["image_grid_thw"] = torch.tensor([[1, 2, 2]])
- return out
-
-
-def _processor(
- monkeypatch: pytest.MonkeyPatch,
- *,
- enabled: bool = True,
- extra_leading_ids: tuple[int, ...] = (),
-) -> MiniMaxM3VLInputProcessor:
- if enabled:
- monkeypatch.setenv(ENABLE_ENV_VAR, "1")
- monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_MIN_CHARS", str(_MIN_CHARS))
- else:
- monkeypatch.delenv(ENABLE_ENV_VAR, raising=False)
- ip = MiniMaxM3VLInputProcessor.__new__(MiniMaxM3VLInputProcessor)
- ip._tokenizer = _Tokenizer()
- ip._processor = _Processor(ip._tokenizer, extra_leading_ids)
- ip._dtype = torch.bfloat16
- ip._prefix_token_cache = ip._create_prefix_token_cache()
- return ip
-
-
-def _turns(n: int) -> list[str]:
- """Chat-shaped prompts: each turn's prompt is the previous prompt (which ends with the assistant generation
- header) followed by the recorded reply, the next user message and the header again, so every prompt extends
- the previous one byte-for-byte, as rendered chat templates do."""
- words = [f"w{i}" for i in range(100 * n)]
- texts, history = [], ""
- for k in range(n):
- history += " ".join(words[100 * k : 100 * (k + 1)]) + " ]~b]ai\n"
- texts.append(history)
- history += " ok done ]~b]user\n "
- return texts
-
-
-def test_text_only_prompts_use_the_cache_and_match_the_processor(
- monkeypatch: pytest.MonkeyPatch,
-) -> None:
- ip = _processor(monkeypatch)
- assert ip._prefix_token_cache is not None
- probe_calls = ip._processor.calls # one processor call from the equivalence probe
- reference = _Processor(_Tokenizer())
- params = SamplingParams(add_special_tokens=False)
- for k, text in enumerate(_turns(4)):
- ids, extra = ip.call_with_text_prompt({"prompt": text}, params)
- assert ids == reference(text=[text])["input_ids"][0].tolist(), k
- assert extra == {"multimodal_data": {}}
- assert ip._processor.calls == probe_calls, "text-only prompts must not reach the HF processor"
- assert ip._prefix_token_cache.hits == 3 and ip._prefix_token_cache.misses == 1
-
-
-def test_disabled_by_default_takes_the_processor_path(monkeypatch: pytest.MonkeyPatch) -> None:
- ip = _processor(monkeypatch, enabled=False)
- assert ip._prefix_token_cache is None
- text = _turns(1)[0]
- ids, _ = ip.call_with_text_prompt({"prompt": text})
- assert ids == _Tokenizer()(text, add_special_tokens=False)["input_ids"]
- assert ip._processor.calls == 1
-
-
-def test_processor_that_adds_tokens_disables_the_cache(monkeypatch: pytest.MonkeyPatch) -> None:
- ip = _processor(monkeypatch, extra_leading_ids=(7,))
- assert ip._prefix_token_cache is None, (
- "a processor that prepends a BOS-like id is not equivalent to the tokenizer"
- )
- text = _turns(1)[0]
- ids, _ = ip.call_with_text_prompt({"prompt": text})
- assert ids[0] == 7
-
-
-def test_image_requests_bypass_the_cache(monkeypatch: pytest.MonkeyPatch) -> None:
- ip = _processor(monkeypatch)
- calls = ip._processor.calls
- ids, extra = ip.call_with_text_prompt(
- {"prompt": _turns(1)[0], "multi_modal_data": {"image": [object()]}}
- )
- assert ip._processor.calls == calls + 1
- assert "image" in extra["multimodal_data"]
- assert ip._prefix_token_cache.hits == 0 and ip._prefix_token_cache.misses == 0
-
-
-@pytest.mark.parametrize(
- "params",
- [
- SamplingParams(add_special_tokens=True),
- SamplingParams(add_special_tokens=False, truncate_prompt_tokens=8),
- ],
-)
-def test_requests_that_add_special_tokens_or_truncate_bypass_the_cache(
- monkeypatch: pytest.MonkeyPatch, params: SamplingParams
-) -> None:
- ip = _processor(monkeypatch)
- calls = ip._processor.calls
- ip.call_with_text_prompt({"prompt": _turns(1)[0]}, params)
- assert ip._processor.calls == calls + 1
- assert ip._prefix_token_cache.hits == 0 and ip._prefix_token_cache.misses == 0
diff --git a/tests/unittest/api_stability/references/llm.yaml b/tests/unittest/api_stability/references/llm.yaml
index f2aefd22f9c6..404170f0415b 100644
--- a/tests/unittest/api_stability/references/llm.yaml
+++ b/tests/unittest/api_stability/references/llm.yaml
@@ -179,6 +179,10 @@ methods:
annotation: bool
default: False
status: prototype
+ enable_tokenization_cache:
+ annotation: bool
+ default: False
+ status: prototype
enable_iter_req_stats:
annotation: bool
default: False
diff --git a/tests/unittest/inputs/test_prefix_token_cache.py b/tests/unittest/inputs/test_prefix_token_cache.py
index 7436c86b462b..474ff2aca44b 100644
--- a/tests/unittest/inputs/test_prefix_token_cache.py
+++ b/tests/unittest/inputs/test_prefix_token_cache.py
@@ -18,11 +18,9 @@
from tensorrt_llm.inputs import prefix_token_cache
from tensorrt_llm.inputs.prefix_token_cache import (
- ENABLE_ENV_VAR,
PrefixTokenCache,
PrefixTokenCacheConfig,
create_prefix_token_cache,
- prefix_cache_enabled,
)
from tensorrt_llm.inputs.registry import DefaultInputProcessor
from tensorrt_llm.sampling_params import SamplingParams
@@ -247,19 +245,6 @@ def test_error_disables_cache_and_still_returns_ids() -> None:
assert not cache._entries
-@pytest.mark.parametrize(
- "value,expected", [(None, False), ("0", False), ("", False), ("true", False), ("1", True)]
-)
-def test_enabled_only_for_exactly_one(
- monkeypatch: pytest.MonkeyPatch, value: str | None, expected: bool
-) -> None:
- if value is None:
- monkeypatch.delenv(ENABLE_ENV_VAR, raising=False)
- else:
- monkeypatch.setenv(ENABLE_ENV_VAR, value)
- assert prefix_cache_enabled() is expected
-
-
def test_config_from_env_reads_overrides(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_ENTRIES", "7")
monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_MIN_CHARS", "100")
@@ -278,17 +263,15 @@ def test_config_from_env_rejects_invalid_values(
PrefixTokenCacheConfig.from_env()
-def test_create_returns_none_for_slow_tokenizer(monkeypatch: pytest.MonkeyPatch) -> None:
+def test_create_returns_none_for_slow_tokenizer() -> None:
class Slow:
is_fast = False
- monkeypatch.setenv(ENABLE_ENV_VAR, "1")
assert create_prefix_token_cache(Slow()) is None
assert create_prefix_token_cache(None) is None
def test_create_raises_on_invalid_env(monkeypatch: pytest.MonkeyPatch) -> None:
- monkeypatch.setenv(ENABLE_ENV_VAR, "1")
monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_ENTRIES", "many")
with pytest.raises(ValueError):
create_prefix_token_cache(_MergeTokenizer())
@@ -298,12 +281,8 @@ def test_create_raises_on_invalid_env(monkeypatch: pytest.MonkeyPatch) -> None:
def _processor(monkeypatch: pytest.MonkeyPatch, enabled: bool = True) -> DefaultInputProcessor:
- if enabled:
- monkeypatch.setenv(ENABLE_ENV_VAR, "1")
- monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_MIN_CHARS", str(_MIN_CHARS))
- else:
- monkeypatch.delenv(ENABLE_ENV_VAR, raising=False)
- return DefaultInputProcessor(None, None, _MergeTokenizer())
+ monkeypatch.setenv("TLLM_PREFIX_TOKEN_CACHE_MIN_CHARS", str(_MIN_CHARS))
+ return DefaultInputProcessor(None, None, _MergeTokenizer(), enable_tokenization_cache=enabled)
def test_processor_uses_cache_for_eligible_prompts(monkeypatch: pytest.MonkeyPatch) -> None: