[Session] Fix image append positions and parent metadata (#39145)
Co-authored-by: Byron Hsu <byron+per@periodiclabs.ai> Co-authored-by: Manik Singhal <manikvsinghal.pub@gmail.com>
This commit is contained in:
co-authored by
Byron Hsu
Manik Singhal
parent
5a132c061b
commit
39e147443b
@@ -1414,11 +1414,60 @@ class Req(ReqDllmMixin):
|
|||||||
self.spec_cap_lens_histogram[cap_len] += 1
|
self.spec_cap_lens_histogram[cap_len] += 1
|
||||||
|
|
||||||
def extend_image_inputs(self, image_inputs):
|
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
|
self.multimodal_inputs = image_inputs
|
||||||
else:
|
else:
|
||||||
self.multimodal_inputs.merge(image_inputs)
|
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:
|
def finished(self) -> bool:
|
||||||
# Whether request reached finished condition
|
# Whether request reached finished condition
|
||||||
return self.finished_reason is not None
|
return self.finished_reason is not None
|
||||||
|
|||||||
@@ -7,11 +7,15 @@ python3 -m unittest test_session_control.TestSessionControlVision.test_session_c
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
import json
|
import json
|
||||||
import unittest
|
import unittest
|
||||||
|
import uuid
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import requests
|
import requests
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
from sglang.srt.utils import is_hip, kill_process_tree
|
from sglang.srt.utils import is_hip, kill_process_tree
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||||
@@ -24,7 +28,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
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")
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user