diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py b/tensorrt_llm/_torch/models/modeling_minimaxm3_vl.py index 32475c5fc1b7..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,6 +1774,15 @@ def __init__( use_fast=self._use_fast, trust_remote_code=trust_remote_code, ) + # 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) @@ -2050,6 +2063,9 @@ def call_with_text_prompt( templated_text = "\n".join(explicit) else: templated_text = text_prompt or "" + 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 # in the BatchFeature output. diff --git a/tensorrt_llm/inputs/registry.py b/tensorrt_llm/inputs/registry.py index 985075db25ff..7c2c2a1b0e30 100644 --- a/tensorrt_llm/inputs/registry.py +++ b/tensorrt_llm/inputs/registry.py @@ -251,6 +251,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 get_mm_encoder_item_metadata( self, prompt_token_ids: List[int], @@ -277,6 +281,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, @@ -1129,7 +1154,8 @@ def create_input_processor( model-provided Python code. enable_tokenization_cache: Whether the ``DefaultInputProcessor`` caches the tokenization of recent prompts. Ignored for model-specific - (multimodal) input processors. + (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). @@ -1177,6 +1203,9 @@ 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 # Input processors build an AutoTokenizer/AutoProcessor with # trust_remote_code; doing so copies the checkpoint's .py files # into the shared HF module cache non-atomically, and a rank that