From 69857bfc96ccd70699d048f6162ba0f772d952c9 Mon Sep 17 00:00:00 2001 From: yechank <161688079+yechank-nvidia@users.noreply.github.com> Date: Fri, 25 Jul 2025 09:39:00 +0900 Subject: [PATCH 1/2] fix: support mixture of text & multimodal prompts Signed-off-by: yechank <161688079+yechank-nvidia@users.noreply.github.com> --- .../_torch/models/modeling_gemma3vl.py | 3 --- .../_torch/models/modeling_hyperclovax.py | 6 +----- tensorrt_llm/_torch/models/modeling_llama.py | 9 +++++---- .../_torch/models/modeling_llava_next.py | 19 ++++++++++--------- .../_torch/models/modeling_mistral.py | 10 +++------- tensorrt_llm/_torch/models/modeling_phi4mm.py | 12 +++++++----- tensorrt_llm/_torch/models/modeling_vila.py | 19 ++++++++++--------- 7 files changed, 36 insertions(+), 42 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_gemma3vl.py b/tensorrt_llm/_torch/models/modeling_gemma3vl.py index 07fb5b5417bb..671f33903587 100644 --- a/tensorrt_llm/_torch/models/modeling_gemma3vl.py +++ b/tensorrt_llm/_torch/models/modeling_gemma3vl.py @@ -187,9 +187,6 @@ def forward( multimodal_param.multimodal_data["image"]["pixel_values"] for multimodal_param in multimodal_params ] - assert pixel_values == [] or len( - pixel_values - ) == num_context_requests, "Number of multimodal features (if provided) should be equal to number of context requests" mm_embeds = [] mm_token_mask = None diff --git a/tensorrt_llm/_torch/models/modeling_hyperclovax.py b/tensorrt_llm/_torch/models/modeling_hyperclovax.py index 9f37759ba03b..56d56f244337 100644 --- a/tensorrt_llm/_torch/models/modeling_hyperclovax.py +++ b/tensorrt_llm/_torch/models/modeling_hyperclovax.py @@ -597,7 +597,7 @@ def _post_process(self, input_ids: torch.Tensor, preprocessed_image: dict[str, any] = None): if not preprocessed_image: - return input_ids + return input_ids[0] vision_query_lengths = preprocessed_image.get("vision_query_lengths", None) @@ -659,7 +659,6 @@ def _preprocess(self, text_prompt: dict[str, any], images: List[Any], mm_processor_kwargs: Dict[str, Any]): preprocessed_image = None - is_video_list = [False] * len(images) if images is not None: is_video_list = [False] * len(images) preprocessed_image = self.processor( @@ -1026,9 +1025,6 @@ def forward( multimodal_params = kwargs.get("multimodal_params", []) mm_embeds = [] if len(multimodal_params) > 0: - assert len(multimodal_params) == num_context_requests == len( - multimodal_params - ), f"Number of multimodal tensors ({len(multimodal_params)}) should be equal to number of context requests ({num_context_requests}) in the batch." if not DISAGG: mm_embeds = self.mm_encoder.forward(multimodal_params) else: diff --git a/tensorrt_llm/_torch/models/modeling_llama.py b/tensorrt_llm/_torch/models/modeling_llama.py index 4af9762d1808..8d3ee666b4e7 100644 --- a/tensorrt_llm/_torch/models/modeling_llama.py +++ b/tensorrt_llm/_torch/models/modeling_llama.py @@ -1060,13 +1060,14 @@ def forward( **kwargs, ) -> torch.Tensor: multimodal_params = kwargs.get("multimodal_params", []) - if multimodal_params: - mm_embed = [ + mm_embeds = [] + if len(multimodal_params) > 0: + mm_embeds = [ multimodal_param.multimodal_data["multimodal_embedding"] for multimodal_param in multimodal_params ] - input_ids, inputs_embeds = fuse_input_embeds( - self.model.embed_tokens, input_ids, mm_embed) + input_ids, inputs_embeds = fuse_input_embeds(self.model.embed_tokens, + input_ids, mm_embeds) return super().forward(attn_metadata, input_ids, position_ids, diff --git a/tensorrt_llm/_torch/models/modeling_llava_next.py b/tensorrt_llm/_torch/models/modeling_llava_next.py index 8af484ce1ab6..b85077f0d6a6 100644 --- a/tensorrt_llm/_torch/models/modeling_llava_next.py +++ b/tensorrt_llm/_torch/models/modeling_llava_next.py @@ -201,11 +201,13 @@ def __call__( ) -> Tuple[List[int], Optional[ExtraProcessedInputs]]: text_prompt, mm_data = inputs.get("prompt"), inputs.get( "multi_modal_data", {}) - assert 'image' in mm_data input_ids = self.tokenizer( text_prompt, return_tensors="pt").input_ids[0].to(self.device) + if not mm_data: + return input_ids.to(torch.int32).tolist(), {} + mm_tensor = self._preprocess(mm_data['image']) mm_features = torch.stack( [self._process(tensor) for tensor in mm_tensor]) @@ -274,16 +276,15 @@ def forward( logger.debug(f"{num_context_requests=}, {num_generation_requests=}") multimodal_params = kwargs.get("multimodal_params", []) - mm_embed = [ - multimodal_param.multimodal_data["multimodal_embedding"] - for multimodal_param in multimodal_params - ] - assert mm_embed == [] or len( - mm_embed - ) == num_context_requests, "Number of multimodal features (if provided) should be equal to number of context requests" + mm_embeds = [] + if len(multimodal_params) > 0: + mm_embeds = [ + multimodal_param.multimodal_data["multimodal_embedding"] + for multimodal_param in multimodal_params + ] input_ids, inputs_embeds = fuse_input_embeds( - self.llm.model.embed_tokens, input_ids, mm_embed) + self.llm.model.embed_tokens, input_ids, mm_embeds) logits = self.llm.forward(attn_metadata, input_ids, position_ids, inputs_embeds, return_context_logits) return logits diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index 785b93fdb67f..f10ea5368c96 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -354,13 +354,9 @@ def forward( logger.debug(f"{num_context_requests=}, {num_generation_requests=}") multimodal_params = kwargs.get("multimodal_params", []) - image_features = [] + mm_embeds = [] multimodal_params_len = len(multimodal_params) if multimodal_params_len > 0: - if multimodal_params_len != num_context_requests: - raise RuntimeError( - f"Number of multimodal tensors ({multimodal_params_len}) should be equal to number of " - f"context requests ({num_context_requests}) in the batch.") pixel_values = [ x.multimodal_data["image"]["pixel_values"] for x in multimodal_params @@ -377,7 +373,7 @@ def forward( f"({multimodal_params_len}).") batched_pixel_values, batched_image_sizes = self._batch_pixel_values( pixel_values=pixel_values, image_sizes=image_sizes) - image_features = [ + mm_embeds = [ self._get_image_features(pixel_values=batched_pixel_values, image_sizes=batched_image_sizes) ] @@ -385,7 +381,7 @@ def forward( input_ids, inputs_embeds = fuse_input_embeds( embedding_layer=self.llm.model.embed_tokens, input_ids=input_ids, - mm_embeds=image_features, + mm_embeds=mm_embeds, mm_token_ids=self._image_token_ids, ) diff --git a/tensorrt_llm/_torch/models/modeling_phi4mm.py b/tensorrt_llm/_torch/models/modeling_phi4mm.py index 8c8982f6e0be..f0fddbef7b86 100644 --- a/tensorrt_llm/_torch/models/modeling_phi4mm.py +++ b/tensorrt_llm/_torch/models/modeling_phi4mm.py @@ -215,14 +215,16 @@ def forward( ) multimodal_params = kwargs.get("multimodal_params", []) - mm_embedding = [ - multimodal_param.multimodal_data["multimodal_embedding"] - for multimodal_param in multimodal_params - ] + mm_embeds = [] + if len(multimodal_params) > 0: + mm_embeds = [ + multimodal_param.multimodal_data["multimodal_embedding"] + for multimodal_param in multimodal_params + ] input_ids, input_embeds = fuse_input_embeds( self.llm.model.embed_tokens, input_ids, - mm_embedding, + mm_embeds, mm_token_ids=self.MM_TOKEN_IDS, ) diff --git a/tensorrt_llm/_torch/models/modeling_vila.py b/tensorrt_llm/_torch/models/modeling_vila.py index c27a88abf5f0..99820c1954c2 100644 --- a/tensorrt_llm/_torch/models/modeling_vila.py +++ b/tensorrt_llm/_torch/models/modeling_vila.py @@ -1102,6 +1102,9 @@ def __call__( input_ids = self.tokenizer( text_prompt, return_tensors="pt").input_ids[0].to(self.device) + if not mm_data: + return input_ids.to(torch.int32).tolist(), {} + mm_tensor, block_sizes = self._preprocess( mm_data, mm_processor_kwargs, use_fast=True ) # use_fast uses Pytorch GPU preprocessing, otherwise uses PIL CPU preprocessing @@ -1164,17 +1167,15 @@ def forward( num_context_requests, num_generation_requests = attn_metadata.num_contexts, attn_metadata.num_generations multimodal_params = kwargs.get("multimodal_params", []) - mm_embed = [ - multimodal_param.multimodal_data["multimodal_embedding"] - for multimodal_param in multimodal_params - ] - - assert mm_embed == [] or len( - mm_embed - ) == num_context_requests, "Number of multimodal features (if provided) should be equal to number of context requests" + mm_embeds = [] + if len(multimodal_params) > 0: + mm_embeds = [ + multimodal_param.multimodal_data["multimodal_embedding"] + for multimodal_param in multimodal_params + ] input_ids, inputs_embeds = fuse_input_embeds( - self.llm.model.embed_tokens, input_ids, mm_embed) + self.llm.model.embed_tokens, input_ids, mm_embeds) logits = self.llm.forward(attn_metadata=attn_metadata, input_ids=input_ids, position_ids=position_ids, From 0d2ee11b76ee65c0aeeeff2233fdd83922c24f29 Mon Sep 17 00:00:00 2001 From: yechank <161688079+yechank-nvidia@users.noreply.github.com> Date: Tue, 29 Jul 2025 13:36:29 +0900 Subject: [PATCH 2/2] add test mixture_text_image Signed-off-by: yechank <161688079+yechank-nvidia@users.noreply.github.com> --- examples/llm-api/quickstart_multimodal.py | 43 +++++++++++++++---- tensorrt_llm/inputs/utils.py | 35 +++++++++------ tests/integration/defs/test_e2e.py | 15 ++++++- .../test_lists/qa/examples_test_list.txt | 1 + .../test_lists/qa/llm_sanity_test.txt | 1 + .../test_lists/test-db/l0_h100.yml | 1 + 6 files changed, 73 insertions(+), 23 deletions(-) diff --git a/examples/llm-api/quickstart_multimodal.py b/examples/llm-api/quickstart_multimodal.py index c4d40655d3dc..1f45da90eb42 100644 --- a/examples/llm-api/quickstart_multimodal.py +++ b/examples/llm-api/quickstart_multimodal.py @@ -55,7 +55,26 @@ "Describe the scene in the image briefly.", "", ] - } + }, + "multiple_image": { + "media": [ + "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/inpaint.png", + "https://huggingface.co/datasets/Sayali9141/traffic_signal_images/resolve/main/61.jpg", + ], + "prompt": ["Describe the difference between the two images."], + }, + "mixture_text_image": { + "media": [ + [], + [ + "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/inpaint.png" + ], + ], + "prompt": [ + "Who invented the internet?", + "Describe the scene in the image briefly.", + ], + }, } @@ -66,7 +85,10 @@ def add_multimodal_args(parser): help="Model type.") parser.add_argument("--modality", type=str, - choices=["image", "video", "audio", "image_audio"], + choices=[ + "image", "video", "audio", "image_audio", + "multiple_image", "mixture_text_image" + ], default="image", help="Media type.") parser.add_argument("--media", @@ -82,6 +104,10 @@ def add_multimodal_args(parser): choices=["pt", "pil"], default="pt", help="The format of the image.") + parser.add_argument("--device", + type=str, + default="cpu", + help="The device to have the input on.") return parser @@ -114,11 +140,6 @@ def parse_arguments(): def main(): args = parse_arguments() - # set prompts and media to example prompts and images if they are not provided - if args.prompt is None: - args.prompt = example_medias_and_prompts[args.modality]["prompt"] - if args.media is None: - args.media = example_medias_and_prompts[args.modality]["media"] lora_config = None if args.load_lora: @@ -138,7 +159,11 @@ def main(): open(os.path.join(llm._hf_model_dir, 'config.json')))['model_type'] assert model_type in ALL_SUPPORTED_MULTIMODAL_MODELS, f"Unsupported model_type: {model_type}" - device = "cpu" + # set prompts and media to example prompts and images if they are not provided + if args.prompt is None: + args.prompt = example_medias_and_prompts[args.modality]["prompt"] + if args.media is None: + args.media = example_medias_and_prompts[args.modality]["media"] inputs = default_multimodal_input_loader(tokenizer=llm.tokenizer, model_dir=llm._hf_model_dir, model_type=model_type, @@ -147,7 +172,7 @@ def main(): media=args.media, image_data_format=image_format, num_frames=args.num_frames, - device=device) + device=args.device) lora_request = None if args.load_lora: diff --git a/tensorrt_llm/inputs/utils.py b/tensorrt_llm/inputs/utils.py index a4bf8570d0ae..a6b984b330e5 100644 --- a/tensorrt_llm/inputs/utils.py +++ b/tensorrt_llm/inputs/utils.py @@ -487,9 +487,9 @@ def convert_to_conversation_message(prompt: str, media: Union[str, modality: str) -> ConversationMessage: if isinstance(media, str): media = [media] - if modality == "image": + if modality in ["image", "multiple_image"]: mm_data = [ - MultimodalData(modality=modality, + MultimodalData(modality="image", data=load_image(i, format=image_data_format, device=device)) for i in media @@ -530,6 +530,15 @@ def convert_to_conversation_message(prompt: str, media: Union[str, if _modal is None: raise ValueError(f"Unknown matching modality: {modality}") mm_data.append(MultimodalData(modality=_modal, data=data)) + elif modality == "mixture_text_image": + mm_data = [] + for m in media: + if m: + mm_data.append( + MultimodalData(modality="image", + data=load_image(m, + format=image_data_format, + device=device))) else: raise ValueError(f"Unknown modality: {modality}") return ConversationMessage(role="user", content=prompt, media=mm_data) @@ -561,16 +570,16 @@ def convert_to_conversation_message(prompt: str, media: Union[str, if mm_placeholder_counts: conv["content"] = add_multimodal_placeholders( model_type, conv["content"], mm_placeholder_counts) - prompt = apply_chat_template( - model_type=model_type, - tokenizer=tokenizer, - processor=processor, - conversation=[conv], - add_generation_prompt=True, - mm_placeholder_counts=mm_placeholder_counts) - inputs.append({ - "prompt": prompt, - "multi_modal_data": mm_data_tracker.retrieve_all_sync() - }) + prompt = apply_chat_template( + model_type=model_type, + tokenizer=tokenizer, + processor=processor, + conversation=[conv], + add_generation_prompt=True, + mm_placeholder_counts=mm_placeholder_counts) + input = {"prompt": prompt} + if mm_placeholder_counts: + input["multi_modal_data"] = mm_data_tracker.retrieve_all_sync() + inputs.append(input) return inputs diff --git a/tests/integration/defs/test_e2e.py b/tests/integration/defs/test_e2e.py index 82d828961b1f..566cf25b6c9a 100644 --- a/tests/integration/defs/test_e2e.py +++ b/tests/integration/defs/test_e2e.py @@ -1939,7 +1939,7 @@ def test_ptp_quickstart_advanced_mixed_precision(llm_root, llm_venv): @pytest.mark.parametrize("use_cuda_graph", [False, True]) -@pytest.mark.parametrize("modality", ["image", "video"]) +@pytest.mark.parametrize("modality", ["image", "video", "mixture_text_image"]) @pytest.mark.parametrize("model_name,model_path", [ ("NVILA-8B-FP16", "vila/NVILA-8B"), ("NVILA-15B-FP16", "NVILA-15B"), @@ -1987,6 +1987,16 @@ def test_ptp_quickstart_multimodal(llm_root, llm_venv, model_name, model_path, str(test_data_root / "world.mp4"), ], }, + "mixture_text_image": { + "prompt": [ + "Who invented the internet?", + "Describe the scene in the image briefly.", + ], + "media": [ + [], + [str(test_data_root / "inpaint.png")], + ], + } } expected_keywords = { @@ -2042,6 +2052,9 @@ def test_ptp_quickstart_multimodal(llm_root, llm_venv, model_name, model_path, ["scenic", "rock", "landscape", "snow", "altitude"], ["highway", "traffic", "directions", "lanes", "Jurong"], ], + "mixture_text_image": + [["invention", "person", "scientists", "Lick", "engineers"], + ["landscape", "dome", "yosemite", "altitude", "scattered"]] }, "gemma-3-27b-it": { "image": [ diff --git a/tests/integration/test_lists/qa/examples_test_list.txt b/tests/integration/test_lists/qa/examples_test_list.txt index 4d9ec1c5deba..eaebfb67c576 100644 --- a/tests/integration/test_lists/qa/examples_test_list.txt +++ b/tests/integration/test_lists/qa/examples_test_list.txt @@ -536,6 +536,7 @@ test_e2e.py::test_ptp_quickstart_multimodal[qwen2.5-vl-7b-instruct-Qwen2.5-VL-7B test_e2e.py::test_ptp_quickstart_multimodal[qwen2.5-vl-7b-instruct-Qwen2.5-VL-7B-Instruct-video-True] test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-image-True] test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-image-False] +test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-mixture_text_image-True] test_e2e.py::test_ptp_quickstart_multimodal[gemma-3-27b-it-gemma/gemma-3-27b-it-image-False] test_e2e.py::test_ptp_quickstart_multimodal[gemma-3-27b-it-gemma/gemma-3-27b-it-image-True] test_e2e.py::test_ptp_quickstart_multimodal_phi4mm[audio] diff --git a/tests/integration/test_lists/qa/llm_sanity_test.txt b/tests/integration/test_lists/qa/llm_sanity_test.txt index 79c51f6d1873..8635973b31dc 100644 --- a/tests/integration/test_lists/qa/llm_sanity_test.txt +++ b/tests/integration/test_lists/qa/llm_sanity_test.txt @@ -102,6 +102,7 @@ test_e2e.py::test_ptp_quickstart_bert[VANILLA-BertForSequenceClassification-bert test_e2e.py::test_ptp_quickstart_multimodal[llava-v1.6-mistral-7b-llava-v1.6-mistral-7b-hf-image-False] test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-image-False] test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-image-True] +test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-mixture_text_image-True] test_e2e.py::test_ptp_quickstart_multimodal[NVILA-8B-FP16-vila/NVILA-8B-image-False] test_e2e.py::test_ptp_quickstart_multimodal[NVILA-8B-FP16-vila/NVILA-8B-video-False] test_e2e.py::test_ptp_quickstart_multimodal[qwen2-vl-7b-instruct-Qwen2-VL-7B-Instruct-image-False] diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 56eb1cf50c14..957c6697c3a1 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -196,6 +196,7 @@ l0_h100: - accuracy/test_llm_api_pytorch.py::TestQwen3_30B_A3B::test_fp8_block_scales[latency] - accuracy/test_llm_api_pytorch.py::TestLlama3_1_8BInstruct::test_guided_decoding[llguidance] - test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-image-True] + - test_e2e.py::test_ptp_quickstart_multimodal[mistral-small-3.1-24b-instruct-Mistral-Small-3.1-24B-Instruct-2503-mixture_text_image-True] - condition: ranges: system_gpu_count: