[EPD] Optimize the Mooncake backend (#22587)
Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"}
|
||||
|
||||
Reference in New Issue
Block a user