[EPD][Perf] parallelize ZMQ send for encode server (#16487)

Co-authored-by: siyu <liusy58@linux.alibaba.com>
This commit is contained in:
Zheng Wengang
2026-01-31 14:30:11 +08:00
committed by GitHub
co-authored by siyu
parent 04efd03dbf
commit a4df95c15f
@@ -40,6 +40,7 @@ from sglang.srt.server_args import (
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
from sglang.srt.utils import ( from sglang.srt.utils import (
config_socket,
get_local_ip_auto, get_local_ip_auto,
get_zmq_socket, get_zmq_socket,
load_audio, load_audio,
@@ -186,6 +187,8 @@ class MMEncoder:
) )
self.context = zmq.asyncio.Context(2) self.context = zmq.asyncio.Context(2)
self.sync_context = zmq.Context() # Reuse sync context for thread pool
self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=10)
embedding_cache_size = int(os.environ.get("SGLANG_VLM_CACHE_SIZE_MB", "4096")) embedding_cache_size = int(os.environ.get("SGLANG_VLM_CACHE_SIZE_MB", "4096"))
self.mm_cache = MultiModalStaticCache(embedding_cache_size * 1024 * 1024) self.mm_cache = MultiModalStaticCache(embedding_cache_size * 1024 * 1024)
@@ -389,25 +392,35 @@ class MMEncoder:
else f"tcp://{prefill_host}:{embedding_port}" else f"tcp://{prefill_host}:{embedding_port}"
) )
logger.info(f"{endpoint = }") logger.info(f"{endpoint = }")
socket = get_zmq_socket(
self.context,
zmq.PUSH,
endpoint,
False,
)
# Serialize data
if self.server_args.encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
socket.send_multipart([pickle.dumps(mm_data)]) serialized_data = pickle.dumps(mm_data)
buffer = None
else: else:
new_mm_data = mm_data.copy_without_embedding() new_mm_data = mm_data.copy_without_embedding()
if new_mm_data.error_msg is not None: if new_mm_data.error_msg is not None:
socket.send_multipart([pickle.dumps(new_mm_data)]) buffer = None
return serialized_data = pickle.dumps(new_mm_data)
else:
embedding_tensor = TensorWrapper(mm_data.embedding) embedding_tensor = TensorWrapper(mm_data.embedding)
socket.send_multipart( serialized_data = pickle.dumps(new_mm_data)
[pickle.dumps(new_mm_data), embedding_tensor.__buffer__()] buffer = embedding_tensor.__buffer__()
)
# Use thread pool executor for parallel ZMQ send operations
def send_with_socket():
sock = self.sync_context.socket(zmq.PUSH)
config_socket(sock, zmq.PUSH)
try:
sock.connect(endpoint)
if buffer is not None:
sock.send_multipart([serialized_data, buffer], copy=False)
else:
sock.send_multipart([serialized_data], copy=False)
finally:
sock.close()
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
async def encode(self, mm_items, req_id, num_parts, part_idx): async def encode(self, mm_items, req_id, num_parts, part_idx):
try: try: