From 6d29d8ab16d28702989af5ee4b518369b1384e7a Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Sun, 18 Jan 2026 14:11:17 +0800 Subject: [PATCH] [VLM][Reland] Refactor load_mm_data to improve performance (#16152) Co-authored-by: luoyuan.luo --- .../multimodal/processors/base_processor.py | 193 +++++++++++++++++- .../srt/multimodal/processors/minicpm.py | 1 + 2 files changed, 192 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 306ef5d19..5b6afb151 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -418,6 +418,47 @@ class BaseMultimodalProcessor(ABC): except Exception as e: raise RuntimeError(f"Error while loading data {data}: {e}") + def _submit_mm_data_loading_tasks_simple( + self, + data_list: Optional[list], + modality: Modality, + audio_sample_rate: Optional[int], + discard_alpha_channel: bool, + ) -> List[Tuple[Modality, int, concurrent.futures.Future]]: + """ + Simple version: For one modal data submit IO load task. + Return: + List[(modality, index_in_that_modality, future)] + """ + futures: List[Tuple[Modality, int, concurrent.futures.Future]] = [] + + if not data_list: + logger.debug( + "[_submit_mm_data_loading_tasks_simple] no data for modality=%s", + modality.name, + ) + return futures + + for idx, data in enumerate(data_list): + logger.debug( + "[_submit_mm_data_loading_tasks_simple] submit load task: " + "modality=%s, index=%d, data_type=%s", + modality.name, + idx, + type(data), + ) + future = self.io_executor.submit( + BaseMultimodalProcessor._load_single_item, + data, + modality, + None, # frame_count_limit: no consider for fast path + audio_sample_rate, + discard_alpha_channel, + ) + futures.append((modality, idx, future)) + + return futures + def submit_data_loading_tasks( self, text_parts: List[str], @@ -572,6 +613,156 @@ class BaseMultimodalProcessor(ABC): discard_alpha_channel: bool = True, audio_sample_rate: Optional[int] = None, ) -> BaseMultiModalProcessorOutput: + + BaseMultimodalProcessor.validate_mm_data(image_data, video_data, audio_data) + + multimodal_tokens_pattern = multimodal_tokens.get_combined_regex() + if isinstance(prompt, list) and return_text: + assert len(prompt) and isinstance(prompt[0], int) + prompt = self._processor.tokenizer.decode(prompt) + else: + prompt = prompt + + assert isinstance(prompt, str) + # split text into list of normal text and special tokens + text_parts = re.split(multimodal_tokens_pattern, prompt) + + cnt = {Modality.IMAGE: 0, Modality.VIDEO: 0, Modality.AUDIO: 0} + for text_part in text_parts: + modality = multimodal_tokens.get_modality_of_token(text_part) + if modality is not None: + cnt[modality] += 1 + + n_image = len(image_data) if image_data else 0 + n_video = len(video_data) if video_data else 0 + n_audio = len(audio_data) if audio_data else 0 + + # For MiniCPMO and MiniCPMV or multimodal_tokens not totally align, legacy show path + if ( + self.server_args.skip_tokenizer_init + or cnt[Modality.IMAGE] != n_image + or cnt[Modality.VIDEO] != n_video + or cnt[Modality.AUDIO] != n_audio + or getattr(self, "support_dynamic_frame_expansion", False) + ): + return self.legacy_load_mm_data( + prompt=prompt, + multimodal_tokens=multimodal_tokens, + image_data=image_data, + video_data=video_data, + audio_data=audio_data, + return_text=return_text, + discard_alpha_channel=discard_alpha_channel, + audio_sample_rate=audio_sample_rate, + ) + # For models other than MiniCPMO and MiniCPMV, + # totally align multimodal_tokens, fast path + return self.fast_load_mm_data( + prompt=prompt, + multimodal_tokens=multimodal_tokens, + image_data=image_data, + video_data=video_data, + audio_data=audio_data, + return_text=return_text, + discard_alpha_channel=discard_alpha_channel, + audio_sample_rate=audio_sample_rate, + ) + + def fast_load_mm_data( + self, + prompt: str, + multimodal_tokens: MultimodalSpecialTokens, + image_data: Optional[list] = None, + video_data: Optional[list] = None, + audio_data: Optional[list] = None, + return_text: Optional[bool] = True, + discard_alpha_channel: bool = True, + audio_sample_rate: Optional[int] = None, + ) -> BaseMultiModalProcessorOutput: + """ + A fast version of `load_mm_data` that loads multimodal data directly. + This version does not scan the prompt to recognize tokens. It assumes + that the caller has already aligned the tokens and data in a 1:1 manner. + The behavior is as follows: + 1. It runs `_load_single_item` for all input data concurrently. + 2. It returns the loaded images, videos, and audios in their original order. + 3. It returns the input prompt as a string. + """ + + # Convert prompt into str + if isinstance(prompt, list) and return_text: + assert len(prompt) and isinstance(prompt[0], int) + prompt_str = self._processor.tokenizer.decode(prompt) + else: + assert isinstance(prompt, str) + prompt_str = prompt + + futures: List[Tuple[Modality, int, concurrent.futures.Future]] = [] + + modalities_data = [ + (image_data, Modality.IMAGE), + (video_data, Modality.VIDEO), + (audio_data, Modality.AUDIO), + ] + + for data_list, modality in modalities_data: + futures.extend( + self._submit_mm_data_loading_tasks_simple( + data_list, modality, audio_sample_rate, discard_alpha_channel + ) + ) + + logger.debug("[load_mm_data(simple)] total futures submitted: %d", len(futures)) + + images: List[Any] = [None] * len(image_data) if image_data else [] + videos: List[Any] = [None] * len(video_data) if video_data else [] + audios: List[Any] = [None] * len(audio_data) if audio_data else [] + + for modality, idx, future in futures: + try: + result = future.result() + except Exception as e: + logger.exception( + "[load_mm_data(simple)] error loading %s data at index=%d", + modality.name, + idx, + ) + raise RuntimeError( + f"An exception occurred while loading {modality.name} data at index {idx}: {e}" + ) + + if modality == Modality.IMAGE: + images[idx] = result + elif modality == Modality.VIDEO: + videos[idx] = result + elif modality == Modality.AUDIO: + audios[idx] = result + + logger.debug( + "[load_mm_data(simple)] loaded counts: images=%d, videos=%d, audios=%d", + len(images), + len(videos), + len(audios), + ) + + return BaseMultiModalProcessorOutput( + images=images, + audios=audios, + videos=videos, + input_text=prompt_str, + ) + + def legacy_load_mm_data( + self, + prompt: str, + multimodal_tokens: MultimodalSpecialTokens, + image_data: Optional[list] = None, + video_data: Optional[list] = None, + audio_data: Optional[list] = None, + return_text: Optional[bool] = True, + discard_alpha_channel: bool = True, + audio_sample_rate: Optional[int] = None, + ) -> BaseMultiModalProcessorOutput: """ Each frame of video/image will be replaced by a single image token @@ -582,8 +773,6 @@ class BaseMultimodalProcessor(ABC): """ - BaseMultimodalProcessor.validate_mm_data(image_data, video_data, audio_data) - multimodal_tokens_pattern = multimodal_tokens.get_combined_regex() if isinstance(prompt, list) and return_text: assert len(prompt) and isinstance(prompt[0], int) diff --git a/python/sglang/srt/multimodal/processors/minicpm.py b/python/sglang/srt/multimodal/processors/minicpm.py index 9ddbf4fb6..defc047aa 100644 --- a/python/sglang/srt/multimodal/processors/minicpm.py +++ b/python/sglang/srt/multimodal/processors/minicpm.py @@ -14,6 +14,7 @@ from sglang.srt.multimodal.processors.base_processor import ( # Compatible with both 'O' and 'V' class MiniCPMMultimodalProcessor(BaseMultimodalProcessor): models = [MiniCPMV, MiniCPMO] + support_dynamic_frame_expansion = True def __init__(self, hf_config, server_args, _processor, *args, **kwargs): super().__init__(hf_config, server_args, _processor, *args, **kwargs)