diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 491ffa9cd..7dca71334 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -198,17 +198,17 @@ _mm_feature_attrs = { def _get_mm_grid_dim(mm_inputs, modality, model_type: Optional[str] = None): - # Kimi K2.5 vision processor only emits `grid_thws`; prefer it over generic keys - # so we never pick a mis-typed or stale `image_grid_hws` field from kwargs. attrs = _mm_grid_attrs[modality] - if (model_type or "").lower() in [ - "kimi_k25", - "kimi_vl", - ] and modality == Modality.IMAGE: - attrs = ("grid_thws", "image_grid_thw", "image_grid_hws") + model_type = (model_type or "").lower() + if modality == Modality.IMAGE: + # Kimi K2.5 emits grid_thws, while Kimi-VL emits image_grid_hws. + if model_type == "kimi_k25": + attrs = ("grid_thws", "image_grid_thw", "image_grid_hws") + elif model_type == "kimi_vl": + attrs = ("image_grid_hws", "image_grid_thw", "grid_thws") for attr in attrs: if attr in mm_inputs and mm_inputs[attr] is not None: - return mm_inputs[attr] + return _convert(mm_inputs[attr]) raise ValueError(f"Grid dim ({_mm_grid_attrs[modality]}) not found in {mm_inputs}") @@ -758,16 +758,33 @@ class MMEncoder: """Calculate number of raw patches (before merge/sampling). Used for pixel_values slicing.""" if modality == Modality.AUDIO: return int(grid.item()) + if self.model_type == "kimi_vl" and modality == Modality.IMAGE: + h, w = self._kimi_hw_from_patch_grid(grid) + return h * w + return int(grid[0] * grid[1] * grid[2]) + + @staticmethod + def _kimi_hw_from_patch_grid( + grid: Union[torch.Tensor, np.ndarray, List[int], Tuple[int, ...]], + ) -> Tuple[int, int]: + """Extract (height, width) from Kimi 2D or 3D patch-grid metadata.""" + if isinstance(grid, torch.Tensor): + values = grid.flatten().tolist() + elif isinstance(grid, np.ndarray): + values = grid.reshape(-1).tolist() else: - return int(grid[0] * grid[1] * grid[2]) + values = np.asarray(grid).reshape(-1).tolist() + + if len(values) not in (2, 3): + raise ValueError( + f"Invalid Kimi image grid metadata: {values}; " + "expected [h, w] or [t, h, w]" + ) + return int(values[-2]), int(values[-1]) def _kimi_tokens_from_patch_grid(self, grid: Union[torch.Tensor, List[int]]) -> int: """MoonViT + tpool: output len is (h//mh)*(w//mw); temporal dim is pooled (not t*h*w/merge^2).""" - if isinstance(grid, torch.Tensor): - flat = grid.flatten() - _t, h, w = (int(x) for x in flat[:3].tolist()) - else: - _t, h, w = int(grid[0]), int(grid[1]), int(grid[2]) + h, w = self._kimi_hw_from_patch_grid(grid) merge_h, merge_w = self.model_config.hf_config.vision_config.merge_kernel_size return (h * w) // (merge_h * merge_w) diff --git a/test/registered/disaggregation/test_epd_disaggregation.py b/test/registered/disaggregation/test_epd_disaggregation.py index 583cc79eb..fbd4a66d7 100644 --- a/test/registered/disaggregation/test_epd_disaggregation.py +++ b/test/registered/disaggregation/test_epd_disaggregation.py @@ -34,9 +34,10 @@ from sglang.test.vlm_utils import ( # Omni model for local testing; override via env var EPD_OMNI_MODEL DEFAULT_OMNI_MODEL = "Qwen/Qwen3-Omni-30B-A3B-Instruct" QWEN35_27B_MODEL = "Qwen/Qwen3.5-27B" +KIMI_VL_MODEL = "moonshotai/Kimi-VL-A3B-Instruct" -register_cuda_ci(est_time=97, stage="base-c", runner_config="4-gpu-h100") +register_cuda_ci(est_time=300, stage="base-c", runner_config="4-gpu-h100") @unittest.skipIf( @@ -750,6 +751,184 @@ class TestEPDDisaggregationOneEncoder(MMMUMixin, PDDisaggregationServerBase): print(f"Error killing process: {e}") +class TestEPDDisaggregationKimiVL(PDDisaggregationServerBase): + """Regression test for Kimi-VL two-dimensional image grids in E/PD mode.""" + + model = KIMI_VL_MODEL + model_args = [ + "--context-length=8192", + "--dtype=bfloat16", + "--mem-fraction-static=0.40", + ] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.process_encode = None + cls.encode_port = f"{int(cls.lb_port) + 307}" + cls.encode_url = f"http://{cls.base_host}:{cls.encode_port}" + + cls.start_encode() + prefill_thread = threading.Thread(target=cls.start_prefill) + decode_thread = threading.Thread(target=cls.start_decode) + prefill_thread.start() + decode_thread.start() + prefill_thread.join() + decode_thread.join() + + cls.wait_server_ready(cls.encode_url + "/health", process=cls.process_encode) + cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill) + cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode) + cls.launch_lb() + + @classmethod + def start_encode(cls): + encode_args = [ + "--trust-remote-code", + "--encoder-only", + "--encoder-transfer-backend", + "zmq_to_scheduler", + "--tp", + "1", + "--base-gpu-id", + "0", + "--port", + cls.encode_port, + *cls.model_args, + ] + cls.process_encode = popen_launch_server( + cls.model, + base_url=cls.encode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=encode_args, + ) + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--language-only", + "--encoder-urls", + cls.encode_url, + "--encoder-transfer-backend", + "zmq_to_scheduler", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "1", + "--base-gpu-id", + "1", + "--port", + cls.prefill_port, + *cls.model_args, + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + cls.process_prefill = popen_launch_server( + cls.model, + base_url=cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "1", + "--base-gpu-id", + "2", + "--port", + cls.decode_port, + *cls.model_args, + ] + decode_args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_server( + cls.model, + base_url=cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + ) + + @classmethod + def tearDownClass(cls): + try: + super().tearDownClass() + finally: + if cls.process_encode: + try: + kill_process_tree(cls.process_encode.pid) + except Exception as e: + print(f"Error killing encode process: {e}") + + def test_multi_image_chat_completion(self): + client = openai.Client(api_key="sk-123456", base_url=f"{self.lb_url}/v1") + response = client.chat.completions.create( + model="default", + messages=[ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": IMAGE_MAN_IRONING_URL}, + "modalities": "multi-images", + }, + { + "type": "image_url", + "image_url": {"url": IMAGE_SGL_LOGO_URL}, + "modalities": "multi-images", + }, + { + "type": "text", + "text": "Describe each image separately.", + }, + ], + }, + ], + temperature=0, + max_tokens=128, + ) + + self.assertEqual(response.choices[0].message.role, "assistant") + text = response.choices[0].message.content + self.assertIsInstance(text, str) + self.assertGreater(len(text), 0) + + text_lower = text.lower() + self.assertTrue( + any( + word in text_lower + for word in ("man", "person", "car", "vehicle", "suv", "iron") + ), + f"First image was not described correctly: {text}", + ) + self.assertTrue( + any( + word in text_lower + for word in ("logo", "sglang", "graphic", "stylized", "letter") + ), + f"Second image was not described correctly: {text}", + ) + + for name, process in ( + ("encoder", self.process_encode), + ("prefill", self.process_prefill), + ("decode", self.process_decode), + ("router", self.process_lb), + ): + self.assertIsNone( + process.poll(), + f"{name} process exited with code {process.returncode}", + ) + + @unittest.skipIf( is_in_ci(), "Qwen3.5 EPD image/video test runs locally only", diff --git a/test/registered/unit/disaggregation/test_encode_server.py b/test/registered/unit/disaggregation/test_encode_server.py new file mode 100644 index 000000000..ae45fd93d --- /dev/null +++ b/test/registered/unit/disaggregation/test_encode_server.py @@ -0,0 +1,108 @@ +import pickle +import unittest +from types import SimpleNamespace + +import numpy as np +import torch + +from sglang.srt.disaggregation.encode_receiver import EmbeddingData +from sglang.srt.disaggregation.encode_server import MMEncoder, _get_mm_grid_dim +from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.utils.common import safe_pickle_loads +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestKimiVLEPDGrid(unittest.TestCase): + @staticmethod + def _make_encoder(model_type="kimi_vl"): + encoder = MMEncoder.__new__(MMEncoder) + encoder.model_type = model_type + encoder.model_config = SimpleNamespace( + hf_config=SimpleNamespace( + vision_config=SimpleNamespace(merge_kernel_size=(2, 2)) + ) + ) + return encoder + + def test_kimi_vl_prefers_and_normalizes_hw_grid(self): + mm_inputs = { + "image_grid_hws": np.array([[40, 60]], dtype=np.int64), + "image_grid_thw": torch.tensor([[1, 20, 30]]), + "grid_thws": torch.tensor([[1, 10, 15]]), + } + + grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_vl") + + self.assertIsInstance(grid, torch.Tensor) + torch.testing.assert_close(grid, torch.tensor([[40, 60]])) + + def test_kimi_k25_keeps_thw_grid_preference(self): + mm_inputs = { + "image_grid_hws": np.array([[40, 60]], dtype=np.int64), + "grid_thws": np.array([[1, 10, 15]], dtype=np.int64), + } + + grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_k25") + + torch.testing.assert_close(grid, torch.tensor([[1, 10, 15]])) + + def test_kimi_vl_2d_grid_counting_and_slicing(self): + encoder = self._make_encoder() + grids = torch.tensor([[40, 60], [20, 40]]) + embedding = torch.arange(800 * 2).reshape(800, 2) + + self.assertEqual( + encoder.get_num_patches(grids[0], Modality.IMAGE), + 2400, + ) + self.assertEqual( + encoder.get_num_tokens(grids[0], Modality.IMAGE), + 600, + ) + + slices = encoder.slice_embedding(embedding, grids, Modality.IMAGE) + + self.assertEqual([item.shape for item in slices], [(600, 2), (200, 2)]) + torch.testing.assert_close(slices[0], embedding[:600]) + torch.testing.assert_close(slices[1], embedding[600:]) + + def test_kimi_3d_grid_remains_supported(self): + encoder = self._make_encoder() + grid = torch.tensor([1, 40, 60]) + + self.assertEqual(encoder.get_num_patches(grid, Modality.IMAGE), 2400) + self.assertEqual(encoder.get_num_tokens(grid, Modality.IMAGE), 600) + + def test_kimi_k25_3d_patch_counting_is_unchanged(self): + encoder = self._make_encoder("kimi_k25") + grid = torch.tensor([2, 12, 16]) + + self.assertEqual(encoder.get_num_patches(grid, Modality.IMAGE), 384) + self.assertEqual(encoder.get_num_tokens(grid, Modality.IMAGE), 48) + + def test_grid_metadata_is_safe_to_deserialize(self): + grid = _get_mm_grid_dim( + {"image_grid_hws": np.array([[40, 60]], dtype=np.int64)}, + Modality.IMAGE, + "kimi_vl", + ) + embedding_data = EmbeddingData( + req_id="test-request", + num_parts=1, + part_idx=0, + grid_dim=grid, + modality=Modality.IMAGE, + embedding=torch.zeros((600, 4)), + ) + + restored = safe_pickle_loads( + pickle.dumps(embedding_data.copy_without_embedding()) + ) + + torch.testing.assert_close(restored.grid_dim, torch.tensor([[40, 60]])) + + +if __name__ == "__main__": + unittest.main()