diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 86eb37e8f..799db26cc 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -555,6 +555,9 @@ class Envs: SGLANG_MM_PRECOMPUTE_HASH = EnvBool(False) SGLANG_VIT_ENABLE_CUDA_GRAPH = EnvBool(False) SGLANG_MM_SKIP_COMPUTE_HASH = EnvBool(False) + # For pre-tokenized (list[int]) multimodal prompts, + # preserve the user's original tokens to avoid retokenization drift. + SGLANG_MM_AVOID_RETOKENIZE = EnvBool(True) # VLM Item CUDA IPC Transport diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index f34e2b54a..3dca4f28d 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -1227,6 +1227,58 @@ class BaseMultimodalProcessor(ABC): return tensor return tensor.cpu() + def resolve_image_token_counts(self, images: List) -> List[int]: + """Per-image expanded token counts, computed without re-tokenizing. + + Default implementation uses the transformers in-tree convention + ``_get_num_multimodal_tokens(image_sizes=...)`` (present on the in-tree + VLM processors, e.g. Qwen-VL, Gemma3, GLM4V). Models whose processor + does not implement it (e.g. Kimi) override this method. + + """ + assert images is not None + image_sizes = [(image.height, image.width) for image in images] + num_image_tokens = self._processor._get_num_multimodal_tokens( + image_sizes=image_sizes + ).num_image_tokens + return [int(count) for count in num_image_tokens] + + @staticmethod + def _expand_input_ids( + original_ids: List[int], + counts: List[int], + placeholder_token_id: Optional[int], + ) -> List[int]: + """Rebuild final input_ids for a pre-tokenized (list[int]) prompt. + + Keep the user's ORIGINAL tokens verbatim and expand the i-th image + placeholder into ``counts[i]`` copies of ``placeholder_token_id``. The HF + processor's re-tokenization is discarded, so non-media tokens cannot + drift. + + """ + if placeholder_token_id is None: + raise ValueError("placeholder_token_id is not set for this processor") + + num_placeholders = sum( + 1 for token_id in original_ids if token_id == placeholder_token_id + ) + if num_placeholders != len(counts): + raise ValueError( + f"prompt has {num_placeholders} image placeholder token(s) but " + f"{len(counts)} image(s) were provided" + ) + + rebuilt: List[int] = [] + next_image_idx = 0 + for token_id in original_ids: + if token_id == placeholder_token_id: + rebuilt.extend([placeholder_token_id] * counts[next_image_idx]) + next_image_idx += 1 + else: + rebuilt.append(token_id) + return rebuilt + def process_and_combine_mm_data( self, base_output: BaseMultiModalProcessorOutput, @@ -1276,6 +1328,48 @@ class BaseMultimodalProcessor(ABC): **kwargs, ) all_collected_items = collected_items + + # When SGLANG_MM_AVOID_RETOKENIZE is on, keep the user's exact tokens to avoid retokenize drift. + # Drift happens when Retokenization is not identity: Decode(X) => String => Re-tokenize => Y, X != Y. + if ( + envs.SGLANG_MM_AVOID_RETOKENIZE.get() + and base_output.input_ids is not None + and input_ids is not None + and raw_images + and not raw_audios + and not raw_videos + ): + assert isinstance( + base_output.input_ids, list + ), f"expected list[int] input_ids, got {type(base_output.input_ids)}" + try: + counts = self.resolve_image_token_counts(raw_images) + image_placeholder_token_id = mm_tokens.image_token_id + if image_placeholder_token_id is None: + raise ValueError( + "image placeholder token id is not set for this processor" + ) + processor_placeholder_count = int( + (input_ids == image_placeholder_token_id).sum().item() + ) + if processor_placeholder_count != sum(counts): + raise ValueError( + "processor image placeholder count mismatch: " + f"processor={processor_placeholder_count}, " + f"resolved={sum(counts)}" + ) + input_ids = torch.tensor( + self._expand_input_ids( + base_output.input_ids, + counts, + image_placeholder_token_id, + ), + dtype=input_ids.dtype, + ) + except Exception as e: + logger.warning( + f"Due to {e}, falling back to decode+retokenize, which may change prompt length (token drift)." + ) else: ret = None diff --git a/python/sglang/srt/multimodal/processors/kimi_common.py b/python/sglang/srt/multimodal/processors/kimi_common.py index c2046d32c..5371fee71 100644 --- a/python/sglang/srt/multimodal/processors/kimi_common.py +++ b/python/sglang/srt/multimodal/processors/kimi_common.py @@ -23,6 +23,21 @@ class KimiGridMMDataMixin: - self._tokenizer (with .encode()) """ + def resolve_image_token_counts(self, images): + """Kimi's processor is remote-code and does not implement the + transformers ``_get_num_multimodal_tokens`` convention; use its + ``media_tokens_calculator`` instead. + + """ + assert images is not None + media_tokens_calculator = ( + self._processor.media_processor.media_tokens_calculator + ) + return [ + int(media_tokens_calculator({"type": "image", "image": image})) + for image in images + ] + def _num_image_tokens_from_grid( self, grid_thw: Union[torch.Tensor, np.ndarray, list, tuple] ) -> int: diff --git a/test/registered/vlm/test_token_id_retokenize_e2e.py b/test/registered/vlm/test_token_id_retokenize_e2e.py new file mode 100644 index 000000000..8a7192bd7 --- /dev/null +++ b/test/registered/vlm/test_token_id_retokenize_e2e.py @@ -0,0 +1,115 @@ +"""E2E test for SGLANG_MM_AVOID_RETOKENIZE on the pre-tokenized VLM path. + +A client may send a multimodal request as input_ids (list[int]) instead of text. +On that path the server decodes the ids back to text and the HF processor +re-tokenizes them. If the original ids were non-canonical (decode -> re-encode is +not identity), that re-tokenization drifts: the reported prompt_tokens changes. + +With SGLANG_MM_AVOID_RETOKENIZE ON (default), the server keeps the user's +original tokens verbatim and only expands the image placeholder, so prompt_tokens +stays faithful to what the client sent. + +For each model we launch a real server twice with the same predefined, +non-canonical prompt ("Describe" split into "D"+"escribe") plus one image: + + * flag OFF -> the prompt re-tokenizes (drift): prompt_tokens shrinks by the + drift delta. + * flag ON -> no drift: prompt_tokens equals the original length (with the + image placeholder expanded). +""" + +import base64 +import io +import unittest + +import requests +from PIL import Image +from transformers import AutoProcessor + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=300, suite="base-b-test-1-gpu-large") + + +def _data_uri(): + img = Image.new("RGB", (64, 64), (128, 128, 128)) + buf = io.BytesIO() + img.save(buf, format="PNG") + return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode() + + +def _build_drift_prompt(model, image_token): + """Return (input_ids, drift_delta). + + input_ids is a predefined non-canonical prompt: "Describe" is split into + "D"+"escribe" (decodes to the same text but re-encodes to the single merged + token), followed by one image placeholder. drift_delta is how many extra + tokens the non-canonical form carries vs. the canonical re-tokenization. + """ + tok = AutoProcessor.from_pretrained( + model, trust_remote_code=True, use_fast=True + ).tokenizer + + def enc(text): + return tok.encode(text, add_special_tokens=False) + + input_ids = enc("D") + enc("escribe") + enc(" the picture: ") + enc(image_token) + canonical = enc(tok.decode(input_ids)) + drift_delta = len(input_ids) - len(canonical) + return input_ids, drift_delta + + +def _prompt_tokens(base_url, input_ids, image): + resp = requests.post( + base_url + "/generate", + json={ + "input_ids": input_ids, + "image_data": [image], + "sampling_params": {"temperature": 0.0, "max_new_tokens": 1}, + }, + timeout=300, + ) + resp.raise_for_status() + return resp.json()["meta_info"]["prompt_tokens"] + + +class TestQwenVLTokenIdRetokenize(CustomTestCase): + model = "Qwen/Qwen2.5-VL-3B-Instruct" + image_token = "<|vision_start|><|image_pad|><|vision_end|>" + other_args = ["--trust-remote-code", "--mem-fraction-static", "0.7"] + + def test_flag_off_drifts_flag_on_does_not(self): + input_ids, drift_delta = _build_drift_prompt(self.model, self.image_token) + self.assertGreater(drift_delta, 0, "prompt is canonical; no drift to exercise") + image = _data_uri() + + prompt_tokens = {} + for flag in ("0", "1"): + process = popen_launch_server( + self.model, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=self.other_args, + env={"SGLANG_MM_AVOID_RETOKENIZE": flag}, + ) + try: + prompt_tokens[flag] = _prompt_tokens( + DEFAULT_URL_FOR_TEST, input_ids, image + ) + finally: + kill_process_tree(process.pid) + + # ON keeps the user's original tokens; OFF loses the drift_delta tokens. + pt_off, pt_on = prompt_tokens["0"], prompt_tokens["1"] + self.assertEqual(pt_on - pt_off, drift_delta, f"on={pt_on}, off={pt_off}") + + +if __name__ == "__main__": + unittest.main()