[Refactor] New EPD (#30398)

Co-authored-by: Yuang Chen <1131578721@qq.com>
Co-authored-by: Yuang Chen <cya539102@antgroup.com>
Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
siyu
2026-08-21 15:22:48 +08:00
committed by GitHub
co-authored by Yuang Chen Yuang Chen ZhengWG
parent 6a12583679
commit 8a123cbd0e
25 changed files with 6375 additions and 5305 deletions
@@ -10,7 +10,7 @@ import zmq
from prometheus_client.parser import text_string_to_metric_families
from prometheus_client.samples import Sample
from sglang.srt.disaggregation.encode_server import MINIMUM_PNG_PICTURE_BASE64
from sglang.srt.disaggregation.encoder.http_server import MINIMUM_PNG_PICTURE_BASE64
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.network import get_zmq_socket_on_host
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
@@ -4,7 +4,7 @@ import unittest
from array import array
from types import SimpleNamespace
from sglang.srt.disaggregation.encode_receiver import MMReceiverBase
from sglang.srt.disaggregation.encoder.receiver import MMReceiverBase
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci
@@ -1,30 +1,53 @@
import asyncio
import pickle
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
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.disaggregation.encoder.preprocessor import EncoderPreprocessor
from sglang.srt.disaggregation.encoder.receiver import EmbeddingData
from sglang.srt.disaggregation.encoder.runtime import execute_encode_pipeline
from sglang.srt.disaggregation.encoder.server import (
EncoderDelivery,
InternalError,
MMEncoder,
MooncakeDelivery,
ReqState,
SendDestination,
ZmqDelivery,
meta_registry,
rid_to_cond,
rid_to_receive_count,
rid_to_receive_endpoint,
)
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
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestKimiVLEPDGrid(unittest.TestCase):
class TestEncoderPreprocessorKimiGrid(CustomTestCase):
@staticmethod
def _make_encoder(model_type="kimi_vl"):
encoder = MMEncoder.__new__(MMEncoder)
encoder.model_type = model_type
encoder.model_config = SimpleNamespace(
def _make_preprocessor(model_type="kimi_vl"):
preprocessor = EncoderPreprocessor.__new__(EncoderPreprocessor)
preprocessor.model_type = model_type
preprocessor.model_config = SimpleNamespace(
hf_config=SimpleNamespace(
vision_config=SimpleNamespace(merge_kernel_size=(2, 2))
)
)
return encoder
preprocessor.image_processor = SimpleNamespace(merge_size=2)
preprocessor._model_preprocessor = None
return preprocessor
@staticmethod
def _make_encoder():
return MMEncoder.__new__(MMEncoder)
def test_kimi_vl_prefers_and_normalizes_hw_grid(self):
mm_inputs = {
@@ -33,7 +56,7 @@ class TestKimiVLEPDGrid(unittest.TestCase):
"grid_thws": torch.tensor([[1, 10, 15]]),
}
grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_vl")
grid = self._make_preprocessor()._get_mm_grid_dim(mm_inputs, Modality.IMAGE)
self.assertIsInstance(grid, torch.Tensor)
torch.testing.assert_close(grid, torch.tensor([[40, 60]]))
@@ -44,49 +67,51 @@ class TestKimiVLEPDGrid(unittest.TestCase):
"grid_thws": np.array([[1, 10, 15]], dtype=np.int64),
}
grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_k25")
grid = self._make_preprocessor("kimi_k25")._get_mm_grid_dim(
mm_inputs, Modality.IMAGE
)
torch.testing.assert_close(grid, torch.tensor([[1, 10, 15]]))
def test_kimi_vl_2d_grid_counting_and_slicing(self):
preprocessor = self._make_preprocessor()
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),
preprocessor.get_num_patches(grids[0], Modality.IMAGE),
2400,
)
self.assertEqual(
encoder.get_num_tokens(grids[0], Modality.IMAGE),
preprocessor.get_num_tokens(grids[0], Modality.IMAGE),
600,
)
slices = encoder.slice_embedding(embedding, grids, Modality.IMAGE)
slices = encoder.slice_embedding(embedding, [600, 200])
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()
preprocessor = self._make_preprocessor()
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)
self.assertEqual(preprocessor.get_num_patches(grid, Modality.IMAGE), 2400)
self.assertEqual(preprocessor.get_num_tokens(grid, Modality.IMAGE), 600)
def test_kimi_k25_3d_patch_counting_is_unchanged(self):
encoder = self._make_encoder("kimi_k25")
preprocessor = self._make_preprocessor("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)
self.assertEqual(preprocessor.get_num_patches(grid, Modality.IMAGE), 384)
self.assertEqual(preprocessor.get_num_tokens(grid, Modality.IMAGE), 48)
def test_grid_metadata_is_safe_to_deserialize(self):
grid = _get_mm_grid_dim(
grid = self._make_preprocessor()._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",
@@ -104,5 +129,411 @@ class TestKimiVLEPDGrid(unittest.TestCase):
torch.testing.assert_close(restored.grid_dim, torch.tensor([[40, 60]]))
class TestEncoderDelivery(CustomTestCase):
def test_contract_has_two_direct_implementations(self):
self.assertEqual(EncoderDelivery.__abstractmethods__, {"send", "release"})
self.assertEqual(
set(EncoderDelivery.__subclasses__()),
{
MooncakeDelivery,
ZmqDelivery,
},
)
def test_zmq_delivery_cleanup_is_configurable(self):
async def run():
req_id = "test-zmq-delivery-cleanup"
rid_to_receive_endpoint[req_id] = {"127.0.0.1:1"}
rid_to_receive_count[req_id] = 1
rid_to_cond[req_id] = asyncio.Condition()
state = ReqState(req_id)
encoder = SimpleNamespace()
await ZmqDelivery(encoder, cleanup_receive_state=False).release(state)
self.assertIn(req_id, rid_to_receive_endpoint)
self.assertIn(req_id, rid_to_receive_count)
self.assertIn(req_id, rid_to_cond)
await ZmqDelivery(encoder, cleanup_receive_state=True).release(state)
self.assertNotIn(req_id, rid_to_receive_endpoint)
self.assertNotIn(req_id, rid_to_receive_count)
self.assertNotIn(req_id, rid_to_cond)
asyncio.run(run())
def test_preprocess_metadata_precedes_embedding(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.req_states = {}
encoder._embedding_dims = {Modality.IMAGE: 8}
encoder._embedding_dtype = torch.float16
encoder._element_size = 2
first_state = encoder._acquire_encode_ref("req-0")
second_state = encoder._acquire_encode_ref("req-1")
ctx = SimpleNamespace(
req_id="req-0",
modality=Modality.IMAGE,
items_per_req=[1, 2],
preprocess_result=SimpleNamespace(
token_counts=[2, 3, 4],
grid_thw=[[1, 2, 3], [1, 4, 5], [1, 6, 7]],
),
)
requests = [
{"req_id": "req-0", "num_parts": 2, "part_idx": 0},
{"req_id": "req-1", "num_parts": 2, "part_idx": 1},
]
publish = AsyncMock()
with patch.object(meta_registry, "publish", publish):
await encoder._publish_preprocess_metadata(ctx, requests)
self.assertIs(encoder.req_states["req-0"], first_state)
self.assertIs(encoder.req_states["req-1"], second_state)
self.assertEqual(first_state.embedding_data.shape, [2, 8])
self.assertEqual(second_state.embedding_data.shape, [7, 8])
self.assertEqual(first_state.embedding_data.grid_dim, [[1, 2, 3]])
self.assertEqual(
second_state.embedding_data.grid_dim,
[[1, 4, 5], [1, 6, 7]],
)
self.assertEqual(first_state.embedding_data.dtype, torch.float16)
self.assertEqual(second_state.embedding_data.dtype, torch.float16)
self.assertFalse(first_state.embedding_ready.is_set())
self.assertFalse(second_state.embedding_ready.is_set())
self.assertEqual(
publish.await_args_list,
[
unittest.mock.call("req-0", 32, 2, 8),
unittest.mock.call("req-1", 112, 7, 8),
],
)
await encoder._release_encode_ref(first_state)
await encoder._release_encode_ref(second_state)
asyncio.run(run())
def test_mooncake_embedding_is_ready_only_after_cuda_sync(self):
class FakeCudaEmbedding:
shape = (2, 4)
dtype = torch.float16
nbytes = 16
is_cuda = True
device = "cuda:0"
def __getitem__(self, key):
return self
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
events = []
encoder._stage_embedding = lambda mm_data: events.append("ready")
ctx = SimpleNamespace(
req_id="req",
modality=Modality.IMAGE,
items_per_req=[1],
preprocess_result=SimpleNamespace(
token_counts=[2],
grid_thw=[[1, 2, 3]],
),
aux_data={},
use_global_cache=True,
)
stream = SimpleNamespace(synchronize=lambda: events.append("sync"))
with patch.object(torch.cuda, "current_stream", return_value=stream):
encoder._stage_embeddings(
ctx,
[{"req_id": "req", "num_parts": 1, "part_idx": 0}],
FakeCudaEmbedding(),
keep_on_gpu=True,
)
self.assertEqual(events, ["sync", "ready"])
def test_stage_embedding_does_not_resurrect_missing_state(self):
encoder = MMEncoder.__new__(MMEncoder)
encoder.req_states = {}
with self.assertRaisesRegex(
InternalError, "No request state exists while encoding request: req"
):
encoder._stage_embedding(
EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding=torch.ones((1, 1)),
)
)
self.assertNotIn("req", encoder.req_states)
def test_stage_embedding_requires_active_encode(self):
encoder = MMEncoder.__new__(MMEncoder)
state = ReqState("req")
encoder.req_states = {"req": state}
with self.assertRaisesRegex(
InternalError, "Request state has no active encode work: req"
):
encoder._stage_embedding(
EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding=torch.ones((1, 1)),
)
)
self.assertIs(encoder.req_states["req"], state)
self.assertIsNone(state.embedding_data)
self.assertFalse(state.embedding_ready.is_set())
def test_release_during_encode_is_deferred_without_resurrecting_state(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.req_states = {}
encoder.delivery = SimpleNamespace(release=AsyncMock())
state = encoder._acquire_encode_ref("req")
state.embedding_data = EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding_shape=[1, 1],
dtype=torch.float32,
)
discard = AsyncMock()
with patch.object(meta_registry, "discard", discard):
await encoder.release_request("req")
self.assertTrue(state.release_requested)
self.assertIn("req", encoder.req_states)
encoder.delivery.release.assert_not_awaited()
embedding = torch.ones((1, 1))
encoder._stage_embedding(
EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding=embedding,
)
)
await encoder._release_encode_ref(state)
encoder.delivery.release.assert_awaited_once_with(state)
discard.assert_awaited_once_with("req")
self.assertIsNone(state.embedding_data)
self.assertNotIn("req", encoder.req_states)
asyncio.run(run())
def test_error_metadata_survives_buffer_release_for_waiter(self):
async def run():
req_id = "test-error-metadata-waiter"
await meta_registry.discard(req_id)
encoder = MMEncoder.__new__(MMEncoder)
encoder.req_states = {}
encoder.delivery = SimpleNamespace(release=AsyncMock())
state = ReqState(
req_id,
EmbeddingData(
req_id,
1,
0,
None,
Modality.IMAGE,
error_msg="encode failed",
),
)
state.embedding_ready.set()
encoder.req_states[req_id] = state
waiter = asyncio.create_task(meta_registry.wait(req_id))
await asyncio.sleep(0)
try:
await meta_registry.publish(req_id, 0, 0, 0, error="encode failed")
await encoder.release_request(req_id, preserve_metadata=True)
meta = await asyncio.wait_for(waiter, timeout=1)
self.assertEqual(meta, {"error": "encode failed"})
self.assertNotIn(req_id, encoder.req_states)
finally:
if not waiter.done():
waiter.cancel()
await meta_registry.discard(req_id)
asyncio.run(run())
def test_zmq_pipeline_sends_only_after_encode_completes(self):
async def run():
events = []
finish_encode = asyncio.Event()
encoder = MMEncoder.__new__(MMEncoder)
encoder.transfer_backend = "zmq_to_tokenizer"
async def encode(**kwargs):
events.append("metadata_published")
await finish_encode.wait()
events.append("encode_completed")
return 16, 2, 4, None, None
async def send(**kwargs):
events.append("send")
return True
async def release_request(req_id, **kwargs):
events.append("release")
encoder.encode = AsyncMock(side_effect=encode)
encoder.send = AsyncMock(side_effect=send)
encoder.release_request = AsyncMock(side_effect=release_request)
publish = AsyncMock()
request = {
"req_id": "req",
"mm_items": ["item"],
"modality": "image",
"num_parts": 1,
"part_idx": 0,
"prefill_host": "127.0.0.1",
"embedding_port": 1234,
}
with patch.object(meta_registry, "publish", publish):
task = asyncio.create_task(
execute_encode_pipeline(encoder, None, request)
)
await asyncio.sleep(0)
self.assertEqual(events, ["metadata_published"])
encoder.send.assert_not_awaited()
finish_encode.set()
self.assertIsNone(await task)
self.assertEqual(
events,
["metadata_published", "encode_completed", "send", "release"],
)
publish.assert_awaited_once_with("req", 16, 2, 4)
asyncio.run(run())
def test_send_waits_for_embedding_published_by_encode(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.rank = 0
encoder.req_states = {}
state = encoder._acquire_encode_ref("req")
state.embedding_data = EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding_shape=[1, 1],
dtype=torch.float32,
)
delivered = []
async def send(current_state, destination):
delivered.append(await encoder._wait_for_embedding(current_state))
encoder.delivery = SimpleNamespace(
send=AsyncMock(side_effect=send),
release=AsyncMock(),
)
send_task = asyncio.create_task(
encoder.send_to_destination(state, SendDestination("127.0.0.1:1"))
)
await asyncio.sleep(0)
self.assertFalse(send_task.done())
embedding = torch.ones((1, 1))
encoder._stage_embedding(
EmbeddingData(
"req",
1,
0,
None,
Modality.IMAGE,
embedding=embedding,
)
)
await send_task
await encoder._release_encode_ref(state)
with patch.object(meta_registry, "discard", AsyncMock()):
await encoder.release_request("req")
self.assertEqual(len(delivered), 1)
self.assertIs(delivered[0].embedding, embedding)
asyncio.run(run())
def test_release_waits_for_send_then_clears_embedding(self):
async def run():
encoder = MMEncoder.__new__(MMEncoder)
encoder.req_states = {}
send_started = asyncio.Event()
finish_send = asyncio.Event()
async def send(state, destination):
send_started.set()
await finish_send.wait()
embedding_seen_by_release = []
async def release(state):
embedding_seen_by_release.append(state.embedding_data.embedding)
encoder.delivery = SimpleNamespace(
send=AsyncMock(side_effect=send),
release=AsyncMock(side_effect=release),
)
embedding = torch.ones((1, 1))
state = ReqState(
"req",
EmbeddingData("req", 1, 0, None, Modality.IMAGE, embedding=embedding),
)
state.embedding_ready.set()
encoder.req_states[state.req_id] = state
send_task = asyncio.create_task(
encoder.send_to_destination(state, SendDestination("127.0.0.1:1"))
)
await send_started.wait()
release_task = asyncio.create_task(encoder.release_request("req"))
await asyncio.sleep(0)
encoder.delivery.release.assert_not_awaited()
self.assertIs(state.embedding_data.embedding, embedding)
finish_send.set()
await send_task
await release_task
encoder.delivery.release.assert_awaited_once_with(state)
self.assertEqual(len(embedding_seen_by_release), 1)
self.assertIs(embedding_seen_by_release[0], embedding)
self.assertIsNone(state.embedding_data)
self.assertNotIn("req", encoder.req_states)
asyncio.run(run())
if __name__ == "__main__":
unittest.main()
@@ -3,7 +3,8 @@ import sys
import pytest
from sglang.srt.disaggregation import encode_server
from sglang.srt.disaggregation.encoder import http_server
from sglang.srt.managers.schedule_batch import Modality
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -17,18 +18,27 @@ class _FakeEncoder:
self.encode_dispatch_lock = asyncio.Lock()
self.encode_calls = []
def has_pending_embeddings(self):
return bool(self.embedding_to_send)
def supports_modality(self, modality):
return modality == Modality.IMAGE
async def encode(self, **kwargs):
self.encode_calls.append(kwargs)
return 1, 1, 1, None, None
async def release_request(self, _req_id):
return None
def _install_tp_encoder(monkeypatch, encoder):
broadcasts = []
monkeypatch.setattr(encode_server, "dp_dispatcher", None)
monkeypatch.setattr(encode_server, "encoder", encoder)
monkeypatch.setattr(encode_server, "send_sockets", [object()])
monkeypatch.setattr(http_server, "dp_dispatcher", None)
monkeypatch.setattr(http_server, "encoder", encoder)
monkeypatch.setattr(http_server, "send_sockets", [object()])
monkeypatch.setattr(
encode_server,
http_server,
"sock_send",
lambda socket, payload: broadcasts.append((socket, payload)),
)
@@ -41,7 +51,7 @@ def test_health_encode_waits_for_collective_dispatch_lock(monkeypatch):
broadcasts = _install_tp_encoder(monkeypatch, encoder)
await encoder.encode_dispatch_lock.acquire()
task = asyncio.create_task(encode_server.health_generate())
task = asyncio.create_task(http_server.health_generate())
await asyncio.sleep(0)
assert broadcasts == []
assert encoder.encode_calls == []
@@ -61,7 +71,7 @@ def test_health_encode_rechecks_busy_state_after_waiting(monkeypatch):
broadcasts = _install_tp_encoder(monkeypatch, encoder)
await encoder.encode_dispatch_lock.acquire()
task = asyncio.create_task(encode_server.health_generate())
task = asyncio.create_task(http_server.health_generate())
await asyncio.sleep(0)
encoder.embedding_to_send["real-request"] = object()
encoder.encode_dispatch_lock.release()
@@ -3,7 +3,7 @@ import sys
import pytest
from sglang.srt.disaggregation.encode_server import (
from sglang.srt.disaggregation.encoder.runtime import (
EncoderScheduler,
PendingRequest,
_resolve_encoder_batch_policy,
@@ -15,14 +15,18 @@ import zmq.asyncio
from fastapi import HTTPException
from PIL import Image
from sglang.srt.disaggregation.encode_receiver import (
from sglang.srt.disaggregation.encoder.preprocessor import (
EncoderPreprocessor,
EncoderPreprocessResult,
)
from sglang.srt.disaggregation.encoder.receiver import (
EmbeddingData,
MMReceiverHTTP,
MultiModalEmbeddingData,
_encoder_media_item,
_select_mm_processor_prompt,
)
from sglang.srt.disaggregation.encode_server import MMEncoder, _get_mm_grid_dim
from sglang.srt.disaggregation.encoder.server import MMEncoder
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.managers.tokenizer_manager import (
_reject_missing_dispatched_encoder_embedding,
@@ -220,16 +224,19 @@ def test_epd_allows_local_processing_when_request_was_not_dispatched():
def _encoder(model_type="kimi_k3"):
encoder = MMEncoder.__new__(MMEncoder)
encoder.model_type = model_type
encoder.model_config = SimpleNamespace(
preprocessor = EncoderPreprocessor.__new__(EncoderPreprocessor)
preprocessor.model_type = model_type
preprocessor.model_config = SimpleNamespace(
hf_config=SimpleNamespace(
vision_config=SimpleNamespace(merge_kernel_size=(2, 2))
)
)
encoder.encoder_media_processor_config = (
preprocessor.encoder_media_processor_config = (
KimiK3ForConditionalGeneration.encoder_media_processor_config
if model_type == "kimi_k3"
else EncoderMediaProcessorConfig()
)
encoder.preprocessor = preprocessor
return encoder
@@ -237,11 +244,11 @@ def test_kimi_k3_encoder_normalizes_pillow_images_to_media_dicts():
image = Image.new("RGB", (2, 2))
encoder = _encoder()
assert encoder._grid_count_per_leaf(
assert encoder.preprocessor._grid_count_per_leaf(
[image, {"type": "image", "image": [image, image]}], Modality.IMAGE
) == [1, 2]
normalized = encoder._normalize_kimi_encoder_images(
normalized = encoder.preprocessor._normalize_kimi_encoder_images(
[image, {"type": "image", "image": [image, image]}]
)
assert len(normalized) == 3
@@ -258,14 +265,15 @@ def test_kimi_k3_encoder_passes_media_dicts_to_image_processor():
return {"pixel_values": torch.ones(1, 3), "grid_thws": [[1, 1, 1]]}
encoder = _encoder()
encoder.image_processor = image_processor
encoder.vision_config = {"image": {"return_tensors": "pt"}}
encoder._flatten_and_load_images = AsyncMock(return_value=[image])
encoder.preproc_executor = ThreadPoolExecutor(max_workers=1)
preprocessor = encoder.preprocessor
preprocessor.image_processor = image_processor
preprocessor.vision_config = {"image": {"return_tensors": "pt"}}
preprocessor._flatten_and_load_images = AsyncMock(return_value=[image])
preprocessor.preproc_executor = ThreadPoolExecutor(max_workers=1)
try:
output = asyncio.run(encoder._process_image_items([image], None))
output = asyncio.run(preprocessor._process_image_items([image], None))
finally:
encoder.preproc_executor.shutdown()
preprocessor.preproc_executor.shutdown()
assert "pixel_values" in output
assert output["original_image_sizes"] == [[3, 2]]
@@ -354,27 +362,28 @@ def test_kimi_k3_epd_model_preprocessor_receives_image_processor():
return prepare_kimi_k3_encoder_inputs(mm_data, image_processor)
encoder = _encoder()
encoder.image_processor = image_processor
encoder.use_image_processor_gpu = False
encoder.vision_config = {"image": {"return_tensors": "pt"}}
encoder._flatten_and_load_images = AsyncMock(return_value=[image])
encoder.preproc_executor = ThreadPoolExecutor(max_workers=1)
preprocessor = encoder.preprocessor
preprocessor.image_processor = image_processor
preprocessor.use_image_processor_gpu = False
preprocessor.vision_config = {"image": {"return_tensors": "pt"}}
preprocessor._flatten_and_load_images = AsyncMock(return_value=[image])
preprocessor.preproc_executor = ThreadPoolExecutor(max_workers=1)
try:
with patch(
"sglang.srt.disaggregation.encode_server.get_parallel",
"sglang.srt.disaggregation.encoder.preprocessor.get_parallel",
return_value=SimpleNamespace(attn_tp_rank=0, attn_tp_size=1),
):
output = asyncio.run(
encoder._process_image_items([image], model_preprocessor)
preprocessor._process_image_items([image], model_preprocessor)
)
finally:
encoder.preproc_executor.shutdown()
preprocessor.preproc_executor.shutdown()
assert len(calls) == 1
assert calls[0][0][0] == {"type": "image", "image": image}
assert calls[0][1:] == (
Modality.IMAGE,
encoder.vision_config,
preprocessor.vision_config,
image_processor,
False,
)
@@ -481,13 +490,13 @@ def test_kimi_k3_epd_selects_matching_jpeg_decode_mode(
):
expected = torch.zeros((3, 2, 3), dtype=torch.uint8)
encoder = _encoder()
encoder.use_image_processor_gpu = use_image_processor_gpu
encoder.preprocessor.use_image_processor_gpu = use_image_processor_gpu
with patch(
"sglang.srt.disaggregation.encode_server.load_image",
"sglang.srt.disaggregation.encoder.preprocessor.load_image",
return_value=(expected, None),
) as load:
output = encoder._load_single_item(b"jpeg", Modality.IMAGE)
output = encoder.preprocessor._load_single_item(b"jpeg", Modality.IMAGE)
assert output is expected
load.assert_called_once_with(b"jpeg", expected_decode_mode)
@@ -498,13 +507,13 @@ def test_kimi_k3_epd_verifies_content_hash_before_decode():
digest = snapshot_media(payload).content_digest
expected = torch.zeros((3, 2, 3), dtype=torch.uint8)
encoder = _encoder()
encoder.use_image_processor_gpu = False
encoder.preprocessor.use_image_processor_gpu = False
with patch(
"sglang.srt.disaggregation.encode_server.load_image",
"sglang.srt.disaggregation.encoder.preprocessor.load_image",
return_value=(expected, None),
) as load:
output = encoder._load_single_item(
output = encoder.preprocessor._load_single_item(
{"url": payload, "content_hash": digest}, Modality.IMAGE
)
@@ -589,8 +598,9 @@ def test_kimi_k3_encoder_prefers_grid_thws_and_uses_temporal_pool_length():
stale_grid = torch.tensor([[1, 2, 2]])
mm_inputs = {"grid_thws": grid_thws, "image_grid_thw": stale_grid}
assert _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_k3") is grid_thws
assert _encoder().get_num_tokens(grid_thws[0], Modality.IMAGE) == 24
preprocessor = _encoder().preprocessor
assert preprocessor._get_mm_grid_dim(mm_inputs, Modality.IMAGE) is grid_thws
assert preprocessor.get_num_tokens(grid_thws[0], Modality.IMAGE) == 24
def test_kimi_k3_encoder_splits_cross_request_batch_into_single_grid_items():
@@ -606,12 +616,14 @@ def test_kimi_k3_encoder_splits_cross_request_batch_into_single_grid_items():
output = encoder._encode_missing(
feature,
{"pixel_values": feature, "grid_thws": grid_thws},
EncoderPreprocessResult(
mm_inputs={"pixel_values": feature, "grid_thws": grid_thws},
grid_thw=grid_thws,
token_counts=[1, 2, 2],
),
indices=[2, 0, 1],
modality=Modality.IMAGE,
get_feature_fn=get_feature_fn,
grid_thw=grid_thws,
keep_on_gpu=True,
)
items = captured["items"]
@@ -643,7 +655,7 @@ def test_encoder_preprocessed_items_follow_dp_owner_selection_order():
{"pixel_values": [item.feature for item in items], "grid_thws": grid_thws},
mm_items=items,
)
embeddings = torch.arange(4, dtype=torch.float32).reshape(4, 1)
embeddings = torch.arange(3, dtype=torch.float32).reshape(3, 1)
captured = {}
def get_feature_fn(selected_items):
@@ -652,17 +664,19 @@ def test_encoder_preprocessed_items_follow_dp_owner_selection_order():
output = encoder._encode_missing(
mm_inputs["pixel_values"],
mm_inputs,
EncoderPreprocessResult(
mm_inputs=mm_inputs,
grid_thw=grid_thws,
token_counts=[1, 2, 2],
),
indices=[2, 0],
modality=Modality.IMAGE,
get_feature_fn=get_feature_fn,
grid_thw=grid_thws,
keep_on_gpu=True,
)
assert captured["items"] == [items[2], items[0]]
assert [part.shape[0] for part in output] == [2, 1]
torch.testing.assert_close(torch.cat(output), embeddings[:3])
torch.testing.assert_close(torch.cat(output), embeddings)
def test_encoder_preprocessed_items_hash_individually():
@@ -767,6 +781,8 @@ def test_epd_encoder_reuses_scheduler_zmq_peer():
)
with config_override as server_args:
encoder.server_args = server_args
encoder.transfer_backend = "zmq_to_scheduler"
encoder.use_mooncake = False
encoder.send_timeout = 3
encoder.context = context
encoder.scheduler_send_sockets = {}
@@ -841,6 +857,8 @@ def test_epd_encoder_pipelines_zero_copy_sends_per_peer():
)
with config_override as server_args:
encoder.server_args = server_args
encoder.transfer_backend = "zmq_to_scheduler"
encoder.use_mooncake = False
encoder.send_timeout = 1
encoder.context = FakeContext(socket)
encoder.scheduler_send_sockets = {}
@@ -108,7 +108,7 @@ _CONFIGURED_SIZE_CALL_SITES = {
"the consumer count is configured fan-out arithmetic (tp_size // "
"dp_size), which is what the record answered before"
),
("srt/disaggregation/encode_server.py", "configured_tp_size"): (
("srt/disaggregation/encoder/runtime.py", "configured_tp_size"): (
"the encode server's launch entry sizes its workers before it has "
"spawned any of them"
),
@@ -77,8 +77,8 @@ _KNOWN_ENTRIES = frozenset(
"run_data_parallel_controller_process",
),
("srt/ray/scheduler_actor.py", "__init__"),
("srt/disaggregation/encode_server.py", "__init__"),
("srt/disaggregation/encode_server.py", "launch_server"),
("srt/disaggregation/encoder/server.py", "__init__"),
("srt/disaggregation/encoder/http_server.py", "launch_server"),
("srt/managers/tokenizer_manager.py", "__init__"),
("srt/entrypoints/engine.py", "_launch_subprocesses"),
(