From 36d0a6e08eb692ec50f70db74fe8e944609bda4c Mon Sep 17 00:00:00 2001 From: LucQueen <32257212+LucQueen@users.noreply.github.com> Date: Fri, 29 May 2026 10:42:24 +0800 Subject: [PATCH] [EPD] Optimize the Mooncake backend (#22587) Co-authored-by: ZhengWG --- .../srt/disaggregation/encode_receiver.py | 806 +++++++++++++++--- .../srt/disaggregation/encode_server.py | 676 ++++++++++++++- python/sglang/srt/environ.py | 6 +- python/sglang/srt/managers/io_struct.py | 2 + python/sglang/srt/managers/mm_utils.py | 16 +- python/sglang/srt/managers/scheduler.py | 4 +- .../scheduler_components/request_receiver.py | 3 +- .../sglang/srt/managers/tokenizer_manager.py | 14 +- python/sglang/srt/utils/common.py | 12 + python/sglang/srt/utils/request_logger.py | 2 + .../disaggregation/test_epd_disaggregation.py | 134 +++ 11 files changed, 1516 insertions(+), 159 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index 72f3982cf..fc554f0b7 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -1,17 +1,17 @@ import asyncio import itertools import logging -import pickle import random import threading import time import uuid +import weakref from abc import ABC, abstractmethod from array import array from collections import OrderedDict, defaultdict from enum import IntEnum from http import HTTPStatus -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING, Dict, List, Optional, Tuple import aiohttp 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.server_args import ServerArgs 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.network import ( NetworkAddress, @@ -153,6 +154,7 @@ class EmbeddingData: self.shape = embedding_shape else: self.shape = list(embedding.shape) if embedding is not None else None + self.cached_embedding = None self.error_msg = error_msg self.error_code = error_code # Store additional metadata (e.g., video_timestamps for qwen3_vl) @@ -182,7 +184,8 @@ class EmbeddingData: error_code=self.error_code, ) 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 setattr(new_data, key, value) return new_data @@ -468,6 +471,7 @@ class WaitingImageRequest: self.start_time = time.time() def send_encode_request(self): + async def _send_single_request(session, url, payload): try: 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 with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=1800) + timeout=aiohttp.ClientTimeout( + total=envs.SGLANG_ENCODER_HTTP_TIMEOUT.get() + ) ) as session: tasks = [] logger.info(f"{self.num_items_assigned = } ") @@ -525,8 +531,18 @@ class WaitingImageRequest: results = await asyncio.gather(*tasks, return_exceptions=True) for i, result in enumerate(results): - if isinstance(result, Exception): - logger.error(f"Request {i} failed: {result}") + if isinstance(result, asyncio.TimeoutError): + 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: logger.debug(f"Request {i} succeeded.") @@ -548,7 +564,7 @@ class WaitingImageRequest: except zmq.Again: # No data available yet, wait a bit and retry 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: logger.warning( 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.recv_socket.close() + def _cleanup_gpu_buffer(self): + pass + class WaitingImageRequestGrpc(WaitingImageRequest): 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): is_cross_node = server_args.dist_init_addr @@ -669,6 +1204,16 @@ class MMReceiverBase(ABC): self.encode_urls = server_args.encoder_urls self.recv_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() 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 = ( getattr(hf_config, "model_type", "").lower() if hf_config is not None @@ -690,63 +1235,90 @@ class MMReceiverBase(ABC): ), ) self.embeddings_buffer = dict() - elif self.encoder_transfer_backend == "zmq_to_scheduler": - 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() - if hf_config is not None: - transport_mode = _determine_tensor_transport_mode(server_args) - import_processors("sglang.srt.multimodal.processors") - _processor = None + self.embedding_pool = None + pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get() + if pool_mb and pool_mb > 0 and scheduler is not None: + gpu_id = getattr(scheduler, "gpu_id", 0) try: - _processor = get_processor( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - use_fast=not server_args.disable_fast_image_processor, - tokenizer_backend=server_args.tokenizer_backend, + self.embedding_pool = MooncakeEmbeddingPool( + self.embeddings_engine, gpu_id, pool_mb * 1024 * 1024 ) - except ValueError as e: - error_message = str(e) - if "does not have a slow version" in error_message: - logger.info( - f"Processor {server_args.tokenizer_path} does not have a slow version. Automatically use fast version" - ) - _processor = get_processor( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - use_fast=True, - tokenizer_backend=server_args.tokenizer_backend, - ) - else: - raise e - - # Skip mm_pool if not adaptive dispatch to encoder - enable_adaptive_dispatch_to_encoder = ( - server_args.enable_adaptive_dispatch_to_encoder - ) - self.mm_processor = get_mm_processor( - hf_config, + except Exception: + logger.exception( + "Failed to allocate MooncakeEmbeddingPool, " + "falling back to per-request register" + ) + self.embedding_pool = 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, - _processor, - transport_mode, + hf_config, 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, ) + 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) + 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 + try: + _processor = get_processor( + server_args.tokenizer_path, + tokenizer_mode=server_args.tokenizer_mode, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + use_fast=not server_args.disable_fast_image_processor, + **extra_kwargs, + ) + except ValueError as e: + error_message = str(e) + if "does not have a slow version" in error_message: + logger.info( + f"Processor {server_args.tokenizer_path} does not have a slow version. Automatically use fast version" + ) + _processor = get_processor( + server_args.tokenizer_path, + tokenizer_mode=server_args.tokenizer_mode, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + use_fast=True, + **extra_kwargs, + ) + else: + raise e + + 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( + hf_config, + server_args, + _processor, + transport_mode, + skip_mm_pool=not enable_adaptive_dispatch_to_encoder, + **mm_processor_kwargs, + ) + @abstractmethod def process_waiting_requests(self, recv_reqs): pass @@ -814,7 +1386,7 @@ class MMReceiverBase(ABC): parts = await recv_socket.recv_multipart(copy=False) if not parts: 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: logger.warning( f"Encoder error for req_id={req_id}: {recv_obj.error_msg} " @@ -859,22 +1431,7 @@ class MMReceiverBase(ABC): return None raw_buffer = self.embeddings_buffer.pop(req_id) self.embeddings_engine.deregister(raw_buffer.data_ptr()) - byte_offset = 0 - 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 + _slice_embedding_buffer(raw_buffer, recv_embedding_data, self.dtype) recv_embedding = recv_embedding_data.get_embedding(is_concat=True) @@ -902,6 +1459,16 @@ class MMReceiverBase(ABC): mm_data, len(self.encode_urls) ) 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( target=self._run_encode_in_thread, args=( @@ -916,7 +1483,7 @@ class MMReceiverBase(ABC): encode_thread.start() # 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 = [] for recv_req in recv_reqs: if ( @@ -931,6 +1498,7 @@ class MMReceiverBase(ABC): model_type=self.model_type, host_name=self.hostname, receive_count=self.tp_size, + **extra_kwargs, ) waiting_req.send_encode_request() self.waiting_list.append(waiting_req) @@ -946,6 +1514,8 @@ class MMReceiverBase(ABC): waiting_req._try_recv_mm_data() if current_time - waiting_req.start_time > self.wait_timeout: waiting_req.status = WaitingImageRequestStatus.TIMEOUT + waiting_req._cleanup_gpu_buffer() + waiting_req.recv_socket.close() local_status.append(waiting_req.status) local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32) @@ -1045,7 +1615,12 @@ class MMReceiverBase(ABC): return req 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( embeddings.data_ptr(), embeddings.nbytes, @@ -1160,10 +1735,48 @@ class MMReceiverHTTP(MMReceiverBase): scheduler=scheduler, ) - # For zmq_to_scheduler + # For zmq_to_scheduler and mooncake 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) + 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( self, req_id, @@ -1228,9 +1841,7 @@ class MMReceiverHTTP(MMReceiverBase): part_idx_offset += num_parts async with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout( - total=1800 - ) # Add timeout for request reliability + timeout=aiohttp.ClientTimeout(total=envs.SGLANG_ENCODER_HTTP_TIMEOUT.get()) ) as session: # Send encode requests @@ -1242,17 +1853,12 @@ class MMReceiverHTTP(MMReceiverBase): for encode_request in encode_requests ] - responses = await asyncio.gather(*tasks) - 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() + responses = await asyncio.gather(*tasks, return_exceptions=True) - logger.error(f"Encoder returned error {response.status}: {msg}") - return + if not await self._check_encoder_responses( + responses, encode_requests, req_id + ): + return response_json_list_unsort = [ await response.json() for response in responses ] @@ -1263,15 +1869,10 @@ class MMReceiverHTTP(MMReceiverBase): # mooncake backend: send bootstrap info - embedding_size_list_sort = [None for _ in range(total_num_parts)] - response_json_list_sort = [None for _ in range(total_num_parts)] - for response_json in response_json_list_unsort: - 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 + embedding_size_list_sort, response_json_list_sort, total_embedding_bytes = ( + _sort_responses_and_compute_total_bytes( + response_json_list_unsort, total_num_parts + ) ) offset = 0 metadata_tasks = [] @@ -1327,7 +1928,7 @@ class MMReceiverGrpc(MMReceiverBase): self.send_encode_request(encode_req) return encode_req - # For zmq_to_scheduler + # For zmq_to_scheduler and mooncake def process_waiting_requests(self, recv_reqs): return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc) @@ -1419,14 +2020,9 @@ class MMReceiverGrpc(MMReceiverBase): if None in response_json_unsorted: return - embedding_size_by_part = [None for _ in range(num_parts)] - response_json_sorted = [None for _ in range(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) + embedding_size_by_part, response_json_sorted, total_embedding_bytes = ( + _sort_responses_and_compute_total_bytes(response_json_unsorted, num_parts) + ) offset = 0 buffer_address = await self.allocate_embedding_buffer( req_id, diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 3300f0b60..604eb35b7 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -56,6 +56,7 @@ from sglang.srt.utils import ( load_video, random_uuid, ) +from sglang.srt.utils.common import configure_logger from sglang.srt.utils.network import ( NetworkAddress, 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): # 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: @@ -282,6 +310,14 @@ class MMEncoder: else: 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: logger.info( f"Using transfer backend: {self.server_args.encoder_transfer_backend}" @@ -306,6 +342,33 @@ class MMEncoder: ) 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 ") @@ -702,18 +765,21 @@ class MMEncoder: offset += num_patches return hashes - async def _encode_missing( + def _encode_missing( self, mm_feature, mm_inputs: dict, indices: List[int], modality: Modality = Modality.IMAGE, get_feature_fn=None, + grid_thw: Optional[List] = None, + keep_on_gpu: bool = False, ) -> List[torch.Tensor]: """ GPU Task: Run ViT inference ONLY on the subset of mm items missing from the cache. """ - grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) + if grid_thw is None: + 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 # N x n_mels x T_max for qwen2_audio); slice by item index and keep @@ -756,7 +822,9 @@ class MMEncoder: mm_item.set(k, val) 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: 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. new_slices = [] if missing_indices: - new_slices = await self._encode_missing( + new_slices = self._encode_missing( 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"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 ) else: @@ -919,6 +987,8 @@ class MMEncoder: mm_embedding, **aux_data, ) + if self.profiler is not None: + self.profiler.step() return ( mm_embedding.nbytes, mm_embedding.shape[0], @@ -927,8 +997,216 @@ class MMEncoder: None, ) else: + if self.profiler is not None: + self.profiler.step() 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): """ Flatten mm_items, load audios concurrently as np.ndarray at @@ -1246,11 +1524,61 @@ class MMEncoder: url=None, ): if self.server_args.encoder_transfer_backend == "mooncake": - self.engine.register(embedding.data_ptr(), embedding.nbytes) - self.engine.transfer_sync( - session_id, embedding.data_ptr(), buffer_address, embedding.nbytes + # Wait for async VIT forward completion if needed + req_id = mm_data.req_id + 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})" ) - self.engine.deregister(embedding.data_ptr()) + + # 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()) + # 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 @@ -1263,7 +1591,9 @@ class MMEncoder: # Serialize data 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 else: 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) - 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: grid_dim, mm_embedding, aux_data = await self._encode(mm_items, modality) @@ -1330,23 +1662,210 @@ class MMEncoder: logger.debug(f"Created error EmbeddingData: {mm_data}") return 0, 0, 0, error_msg, error_code - async def encode_request(self, req: dict, modality: Modality): - """Single-request encode dispatcher: picks cache vs no-cache path.""" - if self.mm_global_cache is not None: - return await self.encode_with_global_cache( - mm_items=req["mm_items"], - modality=modality, - req_id=req["req_id"], - num_parts=req["num_parts"], - part_idx=req["part_idx"], - hashes=req.get("hashes"), + 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, ) - return await self.encode( + 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): + """Single-request encode dispatcher. + + 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"], modality=modality, req_id=req["req_id"], num_parts=req["num_parts"], part_idx=req["part_idx"], + hashes=req.get("hashes"), ) async def batch_encode( @@ -1391,7 +1910,7 @@ class MMEncoder: ), ) - final_slices = await self._encode_missing( + final_slices = self._encode_missing( mm_feature, mm_inputs, list(range(total)), @@ -1925,34 +2444,74 @@ async def handle_encode_request(request: dict): req_id = request["req_id"] start_time = time.monotonic() 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): task = asyncio.create_task(encoder.send_with_url(req_id=req_id)) encoder.background_tasks.add(task) task.add_done_callback(encoder.background_tasks.discard) - request.update({"enter_time": time.time()}) - modality = Modality.from_str(request["modality"]) - if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES: - try: + # 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()}) + modality = Modality.from_str(request["modality"]) + if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES: + try: + nbytes, embedding_len, embedding_dim, error_msg, error_code = ( + await encoder_scheduler.submit(request) + ) + except asyncio.TimeoutError: + return ORJSONResponse( + status_code=HTTPStatus.GATEWAY_TIMEOUT, + content={ + "status": "error", + "message": "encoder batch timed out", + "req_id": req_id, + }, + ) + else: + for socket in send_sockets: + socket.send_pyobj(request) nbytes, embedding_len, embedding_dim, error_msg, error_code = ( - await encoder_scheduler.submit(request) + await encoder.encode_request(request, modality) ) - except asyncio.TimeoutError: - return ORJSONResponse( - status_code=HTTPStatus.GATEWAY_TIMEOUT, - content={ - "status": "error", - "message": "encoder batch timed out", - "req_id": req_id, - }, - ) - else: - for socket in send_sockets: - socket.send_pyobj(request) - nbytes, embedding_len, embedding_dim, error_msg, error_code = ( - await encoder.encode_request(request, modality) - ) if error_msg: if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler": @@ -1965,11 +2524,28 @@ async def handle_encode_request(request: dict): prefill_host=request["prefill_host"], 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( status_code=error_code, content={"status": "error", "message": error_msg, "req_id": req_id}, ) 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"] request.update( { @@ -2016,6 +2592,13 @@ async def handle_encode_request(request: dict): error_msg = str(e) logger.error(f"Unexpected error in encoder logic for {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( status_code=HTTPStatus.INTERNAL_SERVER_ERROR, content={ @@ -2036,7 +2619,10 @@ async def handle_send_request(request: dict): session_id=request["session_id"], 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) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 9821b2b73..f5a0353bf 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -717,10 +717,14 @@ class Envs: # EPD SGLANG_ENCODER_RECV_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_IMAGE_PROCESSOR_USE_GPU = EnvBool(False) 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 SGLANG_BACKUP_PORT_BASE = EnvInt(10000) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 54bb82898..76fd0ab75 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -257,6 +257,7 @@ class GenerateReqInput(BaseReq): # For EPD-disaggregated inference need_wait_for_mm_inputs: Optional[bool] = None num_items_assigned: Optional[Dict[Modality, List[int]]] = None + mm_data_mooncake: Optional[List] = None # Multimodal tiling controls (extensions) max_dynamic_patch: Optional[int] = None @@ -800,6 +801,7 @@ class TokenizedGenerateReqInput(BaseReq): need_wait_for_mm_inputs: bool = False num_items_assigned: Optional[Dict[Modality, List[int]]] = None + mm_data_mooncake: Optional[List] = None # Pre-computed delimiter indices for multi-item scoring multi_item_delimiter_indices: Optional[List[int]] = None diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 174eece93..e7abc915a 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -447,6 +447,16 @@ def _get_precomputed_embedding( raise NotImplementedError( "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) # some models embedding is 3-dim, reshape it to 2-dim (similar to get_embedding_chunk) 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) if language_only: 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) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 135d254d9..ff6ea128e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1108,10 +1108,12 @@ class Scheduler( # Init mm receiver for EPD disaggregation mode if ( 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.server_args, + dtype=self.model_config.dtype, hf_config=self.model_config.hf_config, pp_rank=self.ps.pp_rank, tp_rank=self.ps.tp_rank, diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index 28ec4f098..0c5b83b67 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -189,7 +189,8 @@ class SchedulerRequestReceiver: if ( self.ps.pp_rank == 0 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) for req, error_msg, error_code in abort_reqs: diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index bc33c29ac..48bf059db 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -792,8 +792,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if ( not self.server_args.language_only - or self.server_args.encoder_transfer_backend - in ["zmq_to_tokenizer", "mooncake"] + or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer" ): if self.server_args.language_only: mm_inputs = await self.mm_receiver.recv_mm_data( @@ -817,10 +816,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ) elif ( 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 ): - # 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 mm_inputs = await self.mm_processor.process_mm_data_async( image_data=obj.image_data, @@ -1067,6 +1067,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, num_items_assigned=obj.num_items_assigned, multi_item_delimiter_indices=obj.multi_item_delimiter_indices, + mm_data_mooncake=obj.mm_data_mooncake, ) elif isinstance(obj, EmbeddingReqInput): # 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 if should_dispatch: 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) else: obj.need_wait_for_mm_inputs = False diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index cb6e908b6..feb505d5d 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -2354,6 +2354,8 @@ class SafeUnpickler(pickle.Unpickler): "sglang.srt.model_executor.model_runner.", "sglang.srt.layers.", "sglang.srt.utils.", + "sglang.srt.disaggregation.", + "sglang.srt.managers.", "torch_npu.", } @@ -2394,6 +2396,16 @@ def safe_pickle_load(fp): 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): # todo: replace with a more organized instrumentation def wrapper(*args, **kwargs): diff --git a/python/sglang/srt/utils/request_logger.py b/python/sglang/srt/utils/request_logger.py index 45373103d..2b3b9bd99 100644 --- a/python/sglang/srt/utils/request_logger.py +++ b/python/sglang/srt/utils/request_logger.py @@ -206,6 +206,7 @@ class RequestLogger: "image_data", "audio_data", "video_data", + "mm_data_mooncake", "lora_path", "sampling_params", } @@ -219,6 +220,7 @@ class RequestLogger: "image_data", "audio_data", "video_data", + "mm_data_mooncake", "lora_path", } out_skip_names = {"text", "output_ids", "embedding"} diff --git a/test/registered/disaggregation/test_epd_disaggregation.py b/test/registered/disaggregation/test_epd_disaggregation.py index 0c370fcda..583cc79eb 100644 --- a/test/registered/disaggregation/test_epd_disaggregation.py +++ b/test/registered/disaggregation/test_epd_disaggregation.py @@ -1377,5 +1377,139 @@ class TestEPDDisaggregationGrpcEncoderOnly(PDDisaggregationServerBase): 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__": unittest.main()