From 39e147443bfd750252892e1dc2e46af8439b0679 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Sun, 13 Sep 2026 21:06:47 -0700 Subject: [PATCH] [Session] Fix image append positions and parent metadata (#39145) Co-authored-by: Byron Hsu Co-authored-by: Manik Singhal --- python/sglang/srt/managers/schedule_batch.py | 51 +++++++- .../sessions/test_session_control.py | 116 +++++++++++++++++- 2 files changed, 165 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 178399f72..b2036984b 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1414,11 +1414,60 @@ class Req(ReqDllmMixin): self.spec_cap_lens_histogram[cap_len] += 1 def extend_image_inputs(self, image_inputs): - if self.multimodal_inputs is None: + if self.session is not None: + self._extend_session_image_inputs(image_inputs) + elif self.multimodal_inputs is None: self.multimodal_inputs = image_inputs else: self.multimodal_inputs.merge(image_inputs) + def _extend_session_image_inputs(self, image_inputs): + """Append media while preserving the saved session and its position history.""" + # Padding can change token values without changing their count. + self.full_untruncated_fill_ids = array("q") + if self.multimodal_inputs is not None: + # Branches and aborted turns must leave the parent's metadata intact. + self.multimodal_inputs = dataclasses.replace(self.multimodal_inputs) + + positions = image_inputs.mrope_positions + if positions is not None: + prefix_len = len(self.origin_input_ids) - positions.shape[1] + prefix = ( + self.multimodal_inputs.mrope_positions + if self.multimodal_inputs is not None + else None + ) + if prefix is None: + prefix = positions.new_empty((3, 0)) + prefix = prefix[:, :prefix_len] + next_position = prefix.max() + 1 if prefix.numel() else 0 + text_len = prefix_len - prefix.shape[1] + text_positions = ( + torch.arange( + text_len, dtype=positions.dtype, device=positions.device + ).expand(3, -1) + + next_position + ) + # Fill the reply/text gap, then shift the new turn's media coordinates. + positions = torch.cat( + [prefix, text_positions, positions + next_position + text_len], dim=1 + ) + + if self.multimodal_inputs is None: + self.multimodal_inputs = image_inputs + else: + # Use the full table above, or let the scheduler compute missing positions. + self.multimodal_inputs.mrope_positions = None + self.multimodal_inputs.mrope_position_delta = None + self.multimodal_inputs.merge(image_inputs) + + self.multimodal_inputs.mrope_position_delta_repeated_cache = None + if positions is not None: + self.multimodal_inputs.mrope_positions = positions + self.multimodal_inputs.mrope_position_delta = ( + positions.max() + 1 - positions.shape[1] + ).reshape(1, 1) + def finished(self) -> bool: # Whether request reached finished condition return self.finished_reason is not None diff --git a/test/registered/sessions/test_session_control.py b/test/registered/sessions/test_session_control.py index f7ed8c717..fbe2498d4 100644 --- a/test/registered/sessions/test_session_control.py +++ b/test/registered/sessions/test_session_control.py @@ -7,11 +7,15 @@ python3 -m unittest test_session_control.TestSessionControlVision.test_session_c """ import asyncio +import base64 +import io import json import unittest +import uuid import aiohttp import requests +from PIL import Image from sglang.srt.utils import is_hip, kill_process_tree from sglang.srt.utils.hf_transformers_utils import get_tokenizer @@ -24,7 +28,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=102, stage="extra-a", runner_config="1-gpu-large") +register_cuda_ci(est_time=137, stage="extra-a", runner_config="1-gpu-large") register_amd_ci(est_time=87, suite="stage-b-test-1-gpu-large-amd") @@ -786,5 +790,115 @@ class TestSessionControlVision(CustomTestCase): ) +class TestSessionControlMrope(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + "Qwen/Qwen2.5-VL-3B-Instruct", + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--enable-streaming-session", + "--context-length", + "4096", + "--mem-fraction-static", + "0.5", + "--max-running-requests", + "4", + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def post(self, path, **payload): + response = requests.post(self.base_url + path, json=payload, timeout=180) + response.raise_for_status() + return response.json() if response.content else None + + def test_image_append(self): + images = [] + for color in ("red", "blue"): + buffer = io.BytesIO() + Image.new("RGB", (112, 112), color).save(buffer, format="PNG") + images.append( + "data:image/png;base64," + base64.b64encode(buffer.getvalue()).decode() + ) + prompt = ( + "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>" + "What color is this image?<|im_end|>\n<|im_start|>assistant\n" + ) + sampling = dict( + temperature=0, + max_new_tokens=8, + ignore_eos=True, + skip_special_tokens=False, + no_stop_trim=True, + ) + for streaming in (False, True): + with self.subTest(streaming=streaming): + salt = uuid.uuid4().hex + sid = self.post( + "/open_session", capacity_of_str_len=4096, streaming=streaming + ) + try: + first = self.post( + "/generate", + text=prompt, + image_data=images[:1], + session_params={"id": sid}, + cache_salt=salt, + sampling_params=sampling, + ) + self.assertEqual(len(first["output_ids"]), 8) + suffix = "<|im_end|>\n" + prompt + # Regular mode also branches from the first turn after an append. + for image in images[1:] + (images[:1] if not streaming else []): + # A separate cache namespace prevents the reference from warming the session. + reference = self.post( + "/generate", + text=prompt + first["text"] + suffix, + image_data=[images[0], image], + cache_salt=uuid.uuid4().hex, + sampling_params=sampling, + ) + appended = self.post( + "/generate", + text=suffix, + image_data=[image], + session_params={"id": sid, "rid": first["meta_info"]["id"]}, + cache_salt=salt, + sampling_params=sampling, + ) + self.assertEqual(len(appended["output_ids"]), 8) + self.assertEqual( + appended["output_ids"], reference["output_ids"] + ) + print( + "Session replay example: " + + json.dumps( + { + "streaming": streaming, + "new_image": "red" + if image == images[0] + else "blue", + "first_reply": first["text"], + "session_reply": appended["text"], + "replay_reply": reference["text"], + "session_ids": appended["output_ids"], + "replay_ids": reference["output_ids"], + } + ), + flush=True, + ) + finally: + self.post("/close_session", session_id=sid) + self.assertEqual( + requests.get(self.base_url + "/health", timeout=10).status_code, 200 + ) + + if __name__ == "__main__": unittest.main()