[EPD] Optimize the Mooncake backend (#22587)

Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
LucQueen
2026-05-29 10:42:24 +08:00
committed by GitHub
co-authored by ZhengWG
parent 569ee93357
commit 36d0a6e08e
11 changed files with 1516 additions and 159 deletions
@@ -1,17 +1,17 @@
import asyncio import asyncio
import itertools import itertools
import logging import logging
import pickle
import random import random
import threading import threading
import time import time
import uuid import uuid
import weakref
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from array import array from array import array
from collections import OrderedDict, defaultdict from collections import OrderedDict, defaultdict
from enum import IntEnum from enum import IntEnum
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Dict, List, Optional from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import aiohttp import aiohttp
import numpy as np import numpy as np
@@ -30,6 +30,7 @@ from sglang.srt.managers.multimodal_processor import get_mm_processor, import_pr
from sglang.srt.managers.schedule_batch import Modality, Req from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData from sglang.srt.utils import ImageData
from sglang.srt.utils.common import safe_pickle_loads
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
from sglang.srt.utils.network import ( from sglang.srt.utils.network import (
NetworkAddress, NetworkAddress,
@@ -153,6 +154,7 @@ class EmbeddingData:
self.shape = embedding_shape self.shape = embedding_shape
else: else:
self.shape = list(embedding.shape) if embedding is not None else None self.shape = list(embedding.shape) if embedding is not None else None
self.cached_embedding = None
self.error_msg = error_msg self.error_msg = error_msg
self.error_code = error_code self.error_code = error_code
# Store additional metadata (e.g., video_timestamps for qwen3_vl) # Store additional metadata (e.g., video_timestamps for qwen3_vl)
@@ -182,7 +184,8 @@ class EmbeddingData:
error_code=self.error_code, error_code=self.error_code,
) )
for key, value in self.__dict__.items(): for key, value in self.__dict__.items():
if key.startswith("_") or key == "embedding": # cached_embedding is a GPU tensor used only by mooncake's in-process
if key.startswith("_") or key in ("embedding", "cached_embedding"):
continue continue
setattr(new_data, key, value) setattr(new_data, key, value)
return new_data return new_data
@@ -468,6 +471,7 @@ class WaitingImageRequest:
self.start_time = time.time() self.start_time = time.time()
def send_encode_request(self): def send_encode_request(self):
async def _send_single_request(session, url, payload): async def _send_single_request(session, url, payload):
try: try:
async with session.post(url, json=payload) as response: async with session.post(url, json=payload) as response:
@@ -479,7 +483,9 @@ class WaitingImageRequest:
async def send_embedding_port(req_id, receive_count, host_name, embedding_port): async def send_embedding_port(req_id, receive_count, host_name, embedding_port):
async with aiohttp.ClientSession( async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(total=1800) timeout=aiohttp.ClientTimeout(
total=envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()
)
) as session: ) as session:
tasks = [] tasks = []
logger.info(f"{self.num_items_assigned = } ") logger.info(f"{self.num_items_assigned = } ")
@@ -525,8 +531,18 @@ class WaitingImageRequest:
results = await asyncio.gather(*tasks, return_exceptions=True) results = await asyncio.gather(*tasks, return_exceptions=True)
for i, result in enumerate(results): for i, result in enumerate(results):
if isinstance(result, Exception): if isinstance(result, asyncio.TimeoutError):
logger.error(f"Request {i} failed: {result}") timeout_val = envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()
logger.error(
f"Request {i} to encoder /scheduler_receive_url timed out "
f"({timeout_val}s) for req_id={req_id}"
)
elif isinstance(result, Exception):
logger.error(
f"Request {i} to encoder /scheduler_receive_url failed for "
f"req_id={req_id}: {result}",
exc_info=result,
)
else: else:
logger.debug(f"Request {i} succeeded.") logger.debug(f"Request {i} succeeded.")
@@ -548,7 +564,7 @@ class WaitingImageRequest:
except zmq.Again: except zmq.Again:
# No data available yet, wait a bit and retry # No data available yet, wait a bit and retry
return return
recv_obj: EmbeddingData = pickle.loads(parts[0]) recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None: if getattr(recv_obj, "error_msg", None) is not None:
logger.warning( logger.warning(
f"Received error signal from encoder for {self.rid}: {recv_obj.error_msg} {recv_obj.error_code = }" f"Received error signal from encoder for {self.rid}: {recv_obj.error_msg} {recv_obj.error_code = }"
@@ -595,6 +611,9 @@ class WaitingImageRequest:
self.status = WaitingImageRequestStatus.SUCCESS self.status = WaitingImageRequestStatus.SUCCESS
self.recv_socket.close() self.recv_socket.close()
def _cleanup_gpu_buffer(self):
pass
class WaitingImageRequestGrpc(WaitingImageRequest): class WaitingImageRequestGrpc(WaitingImageRequest):
def send_encode_request(self): def send_encode_request(self):
@@ -643,6 +662,522 @@ class WaitingImageRequestGrpc(WaitingImageRequest):
) )
class WaitingImageRDMARequest(WaitingImageRequest):
def __init__(
self,
rid,
recv_req,
mm_processor,
encoder_urls,
host_name,
receive_count,
embeddings_engine,
dtype,
gpu_id=0,
model_type: Optional[str] = None,
embedding_pool=None,
):
super().__init__(
rid=rid,
recv_req=recv_req,
mm_processor=mm_processor,
encoder_urls=encoder_urls,
model_type=model_type,
host_name=host_name,
receive_count=receive_count,
)
self.embeddings_engine = embeddings_engine
self.dtype = dtype
self.gpu_id = gpu_id
self.embeddings_buffer = None
self.embedding_pool = embedding_pool
self._buffer_from_pool = False
self._pool_slot_id: Optional[int] = None
def send_encode_request(self):
self._encode_thread = threading.Thread(
target=self._run_encode_in_thread, daemon=True
)
self._encode_thread.start()
def _run_encode_in_thread(self):
try:
asyncio.run(self._send_encode_and_rdma_request())
except Exception as e:
logger.error(f"RDMA encode request failed for rid={self.rid}: {e}")
self.status = WaitingImageRequestStatus.FAIL
self.error_msg = str(e)
self._cleanup_gpu_buffer()
self.recv_socket.close()
async def _send_encode_and_rdma_request(self):
modalities = list(self.num_items_assigned.keys())
_, modality_num_parts = calculate_modality_num_parts(
modalities, self.num_items_assigned
)
encode_requests = []
# Use the URL list captured at tokenizer time. TokenizedGenerateReqInput
# has no image_data field, so reading recv_req.image_data here would
# always return None and produce empty mm_items.
mm_data_all = self.recv_req.mm_data_mooncake or []
total_num_parts = sum(modality_num_parts.values())
part_idx_offset = 0
for modality in modalities:
assigned_nums = self.num_items_assigned[modality]
num_parts = modality_num_parts[modality]
mm_data_modality = [d for d in mm_data_all if d["modality"] == modality]
cum_num_items = 0
cum_idx = 0
for idx, assigned_num in enumerate(assigned_nums):
if assigned_num == 0:
continue
part_idx = part_idx_offset + cum_idx
part_req_id = create_part_req_id(self.recv_req.rid, part_idx)
encode_requests.append(
{
"encoder_idx": idx,
"mm_items": [
d["url"]
for d in mm_data_modality[
cum_num_items : cum_num_items + assigned_num
]
],
"num_parts": total_num_parts,
"part_idx": part_idx,
"req_id": part_req_id,
"modality": modality.name,
"prefill_host": self.host_name,
"embedding_port": self.embedding_port,
}
)
cum_idx += 1
cum_num_items += assigned_num
part_idx_offset += num_parts
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(total=envs.SGLANG_ENCODER_HTTP_TIMEOUT.get())
) as session:
# Phase 1: POST /encode to all encoder shards in parallel.
tasks = [
session.post(
f"{self.encoder_urls[r['encoder_idx']]}/encode",
json=r,
)
for r in encode_requests
]
responses = await asyncio.gather(*tasks, return_exceptions=True)
if not await self._check_encoder_responses(responses, "/encode"):
return
response_json_list = [await r.json() for r in responses]
# Sort by part_idx
embedding_sizes, response_sorted, total_bytes = (
_sort_responses_and_compute_total_bytes(
response_json_list, total_num_parts
)
)
# Phase 2: Pre-allocate GPU landing buffer.
# Prefer the pre-registered persistent pool when available; this avoids
# per-request register/deregister and keeps the encoder's openSegment
if total_bytes > 0:
if self.embedding_pool is not None:
alloc_result = await asyncio.to_thread(
self.embedding_pool.alloc, total_bytes
)
if alloc_result is None:
# Either the request exceeds pool capacity outright, or
# the wait timed out. Both are fatal for this request
# — fall through to error handling.
self.status = WaitingImageRequestStatus.FAIL
self.error_msg = (
f"MooncakeEmbeddingPool could not allocate "
f"{total_bytes // (1024 * 1024)}MB (oversize or "
f"timeout). Raise SGLANG_EMBEDDING_POOL_SIZE_MB."
)
self.recv_socket.close()
return
pool_view, buffer_address, slot_id = alloc_result
self.embeddings_buffer = pool_view
self._buffer_from_pool = True
self._pool_slot_id = slot_id
logger.info(
f"Pool-allocated Mooncake GPU landing buffer: "
f"rid={self.rid}, size={total_bytes}, "
f"addr={buffer_address}, slot={slot_id}"
)
else:
gpu_buffer = torch.empty(
total_bytes, dtype=torch.uint8, device=f"cuda:{self.gpu_id}"
)
self.embeddings_engine.register(
gpu_buffer.data_ptr(), gpu_buffer.nbytes
)
self.embeddings_buffer = gpu_buffer
buffer_address = gpu_buffer.data_ptr()
self._buffer_from_pool = False
logger.info(
f"Per-request registered Mooncake GPU landing buffer "
f"(pool disabled): rid={self.rid}, size={total_bytes}, "
f"addr={buffer_address}"
)
else:
self.embeddings_buffer = None
buffer_address = 0
# Phase 2 cont: POST /send with RDMA info.
offset = 0
send_tasks = []
for idx in range(total_num_parts):
rj = response_sorted[idx]
encoder_idx = rj.pop("encoder_idx", None)
rj.update(
{
"session_id": self.embeddings_engine.session_id,
"buffer_address": offset + buffer_address,
}
)
send_tasks.append(
session.post(
f"{self.encoder_urls[encoder_idx]}/send",
json=rj,
)
)
offset += embedding_sizes[idx]
# Phase 3: Wait for RDMA transfers to complete
send_responses = await asyncio.gather(*send_tasks, return_exceptions=True)
if not await self._check_encoder_responses(
send_responses, "/send", on_error=self._cleanup_gpu_buffer
):
return
logger.info(f"RDMA transfers completed for rid={self.rid}")
async def _check_encoder_responses(self, responses, endpoint: str, on_error=None):
"""Validate gathered HTTP responses from the encoder.
Marks the request as FAIL and closes the recv socket on the first error,
invoking ``on_error`` (e.g. GPU buffer cleanup) before closing.
Returns True if all responses succeeded.
"""
for i, resp in enumerate(responses):
msg = None
if isinstance(resp, asyncio.TimeoutError):
timeout_val = envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()
logger.error(
f"Encoder {endpoint} timeout ({timeout_val}s) for rid={self.rid} "
f"(request {i})"
)
msg = f"Encoder {endpoint} timeout ({timeout_val}s)"
elif isinstance(resp, Exception):
logger.error(
f"Encoder {endpoint} failed for rid={self.rid} (request {i}): {resp}",
exc_info=resp,
)
msg = str(resp)
elif resp.status != 200:
try:
err = await resp.json()
msg = err.get("message", "Unknown error")
except Exception:
msg = await resp.text()
logger.error(f"Encoder {endpoint} returned error {resp.status}: {msg}")
if msg is not None:
self.status = WaitingImageRequestStatus.FAIL
self.error_msg = msg
if on_error is not None:
on_error()
self.recv_socket.close()
return False
return True
def _try_recv_mm_data(self):
"""Extract embedding from GPU buffer after RDMA transfer."""
if self.status != WaitingImageRequestStatus.PENDING:
return
while self.recv_embedding_data is None or not self.recv_embedding_data.ready:
try:
parts = self.recv_socket.recv_multipart(flags=zmq.NOBLOCK, copy=False)
except zmq.Again:
return
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None:
logger.warning(f"Received error for {self.rid}: {recv_obj.error_msg}")
self.error_msg = recv_obj.error_msg
self.error_code = recv_obj.error_code
self.status = WaitingImageRequestStatus.FAIL
self._cleanup_gpu_buffer()
self.recv_socket.close()
return
# Extract original req_id
part_req_id = recv_obj.req_id
original_req_id = extract_original_req_id(part_req_id)
if original_req_id != self.recv_req.rid:
logger.warning(
f"Dropping stale embedding data: expected rid={self.recv_req.rid}, "
f"got rid={recv_obj.req_id} (likely from ZMQ port reuse)"
)
continue
recv_obj.req_id = original_req_id
# Embedding was written directly into pre-registered GPU buffer by encode server
# (Mooncake GPU-direct transfer); no ZMQ payload in this message.
# recv_obj.embedding stays None until we extract from GPU buffer below
if self.recv_embedding_data is None:
self.recv_embedding_data = MultiModalEmbeddingData.from_embedding_data(
recv_obj
)
else:
self.recv_embedding_data.add(recv_obj)
# Zero-copy: build per-modality views directly from the pre-registered
# GPU buffer. Skips the per-part split + torch.cat round-trip — both
# the extra GPU allocation and the D2D copy — so mm_item.precomputed_
# embeddings ends up referencing the pool buffer. Slot lifetime is
# bound to mm_inputs GC via weakref.finalize below.
if self.embeddings_buffer is not None:
recv_embedding = _view_pool_buffer_by_modality(
self.embeddings_buffer, self.recv_embedding_data, self.dtype
)
else:
recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True)
mm_inputs = self.mm_processor.get_mm_data(
self.recv_req.input_text,
recv_embedding,
**self.recv_embedding_data.get_mm_extra_meta(),
)
# Bind slot release to mm_inputs GC
if self._buffer_from_pool and mm_inputs is not None:
weakref.finalize(mm_inputs, self.embedding_pool.release, self._pool_slot_id)
for item in getattr(mm_inputs, "mm_items", []) or []:
try:
setattr(item, "_keep_device_embedding", True)
except Exception:
pass
# Detach so _cleanup_gpu_buffer no-ops; finalize now owns release.
self._pool_slot_id = None
self.embeddings_buffer = None
self._buffer_from_pool = False
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = mm_inputs.input_ids
self.status = WaitingImageRequestStatus.SUCCESS
self._cleanup_gpu_buffer()
self.recv_socket.close()
def _cleanup_gpu_buffer(self):
"""Deregister and release the GPU buffer."""
if self.embeddings_buffer is not None:
# Pool-backed views share the pre-registered backing tensor; just
# release the slot back to the pool so a queued alloc can proceed.
if self._buffer_from_pool:
if self._pool_slot_id is not None and self.embedding_pool is not None:
self.embedding_pool.release(self._pool_slot_id)
self._pool_slot_id = None
self.embeddings_buffer = None
return
try:
self.embeddings_engine.deregister(self.embeddings_buffer.data_ptr())
except Exception:
logger.exception("Failed to deregister GPU buffer for rid=%s", self.rid)
self.embeddings_buffer = None
def _sort_responses_and_compute_total_bytes(response_json_list, total_num_parts):
"""Sort responses by part_idx and compute total embedding bytes."""
embedding_sizes = [None] * total_num_parts
response_sorted = [None] * total_num_parts
for rj in response_json_list:
idx = rj["part_idx"]
embedding_sizes[idx] = rj["embedding_size"]
response_sorted[idx] = rj
total_bytes = sum(s for s in embedding_sizes if s is not None)
return embedding_sizes, response_sorted, total_bytes
class MooncakeEmbeddingPool:
"""Persistent GPU buffer pool registered once with the Mooncake engine.
Allocator: first-fit on a free-segment list with 256-byte alignment.
`alloc()` blocks on a Condition when the pool is full and resumes once
a peer `release()`s a slot. Each successful alloc returns a slot_id
that must be passed back to release() when the consumer is done with
the buffer (after RDMA write completes and the data has been read).
"""
_ALIGN = 256
def __init__(self, engine, gpu_id: int, size_bytes: int):
self.engine = engine
self.gpu_id = gpu_id
self.size_bytes = size_bytes
self.buffer = torch.empty(
size_bytes, dtype=torch.uint8, device=f"cuda:{gpu_id}"
)
self.base = self.buffer.data_ptr()
self.engine.register(self.base, self.buffer.nbytes)
self._segments_free: List[Tuple[int, int]] = [(0, size_bytes)]
self._inflight: Dict[int, Tuple[int, int]] = {}
self._next_slot_id = 0
self._total_inflight = 0
self._lock = threading.Lock()
self._cond = threading.Condition(self._lock)
logger.info(
f"MooncakeEmbeddingPool registered: gpu={gpu_id}, "
f"size={size_bytes // (1024 * 1024)}MB, base=0x{self.base:x}"
)
def alloc(
self, nbytes: int, timeout: float = 60.0
) -> Optional[Tuple[torch.Tensor, int, int]]:
"""Allocate `nbytes` from the pool.
Returns ``(tensor_view, gpu_addr, slot_id)`` on success, or ``None``
when (a) the request is bigger than the pool itself or (b) the wait
for a free slot exceeds ``timeout`` seconds.
When the pool is full of in-flight slots, this call blocks the
calling thread on a Condition until a peer ``release()`` opens
enough contiguous space.
NOTE: no ordering guarantee — notify_all + lock race means
large requests can starve behind small ones, plus thundering-herd.
"""
if nbytes > self.size_bytes:
logger.error(
f"MooncakeEmbeddingPool: requested {nbytes // (1024 * 1024)}MB "
f"exceeds pool capacity {self.size_bytes // (1024 * 1024)}MB. "
f"Raise SGLANG_EMBEDDING_POOL_SIZE_MB."
)
return None
aligned = (nbytes + self._ALIGN - 1) & ~(self._ALIGN - 1)
deadline = time.monotonic() + timeout
warned = False
with self._cond:
while True:
slot = self._try_alloc_locked(nbytes, aligned)
if slot is not None:
return slot
if not warned:
inflight_mb = self._total_inflight // (1024 * 1024)
cap_mb = self.size_bytes // (1024 * 1024)
logger.warning(
f"MooncakeEmbeddingPool full: "
f"{inflight_mb}/{cap_mb}MB in-flight across "
f"{len(self._inflight)} requests; queueing a "
f"{nbytes // (1024 * 1024)}MB request. Raise "
f"SGLANG_EMBEDDING_POOL_SIZE_MB if this is frequent."
)
warned = True
remaining = deadline - time.monotonic()
if remaining <= 0:
logger.error(
f"MooncakeEmbeddingPool alloc timed out after "
f"{timeout}s waiting for {nbytes // (1024 * 1024)}MB."
)
return None
self._cond.wait(timeout=remaining)
def _try_alloc_locked(
self, nbytes: int, aligned: int
) -> Optional[Tuple[torch.Tensor, int, int]]:
for i, (off, length) in enumerate(self._segments_free):
if length >= aligned:
if length == aligned:
self._segments_free.pop(i)
else:
self._segments_free[i] = (off + aligned, length - aligned)
slot_id = self._next_slot_id
self._next_slot_id += 1
self._inflight[slot_id] = (off, aligned)
self._total_inflight += aligned
view = self.buffer[off : off + nbytes]
return view, self.base + off, slot_id
return None
def release(self, slot_id: int) -> None:
"""Return a previously-allocated slot to the free list and wake any
blocked alloc() waiters."""
with self._cond:
seg = self._inflight.pop(slot_id, None)
if seg is None:
return
off, aligned = seg
self._total_inflight -= aligned
self._coalesce_free_locked(off, aligned)
self._cond.notify_all()
def _coalesce_free_locked(self, off: int, length: int) -> None:
self._segments_free.append((off, length))
self._segments_free.sort()
merged: List[Tuple[int, int]] = []
for s_off, s_len in self._segments_free:
if merged and merged[-1][0] + merged[-1][1] == s_off:
p_off, p_len = merged[-1]
merged[-1] = (p_off, p_len + s_len)
else:
merged.append((s_off, s_len))
self._segments_free = merged
def _slice_embedding_buffer(raw_buffer, embedding_data, dtype):
"""Slice a flat GPU buffer into per-part embedding tensors in-place."""
elem_size = torch.tensor([], dtype=dtype).element_size()
byte_offset = 0
for i in range(embedding_data.num_parts):
shape = embedding_data.embedding_shape_list[i]
if shape is None:
continue
part_bytes = shape[0] * shape[1] * elem_size
embedding_data.embedding_list[i] = (
raw_buffer[byte_offset : byte_offset + part_bytes]
.view(dtype)
.reshape(shape)
)
byte_offset += part_bytes
def _view_pool_buffer_by_modality(raw_buffer, embedding_data, dtype):
"""Zero-copy view of raw_buffer as {modality: [total_tokens, hidden]}.
Replaces _slice_embedding_buffer + get_embedding(is_concat=True): parts of
the same modality are contiguous in raw_buffer (encoder writes them
modality-outer in _send_encode_and_rdma_request), so we can reshape the
byte range directly — no per-part split, no torch.cat copy.
Caller must keep raw_buffer's storage alive while the returned views are
in use. The pool path binds slot release to mm_inputs GC via finalize.
"""
elem_size = torch.tensor([], dtype=dtype).element_size()
# mod -> [byte_start, byte_end, total_tokens, hidden]
mod_info: Dict[Modality, List[int]] = {}
off = 0
for i in range(embedding_data.num_parts):
shape = embedding_data.embedding_shape_list[i]
if shape is None:
continue
nbytes = shape[0] * shape[1] * elem_size
mod = embedding_data.modality_list[i]
info = mod_info.get(mod)
if info is None:
mod_info[mod] = [off, off + nbytes, shape[0], shape[1]]
else:
assert (
info[3] == shape[1]
), f"hidden_dim mismatch in modality {mod}: {info[3]} vs {shape[1]}"
assert info[1] == off, f"non-contiguous parts in modality {mod}"
info[1] = off + nbytes
info[2] += shape[0]
off += nbytes
return {
mod: raw_buffer[s:e].view(dtype).reshape(tokens, hidden)
for mod, (s, e, tokens, hidden) in mod_info.items()
}
def _determine_tensor_transport_mode(server_args): def _determine_tensor_transport_mode(server_args):
is_cross_node = server_args.dist_init_addr is_cross_node = server_args.dist_init_addr
@@ -669,6 +1204,16 @@ class MMReceiverBase(ABC):
self.encode_urls = server_args.encoder_urls self.encode_urls = server_args.encoder_urls
self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
self.host = get_local_ip_auto(server_args.host) self.host = get_local_ip_auto(server_args.host)
self.pp_rank = pp_rank
self.tp_rank = tp_rank
self.tp_size = server_args.tp_size
self.tp_group = tp_group
self.nnodes = server_args.nnodes
self.hostname = get_local_ip_auto()
self.waiting_list: List[WaitingImageRequest] = []
self.scheduler = scheduler
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
self.model_type = ( self.model_type = (
getattr(hf_config, "model_type", "").lower() getattr(hf_config, "model_type", "").lower()
if hf_config is not None if hf_config is not None
@@ -690,19 +1235,48 @@ class MMReceiverBase(ABC):
), ),
) )
self.embeddings_buffer = dict() self.embeddings_buffer = dict()
elif self.encoder_transfer_backend == "zmq_to_scheduler": self.embedding_pool = None
self.pp_rank = pp_rank pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get()
self.tp_rank = tp_rank if pool_mb and pool_mb > 0 and scheduler is not None:
self.tp_size = server_args.tp_size gpu_id = getattr(scheduler, "gpu_id", 0)
self.tp_group = tp_group try:
self.nnodes = server_args.nnodes self.embedding_pool = MooncakeEmbeddingPool(
self.hostname = get_local_ip_auto() self.embeddings_engine, gpu_id, pool_mb * 1024 * 1024
self.waiting_list: List[WaitingImageRequest] = [] )
self.scheduler = scheduler except Exception:
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() logger.exception(
"Failed to allocate MooncakeEmbeddingPool, "
"falling back to per-request register"
)
self.embedding_pool = None
if hf_config is not None: if hf_config is not None:
self._init_mm_processor(server_args, hf_config)
elif self.encoder_transfer_backend == "zmq_to_scheduler":
if hf_config is not None:
self._init_mm_processor(
server_args,
hf_config,
model_config=(
getattr(self.scheduler, "model_config", None)
if self.scheduler is not None
else None
),
)
def _init_mm_processor(
self,
server_args: "ServerArgs",
hf_config: "PretrainedConfig",
model_config=None,
):
"""Load processor and initialize mm_processor, shared by all backends."""
transport_mode = _determine_tensor_transport_mode(server_args) transport_mode = _determine_tensor_transport_mode(server_args)
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
extra_kwargs = {}
if getattr(server_args, "tokenizer_backend", None) is not None:
extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend
_processor = None _processor = None
try: try:
_processor = get_processor( _processor = get_processor(
@@ -711,7 +1285,7 @@ class MMReceiverBase(ABC):
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, revision=server_args.revision,
use_fast=not server_args.disable_fast_image_processor, use_fast=not server_args.disable_fast_image_processor,
tokenizer_backend=server_args.tokenizer_backend, **extra_kwargs,
) )
except ValueError as e: except ValueError as e:
error_message = str(e) error_message = str(e)
@@ -725,26 +1299,24 @@ class MMReceiverBase(ABC):
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, revision=server_args.revision,
use_fast=True, use_fast=True,
tokenizer_backend=server_args.tokenizer_backend, **extra_kwargs,
) )
else: else:
raise e raise e
# Skip mm_pool if not adaptive dispatch to encoder
enable_adaptive_dispatch_to_encoder = ( enable_adaptive_dispatch_to_encoder = (
server_args.enable_adaptive_dispatch_to_encoder server_args.enable_adaptive_dispatch_to_encoder
) )
mm_processor_kwargs = {}
if model_config is not None:
mm_processor_kwargs["model_config"] = model_config
self.mm_processor = get_mm_processor( self.mm_processor = get_mm_processor(
hf_config, hf_config,
server_args, server_args,
_processor, _processor,
transport_mode, transport_mode,
model_config=(
getattr(self.scheduler, "model_config", None)
if self.scheduler is not None
else None
),
skip_mm_pool=not enable_adaptive_dispatch_to_encoder, skip_mm_pool=not enable_adaptive_dispatch_to_encoder,
**mm_processor_kwargs,
) )
@abstractmethod @abstractmethod
@@ -814,7 +1386,7 @@ class MMReceiverBase(ABC):
parts = await recv_socket.recv_multipart(copy=False) parts = await recv_socket.recv_multipart(copy=False)
if not parts: if not parts:
continue continue
recv_obj: EmbeddingData = pickle.loads(parts[0]) recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None: if getattr(recv_obj, "error_msg", None) is not None:
logger.warning( logger.warning(
f"Encoder error for req_id={req_id}: {recv_obj.error_msg} " f"Encoder error for req_id={req_id}: {recv_obj.error_msg} "
@@ -859,22 +1431,7 @@ class MMReceiverBase(ABC):
return None return None
raw_buffer = self.embeddings_buffer.pop(req_id) raw_buffer = self.embeddings_buffer.pop(req_id)
self.embeddings_engine.deregister(raw_buffer.data_ptr()) self.embeddings_engine.deregister(raw_buffer.data_ptr())
byte_offset = 0 _slice_embedding_buffer(raw_buffer, recv_embedding_data, self.dtype)
for i in range(recv_embedding_data.num_parts):
shape = recv_embedding_data.embedding_shape_list[i]
if shape is None:
continue
part_bytes = (
shape[0]
* shape[1]
* torch.tensor([], dtype=self.dtype).element_size()
)
recv_embedding_data.embedding_list[i] = (
raw_buffer[byte_offset : byte_offset + part_bytes]
.view(self.dtype)
.reshape(shape)
)
byte_offset += part_bytes
recv_embedding = recv_embedding_data.get_embedding(is_concat=True) recv_embedding = recv_embedding_data.get_embedding(is_concat=True)
@@ -902,6 +1459,16 @@ class MMReceiverBase(ABC):
mm_data, len(self.encode_urls) mm_data, len(self.encode_urls)
) )
obj.num_items_assigned = num_items_assigned obj.num_items_assigned = num_items_assigned
# For mooncake, No tokenizer-side thread.
# Save mm_data (extracted URL list) onto obj so the scheduler-side
# WaitingImageRDMARequest can use it. TokenizedGenerateReqInput does
# NOT carry image_data, so re-reading recv_req.image_data at scheduler
# time would always return None.
if self.encoder_transfer_backend == "mooncake":
obj.mm_data_mooncake = mm_data
return
encode_thread = threading.Thread( encode_thread = threading.Thread(
target=self._run_encode_in_thread, target=self._run_encode_in_thread,
args=( args=(
@@ -916,7 +1483,7 @@ class MMReceiverBase(ABC):
encode_thread.start() encode_thread.start()
# For zmq_to_scheduler # For zmq_to_scheduler
def _process_waiting_requests(self, recv_reqs, waiting_cls): def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
new_recv_reqs = [] new_recv_reqs = []
for recv_req in recv_reqs: for recv_req in recv_reqs:
if ( if (
@@ -931,6 +1498,7 @@ class MMReceiverBase(ABC):
model_type=self.model_type, model_type=self.model_type,
host_name=self.hostname, host_name=self.hostname,
receive_count=self.tp_size, receive_count=self.tp_size,
**extra_kwargs,
) )
waiting_req.send_encode_request() waiting_req.send_encode_request()
self.waiting_list.append(waiting_req) self.waiting_list.append(waiting_req)
@@ -946,6 +1514,8 @@ class MMReceiverBase(ABC):
waiting_req._try_recv_mm_data() waiting_req._try_recv_mm_data()
if current_time - waiting_req.start_time > self.wait_timeout: if current_time - waiting_req.start_time > self.wait_timeout:
waiting_req.status = WaitingImageRequestStatus.TIMEOUT waiting_req.status = WaitingImageRequestStatus.TIMEOUT
waiting_req._cleanup_gpu_buffer()
waiting_req.recv_socket.close()
local_status.append(waiting_req.status) local_status.append(waiting_req.status)
local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32) local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32)
@@ -1045,7 +1615,12 @@ class MMReceiverBase(ABC):
return req return req
async def allocate_embedding_buffer(self, req_id, total_bytes): async def allocate_embedding_buffer(self, req_id, total_bytes):
embeddings = torch.empty(total_bytes, dtype=torch.uint8) logger.info(
f"Pre-allocating GPU buffer for mooncake RDMA: "
f"req_id={req_id}, size={total_bytes} bytes"
)
gpu_id = getattr(self.scheduler, "gpu_id", 0)
embeddings = torch.empty(total_bytes, dtype=torch.uint8, device=gpu_id)
self.embeddings_engine.register( self.embeddings_engine.register(
embeddings.data_ptr(), embeddings.data_ptr(),
embeddings.nbytes, embeddings.nbytes,
@@ -1160,10 +1735,48 @@ class MMReceiverHTTP(MMReceiverBase):
scheduler=scheduler, scheduler=scheduler,
) )
# For zmq_to_scheduler # For zmq_to_scheduler and mooncake
def process_waiting_requests(self, recv_reqs): def process_waiting_requests(self, recv_reqs):
if self.encoder_transfer_backend == "mooncake":
gpu_id = getattr(self.scheduler, "gpu_id", 0)
return self._process_waiting_requests(
recv_reqs,
WaitingImageRDMARequest,
embeddings_engine=self.embeddings_engine,
dtype=self.dtype,
gpu_id=gpu_id,
embedding_pool=self.embedding_pool,
)
return self._process_waiting_requests(recv_reqs, WaitingImageRequest) return self._process_waiting_requests(recv_reqs, WaitingImageRequest)
async def _check_encoder_responses(self, responses, encode_requests, req_id):
"""Validate gathered HTTP responses. Returns True if all OK."""
for i, response in enumerate(responses):
if isinstance(response, asyncio.TimeoutError):
timeout_val = envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()
logger.error(
f"Encoder HTTP request timeout ({timeout_val}s) for req_id={req_id} "
f"(request {i}), "
f"encoder={self.encode_urls[encode_requests[i]['encoder_idx']]}"
)
return False
elif isinstance(response, Exception):
logger.error(
f"Encoder HTTP request failed for req_id={req_id} (request {i}): {response}",
exc_info=response,
)
return False
for response in responses:
if response.status != 200:
try:
err_data = await response.json()
msg = err_data.get("message", "Unknown encoder error")
except Exception:
msg = await response.text()
logger.error(f"Encoder returned error {response.status}: {msg}")
return False
return True
async def encode( async def encode(
self, self,
req_id, req_id,
@@ -1228,9 +1841,7 @@ class MMReceiverHTTP(MMReceiverBase):
part_idx_offset += num_parts part_idx_offset += num_parts
async with aiohttp.ClientSession( async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout( timeout=aiohttp.ClientTimeout(total=envs.SGLANG_ENCODER_HTTP_TIMEOUT.get())
total=1800
) # Add timeout for request reliability
) as session: ) as session:
# Send encode requests # Send encode requests
@@ -1242,16 +1853,11 @@ class MMReceiverHTTP(MMReceiverBase):
for encode_request in encode_requests for encode_request in encode_requests
] ]
responses = await asyncio.gather(*tasks) responses = await asyncio.gather(*tasks, return_exceptions=True)
for response in responses:
if response.status != 200:
try:
err_data = await response.json()
msg = err_data.get("message", "Unknown encoder error")
except:
msg = await response.text()
logger.error(f"Encoder returned error {response.status}: {msg}") if not await self._check_encoder_responses(
responses, encode_requests, req_id
):
return return
response_json_list_unsort = [ response_json_list_unsort = [
await response.json() for response in responses await response.json() for response in responses
@@ -1263,15 +1869,10 @@ class MMReceiverHTTP(MMReceiverBase):
# mooncake backend: send bootstrap info # mooncake backend: send bootstrap info
embedding_size_list_sort = [None for _ in range(total_num_parts)] embedding_size_list_sort, response_json_list_sort, total_embedding_bytes = (
response_json_list_sort = [None for _ in range(total_num_parts)] _sort_responses_and_compute_total_bytes(
for response_json in response_json_list_unsort: response_json_list_unsort, total_num_parts
idx = response_json["part_idx"] )
embedding_size_list_sort[idx] = response_json["embedding_size"]
response_json_list_sort[idx] = response_json
total_embedding_bytes = sum(
s for s in embedding_size_list_sort if s is not None
) )
offset = 0 offset = 0
metadata_tasks = [] metadata_tasks = []
@@ -1327,7 +1928,7 @@ class MMReceiverGrpc(MMReceiverBase):
self.send_encode_request(encode_req) self.send_encode_request(encode_req)
return encode_req return encode_req
# For zmq_to_scheduler # For zmq_to_scheduler and mooncake
def process_waiting_requests(self, recv_reqs): def process_waiting_requests(self, recv_reqs):
return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc) return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc)
@@ -1419,14 +2020,9 @@ class MMReceiverGrpc(MMReceiverBase):
if None in response_json_unsorted: if None in response_json_unsorted:
return return
embedding_size_by_part = [None for _ in range(num_parts)] embedding_size_by_part, response_json_sorted, total_embedding_bytes = (
response_json_sorted = [None for _ in range(num_parts)] _sort_responses_and_compute_total_bytes(response_json_unsorted, num_parts)
for response_json in response_json_unsorted: )
idx = response_json["part_idx"]
embedding_size_by_part[idx] = response_json["embedding_size"]
response_json_sorted[idx] = response_json
total_embedding_bytes = sum(s for s in embedding_size_by_part if s is not None)
offset = 0 offset = 0
buffer_address = await self.allocate_embedding_buffer( buffer_address = await self.allocate_embedding_buffer(
req_id, req_id,
+608 -22
View File
@@ -56,6 +56,7 @@ from sglang.srt.utils import (
load_video, load_video,
random_uuid, random_uuid,
) )
from sglang.srt.utils.common import configure_logger
from sglang.srt.utils.network import ( from sglang.srt.utils.network import (
NetworkAddress, NetworkAddress,
config_socket, config_socket,
@@ -179,9 +180,36 @@ def _get_mm_feature(mm_inputs, modality):
) )
def _normalize_aux_value(val):
"""Normalize aux values to pickle types compatible with safe_pickle_loads.
HF multimodal processors (e.g. Qwen3-VL/Omni) emit numpy arrays for
fields like ``video_timestamps`` / ``second_per_grid_ts``. ``numpy.*`` is
not in SafeUnpickler's allowlist, so the receiver would refuse to load
those payloads. Convert numpy values to torch tensors (numeric) or plain
Python lists (object dtype) before pickling.
"""
if val is None:
return None
if isinstance(val, np.ndarray):
if val.dtype == object:
return val.tolist()
return torch.from_numpy(np.ascontiguousarray(val))
if isinstance(val, np.generic):
return val.item()
if isinstance(val, (list, tuple)):
return type(val)(_normalize_aux_value(v) for v in val)
if isinstance(val, dict):
return {k: _normalize_aux_value(v) for k, v in val.items()}
return val
def _build_mm_aux_data(mm_inputs, model_type=None): def _build_mm_aux_data(mm_inputs, model_type=None):
# Video aux metadata, scoped to model_type's video-meta attrs. # Video aux metadata, scoped to model_type's video-meta attrs.
return {attr: mm_inputs.get(attr) for attr in video_meta_attrs_for(model_type)} return {
attr: _normalize_aux_value(mm_inputs.get(attr))
for attr in video_meta_attrs_for(model_type)
}
class MMEncoder: class MMEncoder:
@@ -282,6 +310,14 @@ class MMEncoder:
else: else:
self.mm_global_cache = None self.mm_global_cache = None
# Pre-compute embedding metadata (needed by all ranks for mooncake)
if self.server_args.encoder_transfer_backend == "mooncake":
self._embedding_dims = self._infer_embedding_dims()
self._embedding_dtype = next(self.model.parameters()).dtype
self._element_size = torch.tensor(
[], dtype=self._embedding_dtype
).element_size()
if self.rank == 0: if self.rank == 0:
logger.info( logger.info(
f"Using transfer backend: {self.server_args.encoder_transfer_backend}" f"Using transfer backend: {self.server_args.encoder_transfer_backend}"
@@ -306,6 +342,33 @@ class MMEncoder:
) )
self.embedding_to_send = dict() self.embedding_to_send = dict()
# Need to ensure the NCCL launch order on rank0 matches the dispatch order rank>0
self.encode_dispatch_lock = asyncio.Lock()
# Async mooncake state: track background VIT forward completion
if self.server_args.encoder_transfer_backend == "mooncake":
self._forward_ready_events: Dict[str, asyncio.Event] = {}
self._forward_results: Dict[str, dict] = {}
# when multiple decoder TP ranks call
# POST /encode with the same req_id, only the first triggers
# _run_forward(); subsequent callers wait on the event and
# return the cached metadata.
self._inflight_encode_lock = asyncio.Lock()
self._inflight_encode_events: Dict[str, asyncio.Event] = {}
self._inflight_encode_meta: Dict[str, Tuple] = {}
self._inflight_encode_cleanup_tasks: Dict[str, asyncio.Task] = {}
# Bind unified encode entry point based on backend and cache config
if self.mm_global_cache is not None:
if self.server_args.encoder_transfer_backend == "mooncake":
self._encode_fn = self.encode_with_global_cache_mooncake
else:
self._encode_fn = self.encode_with_global_cache
else:
if self.server_args.encoder_transfer_backend == "mooncake":
self._encode_fn = self.encode_with_mooncake
else:
self._encode_fn = self.encode
logger.info(f"rank {rank} init finish ") logger.info(f"rank {rank} init finish ")
@@ -702,17 +765,20 @@ class MMEncoder:
offset += num_patches offset += num_patches
return hashes return hashes
async def _encode_missing( def _encode_missing(
self, self,
mm_feature, mm_feature,
mm_inputs: dict, mm_inputs: dict,
indices: List[int], indices: List[int],
modality: Modality = Modality.IMAGE, modality: Modality = Modality.IMAGE,
get_feature_fn=None, get_feature_fn=None,
grid_thw: Optional[List] = None,
keep_on_gpu: bool = False,
) -> List[torch.Tensor]: ) -> List[torch.Tensor]:
""" """
GPU Task: Run ViT inference ONLY on the subset of mm items missing from the cache. GPU Task: Run ViT inference ONLY on the subset of mm items missing from the cache.
""" """
if grid_thw is None:
grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type)
# Audio features are per-item (list of mels for mimo_v2, or batched # Audio features are per-item (list of mels for mimo_v2, or batched
@@ -756,7 +822,9 @@ class MMEncoder:
mm_item.set(k, val) mm_item.set(k, val)
with torch.inference_mode(): with torch.inference_mode():
new_embeddings = get_feature_fn([mm_item]).cpu() new_embeddings = get_feature_fn([mm_item])
if not keep_on_gpu:
new_embeddings = new_embeddings.cpu()
if new_embeddings.ndim != 2: if new_embeddings.ndim != 2:
new_embeddings = new_embeddings.reshape(-1, new_embeddings.shape[-1]) new_embeddings = new_embeddings.reshape(-1, new_embeddings.shape[-1])
@@ -819,7 +887,7 @@ class MMEncoder:
# Step 2: All ranks run ViT together on cache-miss images. # Step 2: All ranks run ViT together on cache-miss images.
new_slices = [] new_slices = []
if missing_indices: if missing_indices:
new_slices = await self._encode_missing( new_slices = self._encode_missing(
mm_feature, mm_inputs, missing_indices, modality, get_feature_fn mm_feature, mm_inputs, missing_indices, modality, get_feature_fn
) )
@@ -862,7 +930,7 @@ class MMEncoder:
f"Req {req_id}: Prefetch failed, all ranks running ViT fallback " f"Req {req_id}: Prefetch failed, all ranks running ViT fallback "
f"for {len(hit_indices)} mm items." f"for {len(hit_indices)} mm items."
) )
fallback_slices = await self._encode_missing( fallback_slices = self._encode_missing(
mm_feature, mm_inputs, hit_indices, modality, get_feature_fn mm_feature, mm_inputs, hit_indices, modality, get_feature_fn
) )
else: else:
@@ -919,6 +987,8 @@ class MMEncoder:
mm_embedding, mm_embedding,
**aux_data, **aux_data,
) )
if self.profiler is not None:
self.profiler.step()
return ( return (
mm_embedding.nbytes, mm_embedding.nbytes,
mm_embedding.shape[0], mm_embedding.shape[0],
@@ -927,8 +997,216 @@ class MMEncoder:
None, None,
) )
else: else:
if self.profiler is not None:
self.profiler.step()
return (0, 0, 0, None, None) return (0, 0, 0, None, None)
async def encode_with_global_cache_mooncake(
self,
mm_items,
modality: Modality,
req_id: str,
num_parts: int,
part_idx: int,
hashes: Optional[List[str]] = None,
):
"""Async encode with global cache for mooncake backend.
All ranks participate in VIT forward; tp_size > 1 adds broadcasts for sync."""
try:
mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, modality)
grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type)
mm_feature = _convert(_get_mm_feature(mm_inputs, modality))
num_items = len(grid_thw)
aux_data = _build_mm_aux_data(mm_inputs)
# Setup metadata and event management
nbytes, total_tokens, embedding_dim, event = (
self._setup_mooncake_async_encode(
req_id, num_parts, part_idx, grid_thw, modality, aux_data
)
)
# Rank 0: compute hashes
if self.rank == 0:
if hashes is None:
mm_hashes = self._calculate_hashes_from_features(
mm_feature, grid_thw, modality
)
else:
mm_hashes = hashes
# All ranks: launch background task for cache check + VIT forward.
# Do NOT use run_in_executor: get_feature_fn relies on a session
# context (CUDA / SGLang inference session) that is bound to the
# event-loop main thread and is NOT available inside a
# ThreadPoolExecutor worker thread.
async def _run_forward_with_cache():
try:
# Step 1: Rank 0 checks cache, broadcast mask if TP > 1
if self.rank == 0:
exist_mask = await self.mm_global_cache.batch_is_exist(
mm_hashes
)
mask_tensor = torch.tensor(
[1 if e else 0 for e in exist_mask],
dtype=torch.int32,
)
else:
mask_tensor = torch.zeros(num_items, dtype=torch.int32)
if self.server_args.tp_size > 1:
torch.distributed.broadcast(
mask_tensor,
src=0,
group=self.mm_global_cache.prefetch_tp_group,
)
exist_mask = [m.item() == 1 for m in mask_tensor]
missing_indices = [i for i, e in enumerate(exist_mask) if not e]
hit_indices = [i for i, e in enumerate(exist_mask) if e]
final_slices = [None] * num_items
# Step 2: All ranks run VIT forward for cache misses
# (runs in event loop to preserve session context)
new_slices = []
if missing_indices:
new_slices = self._encode_missing(
mm_feature,
mm_inputs,
missing_indices,
modality,
get_feature_fn,
grid_thw,
keep_on_gpu=True,
)
if self.rank == 0:
for i, idx in enumerate(missing_indices):
final_slices[idx] = new_slices[i]
# Step 3: Rank 0 prefetches cache-hit embeddings
prefetch_status = torch.tensor([1], dtype=torch.int32)
if self.rank == 0 and hit_indices:
hit_hashes = [mm_hashes[i] for i in hit_indices]
hit_tokens = [
self.get_num_tokens(grid_thw[i], modality)
for i in hit_indices
]
self.mm_global_cache.prefetch(
req_id, hit_hashes, hit_tokens, modality
)
try:
async def _wait_prefetch():
while not self.mm_global_cache.check_prefetch_progress(
req_id
):
await asyncio.sleep(0.005)
await asyncio.wait_for(_wait_prefetch(), timeout=60.0)
cached_slices = self.mm_global_cache.get_embeddings(
hit_hashes
)
for i, idx in enumerate(hit_indices):
final_slices[idx] = cached_slices[i]
except (asyncio.TimeoutError, Exception) as e:
logger.error(
f"Prefetch failed for {req_id}: {e}. "
f"Falling back to ViT for "
f"{len(hit_indices)} hit items."
)
prefetch_status[0] = 0
# Broadcast prefetch result if TP > 1
if self.server_args.tp_size > 1:
torch.distributed.broadcast(
prefetch_status,
src=0,
group=self.mm_global_cache.prefetch_tp_group,
)
# Step 4: All ranks fallback VIT for failed prefetch
# (runs in event loop to preserve session context)
fallback_slices = None
if prefetch_status.item() == 0 and hit_indices:
fallback_slices = self._encode_missing(
mm_feature,
mm_inputs,
hit_indices,
modality,
get_feature_fn,
grid_thw,
keep_on_gpu=True,
)
if self.rank == 0:
for i, idx in enumerate(hit_indices):
final_slices[idx] = fallback_slices[i]
# Step 5: Rank 0 assembles and stores result
if self.rank == 0:
mm_embedding = torch.cat(final_slices, dim=0)
# Wait for any pending VIT / cat kernels to finish
# before publishing to /send: mooncake transfer_sync
# is a host-side RDMA read that bypasses CUDA streams
# and would otherwise race with in-flight kernels.
torch.cuda.current_stream(mm_embedding.device).synchronize()
# Background insert new embeddings into cache
all_new_hashes = [mm_hashes[i] for i in missing_indices]
all_new_slices = list(new_slices)
if fallback_slices is not None:
all_new_hashes += [mm_hashes[i] for i in hit_indices]
all_new_slices += list(fallback_slices)
if all_new_hashes:
async def _background_insert():
await asyncio.to_thread(
self.mm_global_cache.insert_batch,
all_new_hashes,
all_new_slices,
)
insert_task = asyncio.create_task(_background_insert())
self.background_tasks.add(insert_task)
insert_task.add_done_callback(self.background_tasks.discard)
self._forward_results[req_id]["embedding"] = mm_embedding
logger.info(
f"Global cache + VIT forward completed for "
f"{req_id}, shape={mm_embedding.shape}"
)
except Exception as e:
logger.error(
f"Global cache + VIT forward failed for " f"{req_id}: {e}"
)
if self.rank == 0:
self._forward_results[req_id]["error"] = str(e)
finally:
if self.rank == 0:
event.set()
if self.profiler is not None:
self.profiler.step()
self._launch_mooncake_background_task(_run_forward_with_cache())
if self.rank == 0:
logger.info(
f"Returning metadata immediately for {req_id}, "
f"global cache + VIT forward running async"
)
return (nbytes, total_tokens, embedding_dim, None, None)
except Exception as e:
error_code = getattr(e, "code", HTTPStatus.INTERNAL_SERVER_ERROR)
error_msg = str(e)
logger.error(
f"Rank {self.rank} encode_with_global_cache_mooncake "
f"failed: {error_msg} {error_code = }"
)
return self._handle_mooncake_encode_error(
req_id, num_parts, part_idx, modality, error_msg, error_code
)
async def _flatten_and_load_audios(self, mm_items): async def _flatten_and_load_audios(self, mm_items):
""" """
Flatten mm_items, load audios concurrently as np.ndarray at Flatten mm_items, load audios concurrently as np.ndarray at
@@ -1246,11 +1524,61 @@ class MMEncoder:
url=None, url=None,
): ):
if self.server_args.encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self.engine.register(embedding.data_ptr(), embedding.nbytes) # Wait for async VIT forward completion if needed
self.engine.transfer_sync( req_id = mm_data.req_id
session_id, embedding.data_ptr(), buffer_address, embedding.nbytes if req_id in self._forward_ready_events:
await self._forward_ready_events[req_id].wait()
result = self._forward_results.get(req_id)
if result is not None:
if "error" in result:
raise InternalError(f"VIT forward failed: {result['error']}")
embedding = result["embedding"]
# Cache the embedding on mm_data so subsequent /send calls
# from other decoder TP ranks can reuse it.
mm_data.cached_embedding = embedding
# Retrieve cached embedding for duplicate /send calls from other
# decoder TP ranks.
if embedding is None:
embedding = mm_data.cached_embedding
if embedding is None:
raise InternalError(
f"No embedding available for Mooncake GPU-direct transfer: {req_id}"
) )
expected_nbytes = mm_data.shape[0] * mm_data.shape[1] * self._element_size
assert embedding.nbytes == expected_nbytes, (
f"Embedding size mismatch for {req_id}: "
f"actual={embedding.nbytes}, expected={expected_nbytes} "
f"(shape={mm_data.shape}, element_size={self._element_size})"
)
# MR was registered once in _run_forward and is shared across all
# sibling-TP /send calls;
mr_already_registered = (
self._forward_results.get(req_id, {}).get("mr_ptr")
== embedding.data_ptr()
)
if not mr_already_registered:
self.engine.register(embedding.data_ptr(), embedding.nbytes)
_t_xfer_start = time.monotonic()
await asyncio.to_thread(
self.engine.transfer_sync,
session_id,
embedding.data_ptr(),
buffer_address,
embedding.nbytes,
)
xfer_ms = (time.monotonic() - _t_xfer_start) * 1000.0
if not mr_already_registered:
self.engine.deregister(embedding.data_ptr()) self.engine.deregister(embedding.data_ptr())
# Only emit at INFO when transfer is slow or fell back
# to per-/send register;
if xfer_ms > 200.0 or not mr_already_registered:
logger.info(
f"[{req_id}] mooncake transfer_sync={xfer_ms:.1f}ms "
f"nbytes={embedding.nbytes} shared_mr={mr_already_registered}"
)
mm_data.embedding = None mm_data.embedding = None
@@ -1263,7 +1591,9 @@ class MMEncoder:
# Serialize data # Serialize data
if self.server_args.encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
serialized_data = pickle.dumps(mm_data) # Mooncake already pushed the embedding via RDMA;
new_mm_data = mm_data.copy_without_embedding()
serialized_data = pickle.dumps(new_mm_data)
buffer = None buffer = None
else: else:
new_mm_data = mm_data.copy_without_embedding() new_mm_data = mm_data.copy_without_embedding()
@@ -1290,7 +1620,9 @@ class MMEncoder:
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
async def encode(self, mm_items, modality: Modality, req_id, num_parts, part_idx): async def encode(
self, mm_items, modality: Modality, req_id, num_parts, part_idx, hashes=None
):
try: try:
grid_dim, mm_embedding, aux_data = await self._encode(mm_items, modality) grid_dim, mm_embedding, aux_data = await self._encode(mm_items, modality)
@@ -1330,10 +1662,204 @@ class MMEncoder:
logger.debug(f"Created error EmbeddingData: {mm_data}") logger.debug(f"Created error EmbeddingData: {mm_data}")
return 0, 0, 0, error_msg, error_code return 0, 0, 0, error_msg, error_code
def _setup_mooncake_async_encode(
self,
req_id: str,
num_parts: int,
part_idx: int,
grid_thw,
modality: Modality,
aux_data: dict,
):
"""Setup metadata and event management for mooncake async encode.
Returns (nbytes, total_tokens, embedding_dim, event)."""
total_tokens = sum(self.get_num_tokens(g, modality) for g in grid_thw)
embedding_dim = self._embedding_dims[modality]
nbytes = total_tokens * embedding_dim * self._element_size
event = None
if self.rank == 0:
mm_data = EmbeddingData(
req_id,
num_parts,
part_idx,
grid_thw,
modality,
embedding=None,
embedding_shape=[total_tokens, embedding_dim],
**aux_data,
)
self.embedding_to_send[req_id] = mm_data
event = asyncio.Event()
self._forward_ready_events[req_id] = event
self._forward_results[req_id] = {}
return nbytes, total_tokens, embedding_dim, event
def _handle_mooncake_encode_error(
self, req_id, num_parts, part_idx, modality, error_msg, error_code
):
"""Handle outer exception for mooncake async encode methods."""
if self.rank == 0:
if req_id in self._forward_ready_events:
self._forward_results[req_id] = {"error": error_msg}
self._forward_ready_events[req_id].set()
mm_data = EmbeddingData(
req_id,
num_parts,
part_idx,
None,
modality,
error_msg=error_msg,
error_code=error_code,
)
self.embedding_to_send[req_id] = mm_data
return 0, 0, 0, error_msg, error_code
def _launch_mooncake_background_task(self, coro):
"""Launch an async background task and track it."""
task = asyncio.create_task(coro)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
return task
async def _cleanup_inflight_encode_state(self, req_id: str):
if not hasattr(self, "_inflight_encode_events"):
return
async with self._inflight_encode_lock:
self._inflight_encode_events.pop(req_id, None)
self._inflight_encode_meta.pop(req_id, None)
task = self._inflight_encode_cleanup_tasks.pop(req_id, None)
if task is not None and not task.done():
task.cancel()
# Also clean up embedding data and forward state
mm_data = self.embedding_to_send.pop(req_id, None)
if mm_data is not None:
mm_data.cached_embedding = None
# Release the rkey after all /send calls have completed.
forward_state = self._forward_results.pop(req_id, None)
if forward_state is not None:
mr_ptr = forward_state.get("mr_ptr")
if mr_ptr is not None:
try:
self.engine.deregister(mr_ptr)
except Exception as dereg_err:
logger.warning(
f"Shared-MR deregister failed for {req_id}: {dereg_err}"
)
self._forward_ready_events.pop(req_id, None)
def _schedule_inflight_encode_cleanup(self, req_id: str):
if not hasattr(self, "_inflight_encode_events"):
return
async def _cleanup_later():
await asyncio.sleep(self.send_timeout)
await self._cleanup_inflight_encode_state(req_id)
old_task = self._inflight_encode_cleanup_tasks.pop(req_id, None)
if old_task is not None and not old_task.done():
old_task.cancel()
task = asyncio.create_task(_cleanup_later())
self._inflight_encode_cleanup_tasks[req_id] = task
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
async def encode_with_mooncake(
self, mm_items, modality: Modality, req_id, num_parts, part_idx, hashes=None
):
"""Async encode for mooncake: all ranks participate in VIT forward via background task,
rank 0 returns metadata immediately."""
try:
mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, modality)
grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type)
aux_data = _build_mm_aux_data(mm_inputs)
# Setup metadata and event management
nbytes, total_tokens, embedding_dim, event = (
self._setup_mooncake_async_encode(
req_id, num_parts, part_idx, grid_thw, modality, aux_data
)
)
# Build mm_item (all ranks)
mm_item = MultimodalDataItem.from_dict(
{
"modality": modality,
"feature": _convert(_get_mm_feature(mm_inputs, modality)),
}
)
for k, v in mm_inputs.items():
if k in _mm_feature_attrs.get(modality, []):
continue
val = _convert(v)
mm_item.set(k, val)
async def _run_forward():
try:
with torch.inference_mode():
emb = get_feature_fn([mm_item])
if len(emb.shape) != 2:
emb = emb.reshape(-1, emb.shape[-1])
# mooncake's transfer_sync is a host-side
# RDMA read that bypasses the CUDA stream. Without an
# explicit sync here, sibling-TP /send handlers can
# invoke transfer_sync while VIT kernels are still
# writing `emb`, producing partial / garbage data on
# the receiver side
if emb.is_cuda:
torch.cuda.current_stream(emb.device).synchronize()
if self.rank == 0:
# Register the MR exactly once here so all sibling-TP /send coroutines share a single registration.
try:
self.engine.register(emb.data_ptr(), emb.nbytes)
self._forward_results[req_id]["mr_ptr"] = emb.data_ptr()
except Exception as reg_err:
logger.warning(
f"Shared-MR register failed for {req_id}, "
f"falling back to per-/send register: {reg_err}"
)
self._forward_results[req_id]["mr_ptr"] = None
self._forward_results[req_id]["embedding"] = emb
except Exception as e:
logger.error(f"VIT forward failed for {req_id}: {e}")
if self.rank == 0:
self._forward_results[req_id]["error"] = str(e)
finally:
if self.rank == 0:
event.set()
if self.profiler is not None:
self.profiler.step()
self._launch_mooncake_background_task(_run_forward())
if self.rank == 0:
logger.info(
f"Returning metadata immediately for {req_id}, "
f"VIT forward running async"
)
return (nbytes, total_tokens, embedding_dim, None, None)
except Exception as e:
error_code = getattr(e, "code", HTTPStatus.INTERNAL_SERVER_ERROR)
error_msg = str(e)
logger.error(
f"Rank {self.rank} encode_with_mooncake failed: "
f"{error_msg} {error_code = }",
exc_info=True,
)
return self._handle_mooncake_encode_error(
req_id, num_parts, part_idx, modality, error_msg, error_code
)
async def encode_request(self, req: dict, modality: Modality): async def encode_request(self, req: dict, modality: Modality):
"""Single-request encode dispatcher: picks cache vs no-cache path.""" """Single-request encode dispatcher.
if self.mm_global_cache is not None:
return await self.encode_with_global_cache( Delegates to ``self._encode_fn``, which is bound at ``__init__``
time to the correct variant (cache / no-cache / mooncake).
"""
return await self._encode_fn(
mm_items=req["mm_items"], mm_items=req["mm_items"],
modality=modality, modality=modality,
req_id=req["req_id"], req_id=req["req_id"],
@@ -1341,13 +1867,6 @@ class MMEncoder:
part_idx=req["part_idx"], part_idx=req["part_idx"],
hashes=req.get("hashes"), hashes=req.get("hashes"),
) )
return await self.encode(
mm_items=req["mm_items"],
modality=modality,
req_id=req["req_id"],
num_parts=req["num_parts"],
part_idx=req["part_idx"],
)
async def batch_encode( async def batch_encode(
self, requests: List[dict], modality: Modality self, requests: List[dict], modality: Modality
@@ -1391,7 +1910,7 @@ class MMEncoder:
), ),
) )
final_slices = await self._encode_missing( final_slices = self._encode_missing(
mm_feature, mm_feature,
mm_inputs, mm_inputs,
list(range(total)), list(range(total)),
@@ -1925,12 +2444,52 @@ async def handle_encode_request(request: dict):
req_id = request["req_id"] req_id = request["req_id"]
start_time = time.monotonic() start_time = time.monotonic()
try: try:
# when multiple decoder TP ranks POST /encode
# with the same req_id, only the first triggers the VIT forward;
# subsequent callers wait and return the same metadata.
if encoder.server_args.encoder_transfer_backend == "mooncake":
async with encoder._inflight_encode_lock:
if req_id in encoder._inflight_encode_events:
event = encoder._inflight_encode_events[req_id]
is_duplicate = True
else:
event = asyncio.Event()
encoder._inflight_encode_events[req_id] = event
is_duplicate = False
if is_duplicate:
await event.wait()
meta = encoder._inflight_encode_meta.get(req_id)
if meta is None:
return ORJSONResponse(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
content={
"status": "error",
"message": "Encode failed on the first request",
"req_id": req_id,
},
)
nbytes, embedding_len, embedding_dim = meta
# Build the same metadata response as the first request
resp = dict(request)
del resp["mm_items"]
resp.update(
{
"embedding_size": nbytes,
"embedding_len": embedding_len,
"embedding_dim": embedding_dim,
}
)
return ORJSONResponse(content=resp)
def start_background_send(req_id): def start_background_send(req_id):
task = asyncio.create_task(encoder.send_with_url(req_id=req_id)) task = asyncio.create_task(encoder.send_with_url(req_id=req_id))
encoder.background_tasks.add(task) encoder.background_tasks.add(task)
task.add_done_callback(encoder.background_tasks.discard) task.add_done_callback(encoder.background_tasks.discard)
# broadcast request, lock together with rank0 await so NCCL
# launch order matches the ZMQ dispatch order rank>0 sees.
async with encoder.encode_dispatch_lock:
request.update({"enter_time": time.time()}) request.update({"enter_time": time.time()})
modality = Modality.from_str(request["modality"]) modality = Modality.from_str(request["modality"])
if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES: if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
@@ -1965,11 +2524,28 @@ async def handle_encode_request(request: dict):
prefill_host=request["prefill_host"], prefill_host=request["prefill_host"],
embedding_port=port, embedding_port=port,
) )
# Signal waiters on failure for mooncake
if encoder.server_args.encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
evt.set()
await encoder._cleanup_inflight_encode_state(req_id)
return ORJSONResponse( return ORJSONResponse(
status_code=error_code, status_code=error_code,
content={"status": "error", "message": error_msg, "req_id": req_id}, content={"status": "error", "message": error_msg, "req_id": req_id},
) )
if encoder.server_args.encoder_transfer_backend == "mooncake": if encoder.server_args.encoder_transfer_backend == "mooncake":
# Store metadata for duplicate callers and signal them
encoder._inflight_encode_meta[req_id] = (
nbytes,
embedding_len,
embedding_dim,
)
evt = encoder._inflight_encode_events.get(req_id)
if evt:
evt.set()
encoder._schedule_inflight_encode_cleanup(req_id)
del request["mm_items"] del request["mm_items"]
request.update( request.update(
{ {
@@ -2016,6 +2592,13 @@ async def handle_encode_request(request: dict):
error_msg = str(e) error_msg = str(e)
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}") logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
rid_to_err_msg[req_id] = error_msg rid_to_err_msg[req_id] = error_msg
# Ensure inflight waiters are unblocked on unexpected errors
if encoder.server_args.encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
evt.set()
await encoder._cleanup_inflight_encode_state(req_id)
return ORJSONResponse( return ORJSONResponse(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR, status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
content={ content={
@@ -2036,7 +2619,10 @@ async def handle_send_request(request: dict):
session_id=request["session_id"], session_id=request["session_id"],
buffer_address=request["buffer_address"], buffer_address=request["buffer_address"],
) )
encoder.embedding_to_send.pop(request["req_id"], None) req_id = request["req_id"]
# Don't pop embedding_to_send here — other decoder TP ranks may still
# need it for their /send calls. Cleanup is handled by the scheduled
# timeout task or _cleanup_inflight_encode_state.
return ORJSONResponse(content=None) return ORJSONResponse(content=None)
+5 -1
View File
@@ -717,10 +717,14 @@ class Envs:
# EPD # EPD
SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0) SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0)
SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0) SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0)
SGLANG_ENCODER_HTTP_TIMEOUT = EnvFloat(1800.0)
SGLANG_ENCODER_REQ_TIMEOUT = EnvFloat(180.0)
SGLANG_ENCODER_DISPATCH_MIN_ITEMS = EnvInt(2) SGLANG_ENCODER_DISPATCH_MIN_ITEMS = EnvInt(2)
SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU = EnvBool(False) SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU = EnvBool(False)
SGLANG_ENCODER_MAX_BATCH_SIZE = EnvInt(8) SGLANG_ENCODER_MAX_BATCH_SIZE = EnvInt(8)
SGLANG_ENCODER_REQ_TIMEOUT = EnvFloat(180.0) # Persistent receiver-side GPU embedding pool size for mooncake EPD transport.
# 0 disables (per-request register/deregister). 4096 = 4GB default per TP
SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096)
# Elastic EP Backup Port # Elastic EP Backup Port
SGLANG_BACKUP_PORT_BASE = EnvInt(10000) SGLANG_BACKUP_PORT_BASE = EnvInt(10000)
+2
View File
@@ -257,6 +257,7 @@ class GenerateReqInput(BaseReq):
# For EPD-disaggregated inference # For EPD-disaggregated inference
need_wait_for_mm_inputs: Optional[bool] = None need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Multimodal tiling controls (extensions) # Multimodal tiling controls (extensions)
max_dynamic_patch: Optional[int] = None max_dynamic_patch: Optional[int] = None
@@ -800,6 +801,7 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Pre-computed delimiter indices for multi-item scoring # Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None multi_item_delimiter_indices: Optional[List[int]] = None
+15 -1
View File
@@ -447,6 +447,16 @@ def _get_precomputed_embedding(
raise NotImplementedError( raise NotImplementedError(
"MM inputs where only some items are precomputed." "MM inputs where only some items are precomputed."
) )
# Normalize device across chunks before concat.
target_device = next(
(t.device for t in precomputed_embeddings if t.is_cuda),
precomputed_embeddings[0].device,
)
precomputed_embeddings = [
t if t.device == target_device else t.to(target_device, non_blocking=True)
for t in precomputed_embeddings
]
result = torch.concat(precomputed_embeddings) result = torch.concat(precomputed_embeddings)
# some models embedding is 3-dim, reshape it to 2-dim (similar to get_embedding_chunk) # some models embedding is 3-dim, reshape it to 2-dim (similar to get_embedding_chunk)
result = result.reshape(-1, result.shape[-1]) result = result.reshape(-1, result.shape[-1])
@@ -1077,7 +1087,11 @@ def offload_mm_features_to_cpu(mm_inputs_list: List[MultimodalInputs]):
item.feature = item.feature.to("cpu", non_blocking=True) item.feature = item.feature.to("cpu", non_blocking=True)
if language_only: if language_only:
pe = item.precomputed_embeddings pe = item.precomputed_embeddings
if isinstance(pe, torch.Tensor) and pe.is_cuda: if (
isinstance(pe, torch.Tensor)
and pe.is_cuda
and not getattr(item, "_keep_device_embedding", False)
):
item.precomputed_embeddings = pe.to("cpu", non_blocking=True) item.precomputed_embeddings = pe.to("cpu", non_blocking=True)
+3 -1
View File
@@ -1108,10 +1108,12 @@ class Scheduler(
# Init mm receiver for EPD disaggregation mode # Init mm receiver for EPD disaggregation mode
if ( if (
self.server_args.language_only self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
): ):
self.mm_receiver = create_mm_receiver( self.mm_receiver = create_mm_receiver(
self.server_args, self.server_args,
dtype=self.model_config.dtype,
hf_config=self.model_config.hf_config, hf_config=self.model_config.hf_config,
pp_rank=self.ps.pp_rank, pp_rank=self.ps.pp_rank,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
@@ -189,7 +189,8 @@ class SchedulerRequestReceiver:
if ( if (
self.ps.pp_rank == 0 self.ps.pp_rank == 0
and self.server_args.language_only and self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
): ):
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs) recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
for req, error_msg, error_code in abort_reqs: for req, error_msg, error_code in abort_reqs:
@@ -792,8 +792,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if ( if (
not self.server_args.language_only not self.server_args.language_only
or self.server_args.encoder_transfer_backend or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer"
in ["zmq_to_tokenizer", "mooncake"]
): ):
if self.server_args.language_only: if self.server_args.language_only:
mm_inputs = await self.mm_receiver.recv_mm_data( mm_inputs = await self.mm_receiver.recv_mm_data(
@@ -817,10 +816,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
elif ( elif (
self.server_args.language_only self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs and not obj.need_wait_for_mm_inputs
): ):
# In language_only mode with zmq_to_scheduler, if we didn't dispatch # In language_only mode with zmq_to_scheduler/mooncake, if we didn't dispatch
# to encoder (e.g., only one image), process locally like non-language_only mode # to encoder (e.g., only one image), process locally like non-language_only mode
mm_inputs = await self.mm_processor.process_mm_data_async( mm_inputs = await self.mm_processor.process_mm_data_async(
image_data=obj.image_data, image_data=obj.image_data,
@@ -1067,6 +1067,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned, num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices, multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
mm_data_mooncake=obj.mm_data_mooncake,
) )
elif isinstance(obj, EmbeddingReqInput): elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available # Resolve unresolved embed overrides now that input_ids are available
@@ -2712,7 +2713,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# This flag will be used in _tokenize_one_request to determine processing path # This flag will be used in _tokenize_one_request to determine processing path
if should_dispatch: if should_dispatch:
obj.need_wait_for_mm_inputs = True obj.need_wait_for_mm_inputs = True
if self.server_args.encoder_transfer_backend == "zmq_to_scheduler": if self.server_args.encoder_transfer_backend in [
"zmq_to_scheduler",
"mooncake",
]:
self.mm_receiver.send_encode_request(obj) self.mm_receiver.send_encode_request(obj)
else: else:
obj.need_wait_for_mm_inputs = False obj.need_wait_for_mm_inputs = False
+12
View File
@@ -2354,6 +2354,8 @@ class SafeUnpickler(pickle.Unpickler):
"sglang.srt.model_executor.model_runner.", "sglang.srt.model_executor.model_runner.",
"sglang.srt.layers.", "sglang.srt.layers.",
"sglang.srt.utils.", "sglang.srt.utils.",
"sglang.srt.disaggregation.",
"sglang.srt.managers.",
"torch_npu.", "torch_npu.",
} }
@@ -2394,6 +2396,16 @@ def safe_pickle_load(fp):
return SafeUnpickler(fp).load() return SafeUnpickler(fp).load()
def safe_pickle_loads(data):
"""Drop-in replacement for pickle.loads() that blocks unsafe class loading."""
if isinstance(data, (bytes, bytearray, memoryview)):
buf = bytes(data)
else:
# zmq.Frame and other buffer-protocol objects
buf = bytes(memoryview(data))
return SafeUnpickler(io.BytesIO(buf)).load()
def debug_timing(func): def debug_timing(func):
# todo: replace with a more organized instrumentation # todo: replace with a more organized instrumentation
def wrapper(*args, **kwargs): def wrapper(*args, **kwargs):
@@ -206,6 +206,7 @@ class RequestLogger:
"image_data", "image_data",
"audio_data", "audio_data",
"video_data", "video_data",
"mm_data_mooncake",
"lora_path", "lora_path",
"sampling_params", "sampling_params",
} }
@@ -219,6 +220,7 @@ class RequestLogger:
"image_data", "image_data",
"audio_data", "audio_data",
"video_data", "video_data",
"mm_data_mooncake",
"lora_path", "lora_path",
} }
out_skip_names = {"text", "output_ids", "embedding"} out_skip_names = {"text", "output_ids", "embedding"}
@@ -1377,5 +1377,139 @@ class TestEPDDisaggregationGrpcEncoderOnly(PDDisaggregationServerBase):
channel.close() channel.close()
@unittest.skipIf(
is_in_ci(),
"TestEPDDisaggregationMooncake test requires RDMA hardware, skipping in CI",
)
class TestEPDDisaggregationMooncake(MMMUMixin, PDDisaggregationServerBase):
"""Test EPD disaggregation with mooncake GPU→GPU transfer.
Validates the async VIT forward + GPU buffer pre-allocation +
GPU-to-GPU mooncake transfer pipeline using MMMU eval (multi-image).
"""
# Qwen2.5-VL-3B-Instruct scores ~0.40 on the 50-sample MMMU subset.
accuracy = 0.40
mmmu_args = ["--limit", "50"]
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = DEFAULT_SMALL_VLM_MODEL_NAME_FOR_TEST
cls.base_url = cls.lb_url # MMMUMixin reads this for OPENAI_API_BASE
cls.encode_port = f"{int(cls.lb_port) + 306}"
cls.encode_url = f"http://{cls.base_host}:{cls.encode_port}"
print(
f"Setting up EPD Mooncake RDMA: encode={cls.encode_port}, "
f"prefill={cls.prefill_port}, decode={cls.decode_port}"
)
# Start servers in order: encode -> prefill/decode
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()
# Wait for all servers to be ready
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):
"""Start encode server with mooncake transfer backend"""
encode_args = [
"--trust-remote-code",
"--encoder-only",
"--encoder-transfer-backend",
"mooncake",
"--tp",
"1",
"--port",
cls.encode_port,
"--enable-prefix-mm-cache",
]
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):
"""Start prefill server with mooncake transfer backend"""
prefill_args = [
"--trust-remote-code",
"--language-only",
"--encoder-urls",
cls.encode_url,
"--encoder-transfer-backend",
"mooncake",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"1",
"--base-gpu-id",
"1",
"--port",
cls.prefill_port,
]
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):
"""Start decode server"""
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
"1",
"--base-gpu-id",
"2",
"--port",
cls.decode_port,
]
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):
"""Clean up all processes"""
for process in [
cls.process_lb,
cls.process_decode,
cls.process_prefill,
cls.process_encode,
]:
if process:
try:
kill_process_tree(process.pid)
except Exception as e:
print(f"Error killing process: {e}")
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()