[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 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,
+631 -45
View File
@@ -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)
+5 -1
View File
@@ -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)
+2
View File
@@ -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
+15 -1
View File
@@ -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)
+3 -1
View File
@@ -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,
@@ -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:
@@ -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
+12
View File
@@ -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):
@@ -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"}
@@ -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()