From 8a123cbd0e1c1ca735615a071d46879b72f6235d Mon Sep 17 00:00:00 2001 From: siyu Date: Fri, 21 Aug 2026 15:22:48 +0800 Subject: [PATCH] [Refactor] New EPD (#30398) Co-authored-by: Yuang Chen <1131578721@qq.com> Co-authored-by: Yuang Chen Co-authored-by: ZhengWG --- .github/CODEOWNERS | 3 +- python/sglang/launch_server.py | 4 +- .../srt/disaggregation/encode_server.py | 4602 ----------------- .../srt/disaggregation/encoder/__init__.py | 0 .../grpc_server.py} | 37 +- .../srt/disaggregation/encoder/http_server.py | 653 +++ .../disaggregation/encoder/preprocessor.py | 827 +++ .../receiver.py} | 1220 ++--- .../srt/disaggregation/encoder/runtime.py | 1643 ++++++ .../srt/disaggregation/encoder/server.py | 2065 ++++++++ python/sglang/srt/managers/io_struct.py | 7 - python/sglang/srt/managers/mm_utils.py | 1 + python/sglang/srt/managers/schedule_batch.py | 2 + python/sglang/srt/managers/scheduler.py | 6 +- .../sglang/srt/managers/tokenizer_manager.py | 5 +- .../srt/multimodal/processors/mimo_v2.py | 2 +- python/sglang/srt/utils/request_logger.py | 2 - .../test_encoder_server_metrics.py | 2 +- .../disaggregation/test_encode_receiver.py | 2 +- .../unit/disaggregation/test_encode_server.py | 473 +- .../disaggregation/test_encoder_health.py | 24 +- .../disaggregation/test_encoder_scheduler.py | 2 +- .../test_kimi_k3_encoder_mode.py | 92 +- .../unit/test_global_config_read_ratchet.py | 2 +- .../unit/test_publish_precedes_bag_reads.py | 4 +- 25 files changed, 6375 insertions(+), 5305 deletions(-) delete mode 100644 python/sglang/srt/disaggregation/encode_server.py create mode 100644 python/sglang/srt/disaggregation/encoder/__init__.py rename python/sglang/srt/disaggregation/{encode_grpc_server.py => encoder/grpc_server.py} (90%) create mode 100644 python/sglang/srt/disaggregation/encoder/http_server.py create mode 100644 python/sglang/srt/disaggregation/encoder/preprocessor.py rename python/sglang/srt/disaggregation/{encode_receiver.py => encoder/receiver.py} (75%) create mode 100644 python/sglang/srt/disaggregation/encoder/runtime.py create mode 100644 python/sglang/srt/disaggregation/encoder/server.py diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 683d3b88e..5ce0a9f86 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -16,8 +16,7 @@ /python/sglang/srt/constrained @hnyls2002 @DarkSharpness @JustinTong0323 /python/sglang/srt/disaggregation @ByronHsu @hnyls2002 @ShangmingCai @HaiShaw @Duyi-Wang @sogalin /python/sglang/srt/disaggregation/ascend @ping1jing2 @iforgetmyname -/python/sglang/srt/disaggregation/encode_receiver.py @ShangmingCai @liusy58 @ZhengWG @gty111 -/python/sglang/srt/disaggregation/encode_server.py @ShangmingCai @liusy58 @ZhengWG @gty111 +/python/sglang/srt/disaggregation/encoder @ShangmingCai @liusy58 @ZhengWG @gty111 /python/sglang/srt/disaggregation/mori @Duyi-Wang @kkHuang-amd @HaiShaw @Lzy17 @billishyahao /python/sglang/srt/distributed @yizhang2077 @merrymercy @ch-wan /python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py @ShangmingCai @stmatengss diff --git a/python/sglang/launch_server.py b/python/sglang/launch_server.py index 5e47dd22b..2a07ceaec 100644 --- a/python/sglang/launch_server.py +++ b/python/sglang/launch_server.py @@ -18,13 +18,13 @@ def run_server(server_args): if server_args.encoder_only: # For encoder disaggregation if server_args.smg_grpc_mode or server_args.grpc_mode: - from sglang.srt.disaggregation.encode_grpc_server import ( + from sglang.srt.disaggregation.encoder.grpc_server import ( serve_grpc_encoder, ) asyncio.run(serve_grpc_encoder(server_args)) else: - from sglang.srt.disaggregation.encode_server import launch_server + from sglang.srt.disaggregation.encoder.http_server import launch_server launch_server(server_args) elif server_args.smg_grpc_mode: diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py deleted file mode 100644 index a1e340d49..000000000 --- a/python/sglang/srt/disaggregation/encode_server.py +++ /dev/null @@ -1,4602 +0,0 @@ -import asyncio -import concurrent.futures -import contextlib -import ctypes -import functools -import logging -import multiprocessing as mp -import os -import pickle -import threading -import time -import traceback -import uuid -from collections import defaultdict -from dataclasses import dataclass -from http import HTTPStatus -from typing import Annotated, Any, Dict, List, Optional, Set, Tuple, Union - -import aiohttp -import numpy as np -import requests as http_requests -import torch -import uvicorn -import zmq -import zmq.asyncio -from fastapi import Body, FastAPI -from fastapi.responses import ORJSONResponse, Response -from transformers import AutoProcessor - -from sglang.srt.configs.device_config import DeviceConfig -from sglang.srt.configs.load_config import LoadConfig -from sglang.srt.configs.model_config import ModelConfig -from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX -from sglang.srt.disaggregation.encode_receiver import ( - EmbeddingData, - video_meta_attrs_for, -) -from sglang.srt.distributed.parallel_state import ( - get_default_distributed_backend, - get_mooncake_transfer_engine, - get_tp_group, - init_distributed_environment, - initialize_model_parallel, -) -from sglang.srt.environ import envs -from sglang.srt.layers.dp_attention import initialize_dp_attention -from sglang.srt.managers.io_struct import ( - ProfileReq, - ProfileReqType, - async_sock_recv, - async_sock_send, - sock_send, - wrap_as_pickle, -) -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem -from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache -from sglang.srt.model_executor.model_runner_components.load_model_utils import ( - maybe_precompile_model_kernels_after_loading, -) -from sglang.srt.model_loader import get_model as load_model -from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media -from sglang.srt.multimodal.encoder_preprocessing import ( - EncoderPreprocessOutput, - get_encoder_preprocessed_items, - invoke_encoder_preprocessor, - resolve_encoder_media_processor_config, -) -from sglang.srt.multimodal.processors.qwen_vl import preprocess_video -from sglang.srt.observability.metrics_collector import EncoderMetricsCollector -from sglang.srt.observability.req_time_stats import EncoderReqTimeStats -from sglang.srt.observability.trace import ( - process_tracing_init, - trace_set_thread_info, -) -from sglang.srt.runtime_context import ( - configured_tp_size, - get_device, - get_disagg, - get_exec, - get_mm, - get_model, - get_observability, - get_parallel, - get_serving, - publish, -) -from sglang.srt.server_args import ( - PortArgs, - ServerArgs, -) -from sglang.srt.utils import ( - CLIENT_MEDIA_EXCEPTIONS, - add_prometheus_middleware, - configure_logger, - configure_media_url_security, - load_audio, - load_image, - load_video, - random_uuid, - set_prometheus_multiproc_dir, -) -from sglang.srt.utils.common import configure_logger, maybe_reindex_device_id -from sglang.srt.utils.hf_transformers_utils import resolve_image_processor_backend -from sglang.srt.utils.network import ( - NetworkAddress, - config_socket, - get_free_port, - get_local_ip_auto, - get_zmq_socket, -) - -logger = logging.getLogger(__name__) - -HEALTH_CHECK_TIMEOUT = 30 - - -def is_health_check_request(rid: Optional[str]) -> bool: - return isinstance(rid, str) and rid.startswith(HEALTH_CHECK_RID_PREFIX) - - -# Minimal 32x32 black PNG for health check dummy encode -MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" - -# Minimal WAV: 16kHz mono 16-bit PCM, 160 samples (0.01s) of silence -MINIMUM_WAV_SILENCE_BASE64 = "UklGRmQBAABXQVZFZm10IBAAAAABAAEAgD4AAAB9AAACABAAZGF0YUABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==" - -rid_lock = asyncio.Lock() -rid_to_receive_endpoint: Dict[str, List[str]] = dict() -rid_to_receive_count: Dict[str, int] = dict() -rid_to_err_msg: Dict[str, str] = dict() -cond_dict_lock = asyncio.Lock() -rid_to_cond: Dict[str, asyncio.Condition] = {} -# mooncake: /send completions per part; release GPU embedding once receive_count reached. -mooncake_send_done_count: Dict[str, int] = dict() - -use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get() - -ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get() -ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set() -# Watchdog: max time to wait for a batched /encode result. Bounds HTTP latency -# if the batch worker stalls (NCCL hang, dead worker proc, etc.). -ENCODER_REQ_TIMEOUT = envs.SGLANG_ENCODER_REQ_TIMEOUT.get() - - -class MMError(Exception): - def __init__(self, message, code=HTTPStatus.INTERNAL_SERVER_ERROR): - self.message = message - self.code = code - super().__init__(self.message) - - -class BadRequestError(MMError): - def __init__(self, message): - super().__init__(message, code=HTTPStatus.BAD_REQUEST) - - -class InternalError(MMError): - def __init__(self, message): - super().__init__(message, code=HTTPStatus.INTERNAL_SERVER_ERROR) - - -@dataclass -class GlobalCacheEncodeContext: - req_id: str - modality: Modality - mm_inputs: dict - get_feature_fn: Any - grid_thw: List - mm_feature: Any - num_items: int - aux_data: dict - str_mm_hashes: Optional[List[str]] - - -class TensorWrapper: - """Wrapper to keep tensor alive while exposing buffer for zero-copy.""" - - def __init__(self, tensor): - # Ensure tensor is on CPU and contiguous - if tensor.is_cuda: - tensor = tensor.cpu() - if not tensor.is_contiguous(): - tensor = tensor.contiguous() - - # Keep tensor reference - self.tensor = tensor - self.shape = list(tensor.shape) - self.dtype = tensor.dtype - - def __buffer__(self): - data_ptr = self.tensor.data_ptr() - total_bytes = self.tensor.numel() * self.tensor.element_size() - c_obj = (ctypes.c_char * total_bytes).from_address(data_ptr) - c_obj._keep_alive_ref = self - return memoryview(c_obj) - - -def _convert(data): - if isinstance(data, torch.Tensor): - return data - elif isinstance(data, np.ndarray): - return torch.tensor(data) - elif isinstance(data, list) and isinstance(data[0], np.ndarray): - return torch.tensor(np.array(data)) - elif isinstance(data, list) and isinstance(data[0], (int, float)): - return torch.tensor(data) - else: - return data - - -_mm_grid_attrs = { - # Kimi K2.5/K3 HF processors use grid_thws (see base_processor.ATTR_NAME_TO_MODALITY). - Modality.IMAGE: ["image_grid_thw", "image_grid_hws", "grid_thws"], - Modality.VIDEO: ["video_grid_thw"], - Modality.AUDIO: ["audio_feature_lens_raw"], -} - -_mm_feature_attrs = { - Modality.IMAGE: ["pixel_values"], - Modality.VIDEO: ["pixel_values_videos"], - Modality.AUDIO: ["input_features"], -} - - -def _get_mm_grid_dim(mm_inputs, modality, model_type: Optional[str] = None): - # Kimi K2.5/K3 vision processors only emit `grid_thws`; prefer it over generic keys - # so we never pick a mis-typed or stale `image_grid_hws` field from kwargs. - attrs = _mm_grid_attrs[modality] - model_type = (model_type or "").lower() - if modality == Modality.IMAGE: - # Kimi K2.5/K3 emit grid_thws, while Kimi-VL emits image_grid_hws. - # Other model types keep the generic attr order above. - if model_type in ("kimi_k25", "kimi_k3"): - attrs = ("grid_thws", "image_grid_thw", "image_grid_hws") - elif model_type == "kimi_vl": - attrs = ("image_grid_hws", "image_grid_thw", "grid_thws") - - for attr in attrs: - if attr in mm_inputs and mm_inputs[attr] is not None: - return _convert(mm_inputs[attr]) - raise ValueError(f"Grid dim ({_mm_grid_attrs[modality]}) not found in {mm_inputs}") - - -def _get_mm_feature(mm_inputs, modality): - for attr in _mm_feature_attrs[modality]: - if attr in mm_inputs: - return mm_inputs[attr] - raise ValueError( - f"Feature attrs ({_mm_feature_attrs[modality]}) not found in {mm_inputs}" - ) - - -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. - aux = { - attr: _normalize_aux_value(mm_inputs.get(attr)) - for attr in video_meta_attrs_for(model_type) - } - if model_type == "kimi_k3": - aux["original_image_sizes"] = _normalize_aux_value( - mm_inputs.get("original_image_sizes") - ) - return aux - - -def _get_original_image_size(image): - """Return an image's original (width, height) before encoder preprocessing.""" - if isinstance(image, dict): - image = image.get("image") - if isinstance(image, torch.Tensor): - if image.ndim < 2: - raise ValueError(f"Invalid image tensor shape: {tuple(image.shape)}") - return [int(image.shape[-1]), int(image.shape[-2])] - if hasattr(image, "size"): - width, height = image.size - return [int(width), int(height)] - raise TypeError(f"Cannot determine original image size from {type(image)}") - - -class MMEncoder: - def __init__( - self, - server_args: ServerArgs, - schedule_path=None, - dist_init_method=None, - rank: int = 0, - gpu_id: Optional[int] = None, - ): - """``gpu_id`` pins this encoder to a device other than - ``base_gpu_id + rank`` — the DP launcher's per-worker placement. It is - this instance's value, not a config change, so it travels as an - argument.""" - # The DP and TP encoder workers are spawned, so this constructor is - # the first publish in those processes. - publish(server_args, role="encoder") - logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") - self.server_args = server_args - configure_media_url_security( - get_mm().allowed_media_domains, - server_args.media_url_max_file_size_mb, - ) - self.rank = rank - # DP rank for metric labels; overridden by run_dp_worker in DP mode. - # 0 in the single-instance (non-DP) path. - self.dp_rank = 0 - self.profiler = EncoderProfiler(rank) - self._load_mm_processor(server_args) - - self.model_config = ModelConfig.from_server_args( - server_args, - ) - self.load_config = LoadConfig( - load_format=get_model().load_format, - download_dir=server_args.download_dir, - model_loader_extra_config=server_args.model_loader_extra_config, - remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, - remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, - remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, - ) - self.model_type = getattr( - self.model_config.hf_config, "model_type", "unknown" - ).lower() - - self.device = get_device().device - self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id - - self.device_config = DeviceConfig( - device=self.device, - gpu_id=self.gpu_id, - ) - - torch.get_device_module(self.device).set_device(self.gpu_id) - - self.use_image_processor_gpu = ( - use_image_processor_gpu - and resolve_image_processor_backend(server_args) != "pil" - ) - self._build_vision_config(get_mm().mm_process_config) - self.model_audio_sr = self._resolve_audio_sr() - logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz") - - init_distributed_environment( - backend=get_default_distributed_backend(self.device), - world_size=server_args.tp_size, - rank=rank, - distributed_init_method=dist_init_method, - local_rank=rank, - ) - initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size) - initialize_dp_attention(server_args, self.model_config) - - self.model = load_model( - model_config=self.model_config, - load_config=self.load_config, - device_config=self.device_config, - ) - self.encoder_media_processor_config = resolve_encoder_media_processor_config( - self.model - ) - maybe_precompile_model_kernels_after_loading(self.model, self.device) - - self.context = zmq.asyncio.Context(2) - self.sync_context = zmq.Context() # Reuse sync context for thread pool - self.scheduler_send_sockets = {} - self.scheduler_send_locks = {} - self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=10) - # Dedicated executor for image preprocessing (resize/normalize). - # Separate from self.executor (ZMQ sends) to avoid contention under high concurrency. - self.preproc_executor = concurrent.futures.ThreadPoolExecutor( - max_workers=envs.SGLANG_ENCODER_PREPROC_WORKERS.get() - ) - - embedding_cache_size = int(os.environ.get("SGLANG_VLM_CACHE_SIZE_MB", "4096")) - self.mm_cache = MultiModalStaticCache(embedding_cache_size * 1024 * 1024) - self.mm_cache_lock = asyncio.Lock() - - self.io_executor = concurrent.futures.ThreadPoolExecutor( - max_workers=int(os.environ.get("SGLANG_ENCODER_MM_LOAD_WORKERS", 4)) - ) - self.send_timeout = envs.SGLANG_ENCODER_SEND_TIMEOUT.get() - - if schedule_path is not None: - self.schedule_socket = get_zmq_socket( - self.context, zmq.PULL, schedule_path, True - ) - self.background_tasks: Set[asyncio.Task] = set() - - # Embedding dtype = model param dtype. Always available (both transfer - # backends and the global-cache pool rely on it). - self._embedding_dtype = next(self.model.parameters()).dtype - self._element_size = torch.tensor( - [], dtype=self._embedding_dtype - ).element_size() - - if get_mm().enable_mm_global_cache: - from sglang.srt.mem_cache.embedding_cache_controller import ( - EmbeddingCacheController, - ) - from sglang.srt.mem_cache.embedding_store import EmbeddingStoreFactory - - embedding_store = EmbeddingStoreFactory.create_backend( - get_mm().mm_global_cache_backend, - ) - hidden_dims = self._infer_embedding_dims() - self.mm_global_cache = EmbeddingCacheController( - rank, - server_args.tp_size, - embedding_store=embedding_store, - hidden_dims=hidden_dims, - tp_group=get_tp_group().cpu_group, - all_rank_get=False, - dtype=self._embedding_dtype, - ) - else: - self.mm_global_cache = None - - # Pre-compute embedding metadata (needed by all ranks for mooncake) - if get_disagg().encoder_transfer_backend == "mooncake": - self._embedding_dims = self._infer_embedding_dims() - - if self.rank == 0: - logger.info( - f"Using transfer backend: {get_disagg().encoder_transfer_backend}" - ) - - if get_disagg().encoder_transfer_backend == "mooncake": - self.local_ip = get_local_ip_auto() - - self.engine = get_mooncake_transfer_engine() - if self.engine is None: - from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( - init_mooncake_transfer_engine, - ) - - self.engine = init_mooncake_transfer_engine( - hostname=self.local_ip, - gpu_id=self.gpu_id, - ib_device=( - get_disagg().disaggregation_ib_device - or get_exec().moe.mooncake_ib_device - ), - ) - - 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 get_disagg().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 get_disagg().encoder_transfer_backend == "mooncake": - self._encode_fn = self.encode_with_global_cache_mooncake - else: - self._encode_fn = self.encode_with_global_cache - else: - if get_disagg().encoder_transfer_backend == "mooncake": - self._encode_fn = self.encode_with_mooncake - else: - self._encode_fn = self.encode - - logger.info(f"rank {rank} init finish ") - - def _infer_embedding_dims(self) -> dict: - """Infer per-modality embedding dimensions from hf_config at init time.""" - default = self.model_config.hidden_size - hf_cfg = self.model_config.hf_config - thinker_cfg = getattr(hf_cfg, "thinker_config", None) - dims = { - Modality.IMAGE: default, - Modality.VIDEO: default, - Modality.AUDIO: default, - } - - vision_cfg = getattr(thinker_cfg, "vision_config", None) or getattr( - hf_cfg, "vision_config", None - ) - if vision_cfg is not None: - out_hs = getattr(vision_cfg, "out_hidden_size", None) - if out_hs is not None: - ds = getattr(vision_cfg, "deepstack_visual_indexes", None) - vis_dim = ( - out_hs * (1 + len(ds)) - if isinstance(ds, (list, tuple)) and ds - else out_hs - ) - dims[Modality.IMAGE] = vis_dim - dims[Modality.VIDEO] = vis_dim - - audio_cfg = getattr(thinker_cfg, "audio_config", None) or getattr( - hf_cfg, "audio_config", None - ) - if audio_cfg is not None: - for attr in ("output_dim", "d_model"): - val = getattr(audio_cfg, attr, None) - if val and int(val) > 0: - dims[Modality.AUDIO] = int(val) - break - - logger.info(f"Global cache embedding dims: {dims}") - return dims - - def _resolve_audio_sr(self) -> int: - # Must match MiMoProcessor.from_hf_config — on drift, mimo tags the - # ndarray with its own audio_sampling_rate and skips resample, so the - # waveform is interpreted at the wrong rate and warped. - def _read(obj, attr): - if obj is None: - return None - if isinstance(obj, dict): - return obj.get(attr) - return getattr(obj, attr, None) - - audio_cfg = self.vision_config.get("audio", {}) - sr = audio_cfg.get("audio_sampling_rate") - if sr: - return int(sr) - - hf_cfg = self.model_config.hf_config - thinker_cfg = _read(hf_cfg, "thinker_config") - pc = _read(thinker_cfg, "processor_config") or _read(hf_cfg, "processor_config") - sr = _read(pc, "audio_sampling_rate") - if sr: - return int(sr) - ac = _read(thinker_cfg, "audio_config") or _read(hf_cfg, "audio_config") - for attr in ("sampling_rate", "sample_rate"): - sr = _read(ac, attr) - if sr: - return int(sr) - - sr = audio_cfg.get("sampling_rate") - if sr: - return int(sr) - logger.warning( - "No audio sampling rate found in mm_config or hf_config; " - "falling back to 16000 Hz. If the model expects a different SR " - "(e.g. MiMo-V2 defaults to 24000), audio will be warped." - ) - return 16000 - - def _build_vision_config(self, mm_process_config): - """ - Validate vision config, used for image/video/audio. - If not provided, keep default values. - """ - self.vision_config = ( - mm_process_config.get("vision_config", {}) - if mm_process_config is not None - else {} - ) - for modality_str in ["image", "video", "audio"]: - if not self.vision_config.get(modality_str, None): - self.vision_config[modality_str] = {} - if self.use_image_processor_gpu: - self.vision_config[modality_str]["device"] = self.device - - if modality_str == "video": - video_defaults = {"fps": 2.0, "max_frames": 768, "min_frames": 4} - for k, v in video_defaults.items(): - self.vision_config["video"].setdefault(k, v) - - if modality_str == "audio": - if "return_attention_mask" not in self.vision_config["audio"]: - self.vision_config["audio"]["return_attention_mask"] = True - if "padding" not in self.vision_config["audio"]: - if self.model_type == "qwen2_audio": - # For Qwen2Audio, use padding="max_length" - # (same as https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_audio/processing_qwen2_audio.py#L93) - self.vision_config["audio"]["padding"] = "max_length" - else: - self.vision_config["audio"]["padding"] = True - if "truncation" not in self.vision_config["audio"]: - # keep same logic as base_processor.py - if ( - hasattr(self, "audio_processor") - and self.audio_processor is not None - ): - if self.audio_processor.__class__.__name__ in { - "Gemma3nProcessor", - "GlmAsrProcessor", - "Qwen2AudioProcessor", - "Qwen3OmniMoeProcessor", - }: - self.vision_config["audio"]["truncation"] = False - - def _load_mm_processor(self, server_args: ServerArgs): - """ - Load image/video/audio processor separately, - avoid issues with AutoProcessor not recognizing certain models - """ - from transformers import AutoImageProcessor, AutoVideoProcessor - - image_processor_backend = resolve_image_processor_backend(server_args) - image_processor_kwargs = ( - {} - if image_processor_backend == "auto" - else {"backend": image_processor_backend} - ) - try: - self.image_processor = AutoImageProcessor.from_pretrained( - get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - **image_processor_kwargs, - ) - except Exception as e: - logger.warning(f"Failed to load image processor: {e}") - self.image_processor = None - - try: - self.video_processor = AutoVideoProcessor.from_pretrained( - get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - ) - except Exception as e: - logger.warning(f"Failed to load video processor: {e}") - self.video_processor = None - - try: - # Note: AutoProcessor is used for audio processor - _audio_proc = AutoProcessor.from_pretrained( - get_serving().tokenizer_path or get_model().model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - ) - if not hasattr(_audio_proc, "feature_extractor"): - logger.warning( - "Loaded AutoProcessor has no feature_extractor attribute, " - "audio processing will be unavailable." - ) - self.audio_processor = None - else: - self.audio_processor = _audio_proc - except Exception as e: - logger.warning(f"Failed to load audio processor: {e}") - self.audio_processor = None - - def _load_single_item( - self, - data, - modality: Modality, - frame_count_limit=None, - discard_alpha_channel=True, - ): - """ - Load a single multimodal data. - If data is precomputed, returns directly. - Static method that can be pickled for multiprocessing""" - media_metadata = {} - content_hash = None - if isinstance(data, dict): - if "url" not in data: - return data - media_metadata = {key: value for key, value in data.items() if key != "url"} - content_hash = parse_content_hash(data.get("content_hash")) - data = data["url"] - try: - if modality == Modality.IMAGE: - if content_hash is not None: - snapshot = snapshot_media(data) - if snapshot.content_digest != content_hash: - raise BadRequestError( - "Encoder media content hash mismatch: " - f"expected {content_hash}, got {snapshot.content_digest}" - ) - data = snapshot.data - gpu_image_decode = ( - self.encoder_media_processor_config.image_decode_mode - if self.use_image_processor_gpu - else False - ) - img, _ = load_image(data, gpu_image_decode) - if ( - discard_alpha_channel - and not isinstance(img, torch.Tensor) - and img.mode != "RGB" - ): - # Needed only when `img` is a PIL image - img = img.convert("RGB") - if ( - media_metadata - and self.encoder_media_processor_config.preserve_media_metadata - ): - return { - "type": "image", - "image": img, - **media_metadata, - } - return img - elif modality == Modality.VIDEO: - return load_video(data, frame_count_limit) - elif modality == Modality.AUDIO: - return load_audio(data, self.model_audio_sr) - - except MMError: - raise - except CLIENT_MEDIA_EXCEPTIONS as e: - # Not ValueError: the DP envelope classifies by `.code`, which only MMError carries. - raise BadRequestError(f"Error while loading data {data}: {e}") from e - except Exception as e: - raise RuntimeError(f"Error while loading data {data}: {e}") - - def submit_data_loading_tasks(self, items, modalities): - futures = [] - task_info = [] - - for data, modality in zip(items, modalities): - if modality is not None: - futures.append( - self.io_executor.submit( - self._load_single_item, - data, - modality, - ) - ) - task_info.append((modality, data)) - return futures, task_info - - def _get_feat_extract_output_lengths(self, feature_lens): - """ - Computes the output length of the convolutional layers and the output length of the audio encoder - """ - # qwen2_audio/qwen2.5_omni - if self.model_type in ["qwen2_audio", "qwen2_5_omni"]: - input_length = (feature_lens - 1) // 2 + 1 - return (input_length - 2) // 2 + 1 - # qwen3_asr / qwen3_omni_moe (same audio encoder architecture) - elif self.model_type in ["qwen3_asr", "qwen3_omni_moe"]: - input_lengths_leave = feature_lens % 100 - feat_lengths = (input_lengths_leave - 1) // 2 + 1 - output_lengths = ( - ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (feature_lens // 100) * 13 - ) - return output_lengths - elif self.model_type == "mimo_v2": - # MiMo-V2's preprocess_audio returns audio_token_len (already - # post-encoder/avg-pooler/group-size). Stored in audio_feature_lens_raw, - # so no further reduction here. - return feature_lens - else: - # fallback to original HF audio sample logic for other models - logger.warning( - f"Fallback to original HF audio sample logic for {self.model_type}" - ) - input_length = (feature_lens - 1) // 2 + 1 - return (input_length - 2) // 2 + 1 - - async def _flatten_and_load_videos(self, mm_items): - if not isinstance(mm_items, (list, tuple)): - mm_items = [mm_items] - - futures, _ = self.submit_data_loading_tasks( - mm_items, [Modality.VIDEO] * len(mm_items) - ) - async_futures = [asyncio.wrap_future(f) for f in futures] - video_items = await asyncio.gather(*async_futures) - - video_processor_kwargs = {} - if "qwen" in self.model_type: - # for qwen-series model, do sample frames before preprocess - video_processed = [ - await preprocess_video( - video, video_config=self.vision_config.get("video", {}) - ) - for video in video_items - ] - videos, video_metadata = map(list, zip(*video_processed)) - video_processor_kwargs["do_sample_frames"] = False - if video_metadata: - video_processor_kwargs["video_metadata"] = video_metadata - return videos, video_processor_kwargs - else: - raise NotImplementedError( - f"Video processing is not supported for {self.model_type} model." - ) - - async def _flatten_and_load_data_by_modality(self, mm_items, modality): - """ - Flatten mm_items structure, load multimodal data concurrently, and restore original structure. - - Returns: - Same structure as load_mm_items would return, support for image/audio - """ - # Handle single mm_item (not a list) - if not isinstance(mm_items, (list, tuple)): - futures, _ = self.submit_data_loading_tasks([mm_items], [modality]) - return await asyncio.wrap_future(futures[0]) - - # Handle nested list (list of lists) - if len(mm_items) > 0 and isinstance(mm_items[0], (list, tuple)): - # Flatten nested structure - flat_data = [] - flat_indices = [] # Track which group each item belongs to - for group_idx, item_group in enumerate(mm_items): - for item in item_group: - flat_data.append(item) - flat_indices.append(group_idx) - - # Submit all tasks concurrently - futures, _ = self.submit_data_loading_tasks( - flat_data, [modality] * len(flat_data) - ) - - # Wait for all tasks to complete asynchronously - async_futures = [asyncio.wrap_future(f) for f in futures] - results = await asyncio.gather(*async_futures) - - # Restore nested structure - nested_results = [[] for _ in range(len(mm_items))] - for idx, result in zip(flat_indices, results): - nested_results[idx].append(result) - - return nested_results - - # Handle simple list - else: - futures, _ = self.submit_data_loading_tasks( - mm_items, [modality] * len(mm_items) - ) - # Wait for all tasks to complete asynchronously - async_futures = [asyncio.wrap_future(f) for f in futures] - return await asyncio.gather(*async_futures) - - def get_num_patches( - self, grid: Union[torch.Tensor, List[int]], modality: Modality - ) -> int: - """Calculate number of raw patches (before merge/sampling). Used for pixel_values slicing.""" - if modality == Modality.AUDIO: - return int(grid.item()) - if self.model_type == "kimi_vl" and modality == Modality.IMAGE: - h, w = self._kimi_hw_from_patch_grid(grid) - return h * w - return int(grid[0] * grid[1] * grid[2]) - - @staticmethod - def _kimi_hw_from_patch_grid( - grid: Union[torch.Tensor, np.ndarray, List[int], Tuple[int, ...]], - ) -> Tuple[int, int]: - """Extract (height, width) from Kimi 2D or 3D patch-grid metadata.""" - if isinstance(grid, torch.Tensor): - values = grid.flatten().tolist() - elif isinstance(grid, np.ndarray): - values = grid.reshape(-1).tolist() - else: - values = np.asarray(grid).reshape(-1).tolist() - - if len(values) not in (2, 3): - raise ValueError( - f"Invalid Kimi image grid metadata: {values}; " - "expected [h, w] or [t, h, w]" - ) - return int(values[-2]), int(values[-1]) - - def _kimi_tokens_from_patch_grid(self, grid: Union[torch.Tensor, List[int]]) -> int: - """MoonViT + tpool: output len is (h//mh)*(w//mw); temporal dim is pooled (not t*h*w/merge^2).""" - h, w = self._kimi_hw_from_patch_grid(grid) - merge_h, merge_w = self.model_config.hf_config.vision_config.merge_kernel_size - return (h * w) // (merge_h * merge_w) - - def get_num_tokens( - self, grid: Union[torch.Tensor, List[int]], modality: Modality - ) -> int: - """Calculate number of tokens (after 2x2 merge). Used for mm_embedding slicing.""" - if modality == Modality.AUDIO: - input_length = self.get_num_patches(grid, modality) - return self._get_feat_extract_output_lengths(input_length) - else: - if ( - self.model_type in ["kimi_k25", "kimi_k3", "kimi_vl"] - and modality == Modality.IMAGE - ): - return self._kimi_tokens_from_patch_grid(grid) - merge_size = getattr(self.image_processor, "merge_size", 2) - return self.get_num_patches(grid, modality) // (merge_size**2) - - def slice_embedding( - self, mm_embedding: torch.Tensor, grid_thw: List, modality: Modality - ) -> List[torch.Tensor]: - """Slice a concatenated embedding tensor into individual image embeddings.""" - slices, offset = [], 0 - for grid in grid_thw: - count = self.get_num_tokens(grid, modality) - slices.append(mm_embedding[offset : offset + count]) - offset += count - return slices - - def _calculate_hashes_from_features( - self, mm_feature, grid_thw: List, modality: Modality, mm_inputs=None - ) -> List[int]: - """CPU Task: Compute hashes based on processed feature patches.""" - preprocessed_items = ( - get_encoder_preprocessed_items(mm_inputs) if mm_inputs is not None else None - ) - if preprocessed_items is not None: - if len(preprocessed_items) != len(grid_thw): - raise ValueError( - "Encoder preprocess item/grid mismatch: " - f"{len(preprocessed_items)} items != {len(grid_thw)} grids" - ) - hashes = [] - for item in preprocessed_items: - item.set_pad_value() - hashes.append(item.hash) - return hashes - - hashes = [] - if modality == Modality.AUDIO and isinstance(mm_feature, list): - for feature in mm_feature: - tmp_item = MultimodalDataItem(modality=modality, feature=feature) - tmp_item.set_pad_value() - hashes.append(tmp_item.hash) - return hashes - - offset = 0 - logger.info(f"{mm_feature.shape=} with {modality=}") - for grid in grid_thw: - num_patches = self.get_num_patches(grid, modality) - feature_slice = mm_feature[offset : offset + num_patches] - tmp_item = MultimodalDataItem(modality=modality, feature=feature_slice) - tmp_item.set_pad_value() - hashes.append(tmp_item.hash) - offset += num_patches - return hashes - - def _build_mm_data_items( - self, - mm_feature, - mm_inputs: dict, - indices: List[int], - modality: Modality, - grid_thw: Optional[List] = None, - ) -> List[MultimodalDataItem]: - """Build the model-facing items selected for one encoder forward. - - A model preprocessor can preserve an item-wise representation with - ``EncoderPreprocessOutput``. This path avoids concatenating and then - re-slicing features before encoder-DP knows which rank owns each item. - Legacy Hugging Face processor outputs retain their existing aggregate - tensor behavior. - """ - if grid_thw is None: - grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) - - preprocessed_items = get_encoder_preprocessed_items(mm_inputs) - if preprocessed_items is not None: - if len(preprocessed_items) != len(grid_thw): - raise ValueError( - "Encoder preprocess item/grid mismatch: " - f"{len(preprocessed_items)} items != {len(grid_thw)} grids" - ) - selected = [preprocessed_items[index] for index in indices] - if any(item.modality != modality for item in selected): - raise ValueError("Encoder preprocess output contains wrong modality") - return selected - - split_kimi_k3_images = ( - self.model_type == "kimi_k3" and modality == Modality.IMAGE - ) - - # 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 - # per-item shape. Image/video features are concatenated along the - # patch dim; slice by cumulative patch offsets and cat. - if modality == Modality.AUDIO: - if isinstance(mm_feature, list): - sub_feature = [mm_feature[i] for i in indices] - else: - sub_feature = mm_feature[list(indices)] - else: - sub_feature_list = [] - offsets = [0] - curr = 0 - for grid in grid_thw: - curr += self.get_num_patches(grid, modality) - offsets.append(curr) - for index in indices: - sub_feature_list.append(mm_feature[offsets[index] : offsets[index + 1]]) - if not split_kimi_k3_images: - sub_feature = torch.cat(sub_feature_list, dim=0) - - if split_kimi_k3_images: - mm_items = [ - MultimodalDataItem.from_dict( - { - "modality": modality, - "feature": _convert(feature), - } - ) - for feature in sub_feature_list - ] - else: - mm_items = [ - MultimodalDataItem.from_dict( - { - "modality": modality, - "feature": ( - sub_feature - if isinstance(sub_feature, list) - else _convert(sub_feature) - ), - } - ) - ] - - for key, value in mm_inputs.items(): - if key in _mm_feature_attrs.get(modality, []): - continue - value = _convert(value) - if key in _mm_grid_attrs.get(modality, []): - if split_kimi_k3_images: - for mm_item, index in zip(mm_items, indices): - mm_item.set(key, value[index : index + 1]) - else: - mm_items[0].set(key, value[indices]) - else: - for mm_item in mm_items: - mm_item.set(key, value) - return mm_items - - 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. - """ - if grid_thw is None: - grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) - mm_items = self._build_mm_data_items( - mm_feature, mm_inputs, indices, modality, grid_thw - ) - - forward_start = time.perf_counter() - with torch.inference_mode(): - new_embeddings = get_feature_fn(mm_items) - 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]) - if encoder_metrics_collector is not None: - encoder_metrics_collector.observe_model_forward( - time.perf_counter() - forward_start, modality=modality.name.lower() - ) - - sub_grids = [grid_thw[i] for i in indices] - return self.slice_embedding(new_embeddings, sub_grids, modality) - - async def _prepare_global_cache_context( - self, - mm_items, - modality: Modality, - req_id: str, - hashes: Optional[List[str]] = None, - ) -> GlobalCacheEncodeContext: - 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) - - # Hashes must be grid-space; a leaf-space list would size-mismatch - # rank>0's mask (zeros(num_items)) and deadlock TP. - if hashes is not None and len(hashes) != num_items: - raise BadRequestError( - f"User-supplied hashes length {len(hashes)} != grid count " - f"{num_items} for {self.model_type}/{modality.name}; hashes " - f"must be in grid space (1 per encoder grid entry)." - ) - - str_mm_hashes = None - if self.rank == 0: - if hashes is None: - mm_hashes = self._calculate_hashes_from_features( - mm_feature, grid_thw, modality, mm_inputs - ) - else: - mm_hashes = hashes - # L2 cache expects string keys for Mooncake. - str_mm_hashes = [str(h) for h in mm_hashes] - - return GlobalCacheEncodeContext( - req_id=req_id, - modality=modality, - mm_inputs=mm_inputs, - get_feature_fn=get_feature_fn, - grid_thw=grid_thw, - mm_feature=mm_feature, - num_items=num_items, - aux_data=_build_mm_aux_data(mm_inputs, self.model_type), - str_mm_hashes=str_mm_hashes, - ) - - def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): - if self.server_args.tp_size > 1: - torch.distributed.broadcast( - mask_tensor, - src=0, - group=self.mm_global_cache.prefetch_tp_group, - ) - - async def _lookup_global_cache( - self, - ctx: GlobalCacheEncodeContext, - ) -> Tuple[List[int], List[int]]: - if self.rank == 0: - exist_mask = await self.mm_global_cache.batch_is_exist(ctx.str_mm_hashes) - mask_tensor = torch.tensor( - [1 if e else 0 for e in exist_mask], dtype=torch.int32 - ) - else: - mask_tensor = torch.zeros(ctx.num_items, dtype=torch.int32) - - self._broadcast_global_cache_mask(mask_tensor) - - 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] - return missing_indices, hit_indices - - def _prefetch_global_cache_hits( - self, - ctx: GlobalCacheEncodeContext, - hit_indices: List[int], - ) -> List[str]: - if self.rank != 0 or not hit_indices: - return [] - - hit_hashes = [ctx.str_mm_hashes[i] for i in hit_indices] - hit_tokens = [ - self.get_num_tokens(ctx.grid_thw[i], ctx.modality) for i in hit_indices - ] - self.mm_global_cache.prefetch(ctx.req_id, hit_hashes, hit_tokens, ctx.modality) - return hit_hashes - - async def _wait_global_cache_prefetch( - self, - ctx: GlobalCacheEncodeContext, - hit_indices: List[int], - hit_hashes: List[str], - ) -> List[int]: - fallback_mask = torch.zeros(ctx.num_items, dtype=torch.int32) - if self.rank == 0 and hit_indices: - try: - - async def _wait_prefetch(): - while not self.mm_global_cache.check_prefetch_progress(ctx.req_id): - await asyncio.sleep(0.005) - - await asyncio.wait_for(_wait_prefetch(), timeout=60.0) - - for i, idx in enumerate(hit_indices): - if not self.mm_global_cache.has_local_embedding(hit_hashes[i]): - fallback_mask[idx] = 1 - num_partial_fail = int(fallback_mask.sum().item()) - if num_partial_fail > 0: - logger.warning( - f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} " - f"cache-hit items failed to load, falling back to ViT" - ) - except (asyncio.TimeoutError, Exception) as e: - logger.error( - f"Prefetch failed for req {ctx.req_id}: {e}. " - f"Falling back to ViT for {len(hit_indices)} hit items." - ) - for idx in hit_indices: - fallback_mask[idx] = 1 - - self._broadcast_global_cache_mask(fallback_mask) - fallback_indices = [ - i for i in range(ctx.num_items) if fallback_mask[i].item() == 1 - ] - return fallback_indices - - def _launch_global_cache_insert( - self, - ctx: GlobalCacheEncodeContext, - hashes: List[str], - d2h_handles: List[Any], - ): - if not hashes: - return - - async def _background_insert(): - await asyncio.to_thread( - self.mm_global_cache.wait_store_to_pool, - d2h_handles, - ) - await asyncio.to_thread( - self.mm_global_cache.insert_batch, - hashes, - ctx.modality, - ) - - task = asyncio.create_task(_background_insert()) - self.background_tasks.add(task) - task.add_done_callback(self.background_tasks.discard) - - @staticmethod - def _as_2d_tensor(tensor: torch.Tensor) -> torch.Tensor: - if tensor.ndim != 2: - tensor = tensor.reshape(-1, tensor.shape[-1]) - return tensor - - def _assemble_global_cache_cpu( - self, - ctx: GlobalCacheEncodeContext, - hit_indices: List[int], - missing_indices: List[int], - fallback_indices: List[int], - new_slices: List[torch.Tensor], - fallback_slices: List[torch.Tensor], - ) -> torch.Tensor: - miss_slice_pos = {idx: pos for pos, idx in enumerate(missing_indices)} - fallback_slice_pos = {idx: pos for pos, idx in enumerate(fallback_indices)} - fallback_index_set = set(fallback_indices) - token_counts = [ - self.get_num_tokens(grid, ctx.modality) for grid in ctx.grid_thw - ] - dim = self.mm_global_cache.get_embedding_dim(ctx.modality) - - mm_embedding = torch.empty( - (sum(token_counts), dim), - dtype=self._embedding_dtype, - pin_memory=True, - ) - - hit_view_hashes = [ - ctx.str_mm_hashes[idx] - for idx in hit_indices - if idx not in fallback_index_set - ] - hit_views = {} - try: - if hit_view_hashes: - cached_slice_lists = self.mm_global_cache.get_pool_views( - hit_view_hashes - ) - for h, slices in zip(hit_view_hashes, cached_slice_lists): - if slices is None: - raise InternalError( - f"Cached embedding {h} not available for req {ctx.req_id}" - ) - hit_views[h] = slices - - offset = 0 - for idx, num_tokens in enumerate(token_counts): - if idx in miss_slice_pos: - src = self._as_2d_tensor(new_slices[miss_slice_pos[idx]]) - mm_embedding[offset : offset + num_tokens].copy_( - src, non_blocking=True - ) - elif idx in fallback_slice_pos: - src = self._as_2d_tensor(fallback_slices[fallback_slice_pos[idx]]) - mm_embedding[offset : offset + num_tokens].copy_( - src, non_blocking=True - ) - else: - copied = 0 - for view in hit_views[ctx.str_mm_hashes[idx]]: - n = view.shape[0] - mm_embedding[offset + copied : offset + copied + n].copy_(view) - copied += n - offset += num_tokens - - torch.cuda.current_stream(self.device).synchronize() - return mm_embedding - finally: - if hit_view_hashes: - self.mm_global_cache.release_pool_views(hit_view_hashes) - - def _assemble_global_cache_gpu( - self, - ctx: GlobalCacheEncodeContext, - missing_indices: List[int], - fallback_indices: List[int], - new_slices: List[torch.Tensor], - fallback_slices: List[torch.Tensor], - ) -> torch.Tensor: - miss_slice_pos = {idx: pos for pos, idx in enumerate(missing_indices)} - fallback_slice_pos = {idx: pos for pos, idx in enumerate(fallback_indices)} - token_counts = [ - self.get_num_tokens(grid, ctx.modality) for grid in ctx.grid_thw - ] - embedding_dim = self.mm_global_cache.get_embedding_dim(ctx.modality) - mm_embedding = torch.empty( - (sum(token_counts), embedding_dim), - dtype=self._embedding_dtype, - device=self.device, - ) - - offset = 0 - copy_handles = [] - for idx, num_tokens in enumerate(token_counts): - if idx in miss_slice_pos: - mm_embedding[offset : offset + num_tokens].copy_( - new_slices[miss_slice_pos[idx]], - non_blocking=True, - ) - elif idx in fallback_slice_pos: - mm_embedding[offset : offset + num_tokens].copy_( - fallback_slices[fallback_slice_pos[idx]], - non_blocking=True, - ) - else: - handle = self.mm_global_cache.load_to_device_async( - ctx.str_mm_hashes[idx], mm_embedding, offset - ) - if handle is None: - raise InternalError( - f"Cached embedding {ctx.str_mm_hashes[idx]} disappeared " - f"during assembly for req {ctx.req_id}" - ) - copy_handles.append(handle) - offset += num_tokens - - self.mm_global_cache.wait_load_to_device(copy_handles) - torch.cuda.current_stream(mm_embedding.device).synchronize() - return mm_embedding - - async def encode_with_global_cache( - self, - mm_items, - modality: Modality, - req_id: str, - num_parts: int, - part_idx: int, - hashes: Optional[List[str]] = None, - ) -> torch.Tensor: - ctx = await self._prepare_global_cache_context( - mm_items, modality, req_id, hashes - ) - - missing_indices, hit_indices = await self._lookup_global_cache(ctx) - hit_hashes = self._prefetch_global_cache_hits(ctx, hit_indices) - - new_slices = [] - if missing_indices: - new_slices = self._encode_missing( - ctx.mm_feature, - ctx.mm_inputs, - missing_indices, - ctx.modality, - ctx.get_feature_fn, - ctx.grid_thw, - keep_on_gpu=True, - ) - - miss_d2h_handles = [] - if self.rank == 0 and new_slices: - miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] - miss_d2h_handles = self.mm_global_cache.store_to_pool_async( - miss_hashes, new_slices, ctx.modality - ) - - fallback_indices = await self._wait_global_cache_prefetch( - ctx, hit_indices, hit_hashes - ) - - fallback_slices = [] - fallback_d2h_handles = [] - if fallback_indices: - logger.info( - f"Req {ctx.req_id}: All ranks running ViT fallback " - f"for {len(fallback_indices)} items." - ) - fallback_slices = self._encode_missing( - ctx.mm_feature, - ctx.mm_inputs, - fallback_indices, - ctx.modality, - ctx.get_feature_fn, - ctx.grid_thw, - keep_on_gpu=True, - ) - if self.rank == 0: - fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] - fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( - fallback_hashes, fallback_slices, ctx.modality - ) - - if self.rank == 0: - mm_embedding = self._assemble_global_cache_cpu( - ctx, - hit_indices, - missing_indices, - fallback_indices, - new_slices, - fallback_slices, - ) - - new_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] - new_hashes += [ctx.str_mm_hashes[i] for i in fallback_indices] - self._launch_global_cache_insert( - ctx, - new_hashes, - miss_d2h_handles + fallback_d2h_handles, - ) - - self.embedding_to_send[ctx.req_id] = EmbeddingData( - ctx.req_id, - num_parts, - part_idx, - ctx.grid_thw, - ctx.modality, - mm_embedding, - **ctx.aux_data, - ) - if self.profiler is not None: - self.profiler.step() - return ( - mm_embedding.nbytes, - mm_embedding.shape[0], - mm_embedding.shape[1], - None, - 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: - ctx = await self._prepare_global_cache_context( - mm_items, modality, req_id, hashes - ) - - nbytes, total_tokens, embedding_dim, event = ( - self._setup_mooncake_async_encode( - ctx.req_id, - num_parts, - part_idx, - ctx.grid_thw, - ctx.modality, - ctx.aux_data, - ) - ) - - # 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: - missing_indices, hit_indices = await self._lookup_global_cache(ctx) - hit_hashes = self._prefetch_global_cache_hits(ctx, hit_indices) - - new_slices = [] - if missing_indices: - new_slices = self._encode_missing( - ctx.mm_feature, - ctx.mm_inputs, - missing_indices, - ctx.modality, - ctx.get_feature_fn, - ctx.grid_thw, - keep_on_gpu=True, - ) - - fallback_indices = await self._wait_global_cache_prefetch( - ctx, hit_indices, hit_hashes - ) - - fallback_slices = [] - if fallback_indices: - logger.info( - f"Req {ctx.req_id}: All ranks running ViT fallback " - f"for {len(fallback_indices)} items." - ) - fallback_slices = self._encode_missing( - ctx.mm_feature, - ctx.mm_inputs, - fallback_indices, - ctx.modality, - ctx.get_feature_fn, - ctx.grid_thw, - keep_on_gpu=True, - ) - - if self.rank == 0: - d2h_handles = [] - if new_slices: - miss_hashes = [ - ctx.str_mm_hashes[i] for i in missing_indices - ] - miss_handles = self.mm_global_cache.store_to_pool_async( - miss_hashes, new_slices, ctx.modality - ) - d2h_handles.extend(miss_handles) - if fallback_slices: - fallback_hashes = [ - ctx.str_mm_hashes[i] for i in fallback_indices - ] - fb_handles = self.mm_global_cache.store_to_pool_async( - fallback_hashes, fallback_slices, ctx.modality - ) - d2h_handles.extend(fb_handles) - - mm_embedding = self._assemble_global_cache_gpu( - ctx, - missing_indices, - fallback_indices, - new_slices, - fallback_slices, - ) - - new_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] - new_hashes += [ctx.str_mm_hashes[i] for i in fallback_indices] - self._launch_global_cache_insert( - ctx, - new_hashes, - d2h_handles, - ) - - self._forward_results[ctx.req_id]["embedding"] = mm_embedding - logger.info( - f"Global cache + VIT forward completed for " - f"{ctx.req_id}, shape={mm_embedding.shape}" - ) - except Exception as e: - logger.error( - f"Global cache + VIT forward failed for {ctx.req_id}: {e}" - ) - if self.rank == 0: - self._forward_results[ctx.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 {ctx.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 - self.model_audio_sr, restore original structure. - """ - return await self._flatten_and_load_data_by_modality(mm_items, Modality.AUDIO) - - async def _flatten_and_load_images(self, mm_items): - """ - Flatten mm_items structure, load images concurrently, and restore original structure. - """ - return await self._flatten_and_load_data_by_modality(mm_items, Modality.IMAGE) - - def _calculate_timestamps(self, indices, video_fps: float, merge_size: int = 2): - """Calculate timestamps for video frames, used for qwen3_vl models.""" - # refer to https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen3_vl/processing_qwen3_vl.py#L255 - if not isinstance(indices, list): - indices = indices.tolist() - if len(indices) % merge_size != 0: - indices.extend( - indices[-1] for _ in range(merge_size - len(indices) % merge_size) - ) - timestamps = [idx / video_fps for idx in indices] - # Frames are merged by merge_size, so we need to average the timestamps - # between the first/last frame within the temporal patch - timestamps = [ - (timestamps[i] + timestamps[i + merge_size - 1]) / 2 - for i in range(0, len(timestamps), merge_size) - ] - return timestamps - - @staticmethod - def _flatten_nested_items(items): - if not isinstance(items, (list, tuple)): - return [items] - - flat = [] - for item in items: - if isinstance(item, (list, tuple)): - flat.extend(MMEncoder._flatten_nested_items(item)) - else: - flat.append(item) - return flat - - def _grid_count_per_leaf(self, leaves: List, modality: Modality) -> List[int]: - """Number of grid entries each leaf produces under the model's processor. - - Most processors map 1 leaf → 1 grid. Kimi-VL/K2.5/K3 image processors expand - a leaf shaped {"type": "image", "image": [pil1, pil2, ...]} into N grids - (see _normalize_kimi_encoder_images). Cross-request batching needs these - counts to keep per-request boundaries aligned with grid_dim. - """ - if ( - self.model_type - not in ( - "kimi_k25", - "kimi_k3", - "kimi_vl", - ) - or modality != Modality.IMAGE - ): - return [1] * len(leaves) - - def count(leaf): - if ( - isinstance(leaf, dict) - and leaf.get("type") == "image" - and isinstance(leaf.get("image"), (list, tuple)) - ): - return len(self._flatten_nested_items(leaf["image"])) - return 1 - - return [count(leaf) for leaf in leaves] - - def _normalize_kimi_encoder_images(self, images): - """Normalize Kimi image inputs for the image processor call.""" - from PIL import Image as PILImage - - def wrap_one(img): - if isinstance(img, dict) and img.get("type") in ("image", "video_chunk"): - return [img] - if isinstance(img, PILImage.Image): - return [{"type": "image", "image": img}] - return [img] - - if not images: - return images - - # Disagg may supply nested lists from grouped routing. - images = self._flatten_nested_items(images) - - # Kimi-VL image processor expects a flat list of concrete images. - if self.model_type == "kimi_vl": - normalized = [] - for img in images: - if ( - isinstance(img, dict) - and img.get("type") == "image" - and "image" in img - ): - inner = img["image"] - if isinstance(inner, (list, tuple)): - normalized.extend(self._flatten_nested_items(inner)) - else: - normalized.append(inner) - else: - normalized.append(img) - return normalized - - # Kimi-K2.5/K3 vision processors expect media dicts. - normalized = [] - for img in images: - wrapped = wrap_one(img) - for media in wrapped: - # Some pipelines may produce {"type": "image", "image": [PIL]}. - # Split it into one media item per concrete image object. - if ( - isinstance(media, dict) - and media.get("type") == "image" - and isinstance(media.get("image"), (list, tuple)) - ): - for inner in self._flatten_nested_items(media["image"]): - normalized.append({**media, "image": inner}) - else: - normalized.append(media) - - return normalized - - async def _process_mm_items(self, mm_items, modality, log_metrics: bool = True): - model_preprocessor = getattr(self.model, "preprocess_mm_for_encoder", None) - - preprocess_start = time.perf_counter() - if modality == Modality.IMAGE: - processor_input = await self._process_image_items( - mm_items, model_preprocessor - ) - elif modality == Modality.VIDEO: - processor_input = await self._process_video_items( - mm_items, model_preprocessor - ) - elif modality == Modality.AUDIO: - processor_input = await self._process_audio_items( - mm_items, model_preprocessor - ) - else: - raise ValueError(f"Unsupported modality: {modality}") - if encoder_metrics_collector is not None and log_metrics: - encoder_metrics_collector.observe_preprocess( - time.perf_counter() - preprocess_start, modality=modality.name.lower() - ) - - target = self.model.thinker if hasattr(self.model, "thinker") else self.model - get_feature_method = getattr(target, f"get_{modality.name.lower()}_feature") - return processor_input, get_feature_method - - async def _process_image_items(self, mm_items, model_preprocessor): - if not (self.image_processor or model_preprocessor): - raise ValueError("No image processor available") - images = await self._flatten_and_load_images(mm_items) - if self.model_type in ["kimi_k25", "kimi_k3", "kimi_vl"]: - images = self._normalize_kimi_encoder_images(images) - original_image_sizes = [_get_original_image_size(item) for item in images] - if model_preprocessor: - processor_output = invoke_encoder_preprocessor( - model_preprocessor, - images, - Modality.IMAGE, - self.vision_config, - image_processor=self.image_processor, - use_gpu_preprocessing=self.use_image_processor_gpu, - ) - if ( - isinstance(processor_output, EncoderPreprocessOutput) - and processor_output.materialize_local_items is not None - ): - parallel = get_parallel() - await asyncio.get_running_loop().run_in_executor( - self.preproc_executor, - processor_output.materialize_for_rank, - parallel.attn_tp_rank, - parallel.attn_tp_size, - ) - return processor_output - image_config = self.vision_config.get("image", {}) - processor_input = await asyncio.get_running_loop().run_in_executor( - self.preproc_executor, - functools.partial(self.image_processor, images=images, **image_config), - ) - if self.model_type == "kimi_k3": - processor_input["original_image_sizes"] = original_image_sizes - return processor_input - - async def _process_video_items(self, mm_items, model_preprocessor): - if model_preprocessor: - return model_preprocessor(mm_items, Modality.VIDEO, self.vision_config) - if not self.video_processor: - raise ValueError("No video processor available") - - videos, video_processor_kwargs = await self._flatten_and_load_videos(mm_items) - processor_input = await asyncio.get_running_loop().run_in_executor( - self.preproc_executor, - functools.partial( - self.video_processor, videos=videos, **video_processor_kwargs - ), - ) - - # Get additional video metadata - if ( - self.model_type - in [ - "qwen3_vl", - "qwen3_vl_moe", - "qwen3_5", - "qwen3_5_moe", - "intern_s2_preview", - ] - and video_processor_kwargs.get("video_metadata", None) is not None - ): - video_metadata = video_processor_kwargs["video_metadata"] - try: - merge_size = ( - self.model_config.hf_config.vision_config.spatial_merge_size - ) - except (AttributeError, KeyError): - merge_size = 2 # Default merge_size - - video_timestamps = [] - for metadata in video_metadata: - video_fps = metadata.get("fps", None) or 24 # original video fps - frames_indices = metadata.get("frames_indices", None) - timestamps = self._calculate_timestamps( - frames_indices, video_fps, merge_size - ) - video_timestamps.append(timestamps) - processor_input["video_timestamps"] = video_timestamps - elif ( - self.model_type in ["qwen2_5_vl", "qwen2_5_omni", "qwen3_omni_moe"] - and processor_input.get("video_grid_thw", None) is not None - ): - video_grid_thw = processor_input["video_grid_thw"] - try: - temporal_patch_size = self.video_processor.temporal_patch_size - except AttributeError: - temporal_patch_size = 2 # Default temporal_patch_size - fps_list = [ - self.vision_config.get("video", {}).get("fps", None) or 2 - ] * len(video_grid_thw) - second_per_grid_ts = [(temporal_patch_size / fps) for fps in fps_list] - second_per_grid_ts_tensor = torch.tensor( - second_per_grid_ts, dtype=torch.float32 - ) - processor_input["second_per_grid_ts"] = second_per_grid_ts_tensor - - return processor_input - - async def _process_audio_items(self, mm_items, model_preprocessor): - # Await off the event loop so EncoderScheduler can accumulate - # cross-request batches during download. - audios = await self._flatten_and_load_audios(mm_items) - - if model_preprocessor: - return model_preprocessor(audios, Modality.AUDIO, self.vision_config) - - if not self.audio_processor: - raise ValueError("No audio processor available") - - audio_config = self.vision_config.get("audio", {}) - processor_input = await asyncio.get_running_loop().run_in_executor( - self.preproc_executor, - functools.partial( - self.audio_processor.feature_extractor, audios, **audio_config - ), - ) - processor_input["feature_attention_mask"] = processor_input.pop( - "attention_mask" - ) - input_lengths = torch.tensor( - processor_input["feature_attention_mask"].sum(-1), dtype=torch.long - ) - processor_input["audio_feature_lens_raw"] = input_lengths - output_lengths = self._get_feat_extract_output_lengths(input_lengths) - processor_input["audio_feature_lens"] = output_lengths - return processor_input - - async def _encode( - self, mm_items, modality: Modality, log_metrics: bool = True - ) -> torch.Tensor: - modality_str = modality.name.lower() - try: - # preprocess latency is observed inside _process_mm_items so all - # callers (encode / batch_encode / global-cache) are covered. - mm_inputs, get_feature_fn = await self._process_mm_items( - mm_items, modality, log_metrics=log_metrics - ) - except NotImplementedError as e: - raise InternalError(f"Not implemented error: {str(e)}") - except Exception as e: - raise BadRequestError(f"Failed to process mm items: {str(e)}") - try: - # support mm_cache - mm_embedding = None - mm_hash = None - mm_feature = _convert(_get_mm_feature(mm_inputs, modality)) - grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) - model_mm_items = self._build_mm_data_items( - mm_feature, - mm_inputs, - list(range(len(grid_thw))), - modality, - grid_thw, - ) - - cache_hit = False - use_mm_cache = get_mm().enable_prefix_mm_cache and log_metrics - if use_mm_cache: - for item in model_mm_items: - item.set_pad_value() - item_hashes = [item.hash for item in model_mm_items] - mm_hash = MultiModalStaticCache.combine_hashes(item_hashes) - async with self.mm_cache_lock: - mm_cache = self.mm_cache.get(item_hashes) - if mm_cache is not None: - mm_embedding = mm_cache.embedding - cache_hit = True - - if mm_embedding is None: - forward_start = time.perf_counter() - with torch.inference_mode(): - mm_embedding: torch.Tensor = get_feature_fn(model_mm_items) - mm_embedding = mm_embedding.cpu() - if len(mm_embedding.shape) != 2: - mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1]) - if encoder_metrics_collector is not None and log_metrics: - encoder_metrics_collector.observe_model_forward( - time.perf_counter() - forward_start, modality=modality_str - ) - - # Per-request cache hit metrics: tokens = embedding rows, files = - # logical multimodal items (not the legacy aggregate tensor count). - if use_mm_cache and encoder_metrics_collector is not None: - total_tokens = int(mm_embedding.shape[0]) - hit_tokens = total_tokens if cache_hit else 0 - encoder_metrics_collector.record_cache_tokens( - hit_tokens, total_tokens, modality=modality_str - ) - encoder_metrics_collector.record_cache_files( - len(model_mm_items) if cache_hit else 0, - len(model_mm_items), - modality=modality_str, - ) - - if use_mm_cache: - async with self.mm_cache_lock: - entries_before = len(self.mm_cache) - already_present = self.mm_cache.has(mm_hash) - inserted = self.mm_cache.set( - mm_hash, EmbeddingResult(embedding=mm_embedding) - ) - entries_after = len(self.mm_cache) - if encoder_metrics_collector is not None: - added = 0 if already_present else (1 if inserted else 0) - evictions = max(0, added - (entries_after - entries_before)) - if evictions > 0: - encoder_metrics_collector.inc_cache_evictions( - modality=modality_str, count=evictions - ) - encoder_metrics_collector.set_cache_state( - self.mm_cache.current_size, entries_after - ) - if self.profiler is not None: - self.profiler.step() - - aux_data = _build_mm_aux_data(mm_inputs, self.model_type) - - if modality == Modality.VIDEO and mm_inputs.get("video_audio_features"): - target = ( - self.model.thinker if hasattr(self.model, "thinker") else self.model - ) - encode_video_audio_fn = getattr(target, "encode_video_audio", None) - if encode_video_audio_fn is not None: - audio_forward_start = time.perf_counter() - audio_embedding = encode_video_audio_fn(mm_inputs) - if encoder_metrics_collector is not None and log_metrics: - encoder_metrics_collector.observe_model_forward( - time.perf_counter() - audio_forward_start, modality="audio" - ) - if audio_embedding is not None: - aux_data["video_audio_embedding"] = audio_embedding - else: - logger.warning( - "Videos carry audio tracks but model has no " - "encode_video_audio; dropping audio for EPD encoding." - ) - - return ( - grid_thw, - mm_embedding, - aux_data, - ) - except BadRequestError as e: - raise BadRequestError(f"Bad request error: {str(e)}") - except Exception as e: - raise InternalError(f"Internal encoding error: {str(e)}") - - async def _send( - self, - embedding: torch.Tensor, - mm_data: EmbeddingData, - session_id=None, - buffer_address=None, - prefill_host=None, - embedding_port=None, - url=None, - ): - if get_disagg().encoder_transfer_backend == "mooncake": - # 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})" - ) - - # Request-level shared MR, registered lazily on the first /send; - # deregistration is deferred to _cleanup_inflight_encode_state. - fwd_state = self._forward_results.setdefault(req_id, {}) - mr_already_registered = fwd_state.get("mr_ptr") == embedding.data_ptr() - if not mr_already_registered: - self.engine.register(embedding.data_ptr(), embedding.nbytes) - self._forward_results[req_id]["mr_ptr"] = embedding.data_ptr() - _t_xfer_start = time.monotonic() - xfer_ret = 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 encoder_metrics_collector is not None: - encoder_metrics_collector.observe_transfer( - xfer_ms / 1000.0, backend="mooncake" - ) - if xfer_ret < 0: - raise InternalError( - f"Mooncake transfer_sync failed for {req_id} " - f"(session={session_id}, nbytes={embedding.nbytes}, " - f"ret={xfer_ret})" - ) - # Only emit at INFO when transfer is slow or the MR was - # registered lazily by this /send; - 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}" - ) - - # Send ack/data - if url is not None: - endpoint = NetworkAddress.parse(url).to_tcp() - else: - endpoint = NetworkAddress(prefill_host, embedding_port).to_tcp() - logger.info(f"{endpoint = }") - - # Serialize data - if get_disagg().encoder_transfer_backend == "mooncake": - # 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() - if new_mm_data.error_msg is not None: - buffer = None - serialized_data = pickle.dumps(new_mm_data) - else: - embedding_tensor = TensorWrapper(mm_data.embedding) - serialized_data = pickle.dumps(new_mm_data) - buffer = embedding_tensor.__buffer__() - - _zmq_xfer_start = time.perf_counter() - if ( - get_disagg().encoder_transfer_backend == "zmq_to_scheduler" - and url is not None - ): - lock = self.scheduler_send_locks.get(endpoint) - if lock is None: - lock = asyncio.Lock() - self.scheduler_send_locks[endpoint] = lock - - async with lock: - sock = self.scheduler_send_sockets.get(endpoint) - if sock is None: - sock = self.context.socket(zmq.PUSH) - config_socket(sock, zmq.PUSH) - sock.setsockopt(zmq.IMMEDIATE, 1) - sock.setsockopt(zmq.SNDTIMEO, int(self.send_timeout * 1000)) - sock.connect(endpoint) - self.scheduler_send_sockets[endpoint] = sock - try: - frames = ( - [serialized_data, buffer] - if buffer is not None - else [serialized_data] - ) - tracker = await sock.send_multipart(frames, copy=False, track=True) - except Exception: - if self.scheduler_send_sockets.get(endpoint) is sock: - self.scheduler_send_sockets.pop(endpoint, None) - sock.close(linger=0) - raise - - # MessageTracker.wait() protects the zero-copy source buffer; it - # is not a receiver acknowledgement. Waiting under the per-peer - # lock serialized every large embedding on that TCP connection. - # Queue sends in order under the lock, then wait for buffer - # ownership independently so libzmq can pipeline the connection. - try: - await asyncio.to_thread(tracker.wait, self.send_timeout) - except Exception: - if self.scheduler_send_sockets.get(endpoint) is sock: - self.scheduler_send_sockets.pop(endpoint, None) - sock.close(linger=0) - raise - - if encoder_metrics_collector is not None: - encoder_metrics_collector.observe_transfer( - time.perf_counter() - _zmq_xfer_start, - backend=get_disagg().encoder_transfer_backend, - ) - return - - # Per-request sockets remain for zmq_to_tokenizer and legacy direct - # scheduler sends. Scheduler URL sends use persistent sockets above. - def send_with_socket(): - sock = self.sync_context.socket(zmq.PUSH) - config_socket(sock, zmq.PUSH) - sock.setsockopt(zmq.IMMEDIATE, 1) - sock.setsockopt(zmq.SNDTIMEO, int(self.send_timeout * 1000)) - try: - sock.connect(endpoint) - if buffer is not None: - tracker = sock.send_multipart( - [serialized_data, buffer], copy=False, track=True - ) - else: - tracker = sock.send_multipart( - [serialized_data], copy=False, track=True - ) - tracker.wait(timeout=self.send_timeout) - finally: - sock.close(linger=5000) - - await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) - if ( - encoder_metrics_collector is not None - and get_disagg().encoder_transfer_backend != "mooncake" - ): - encoder_metrics_collector.observe_transfer( - time.perf_counter() - _zmq_xfer_start, - backend=get_disagg().encoder_transfer_backend, - ) - - async def encode( - self, mm_items, modality: Modality, req_id, num_parts, part_idx, hashes=None - ): - try: - log_metrics = not is_health_check_request(req_id) - grid_dim, mm_embedding, aux_data = await self._encode( - mm_items, modality, log_metrics=log_metrics - ) - - if self.rank == 0: - mm_data = EmbeddingData( - req_id, - num_parts, - part_idx, - grid_dim, - modality, - mm_embedding, - **aux_data, - ) - self.embedding_to_send[req_id] = mm_data - return ( - mm_embedding.nbytes, - mm_embedding.shape[0], - mm_embedding.shape[1], - 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 failed: {error_msg} {error_code = }") - if self.rank == 0: - 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 - logger.debug(f"Created error EmbeddingData: {mm_data}") - return 0, 0, 0, error_msg, error_code - - def _setup_mooncake_async_encode( - self, - req_id: str, - num_parts: int, - part_idx: int, - grid_thw, - modality: Modality, - aux_data: dict, - ): - """Setup metadata and event management for mooncake async encode. - Returns (nbytes, total_tokens, embedding_dim, event).""" - total_tokens = sum(self.get_num_tokens(g, modality) for g in grid_thw) - embedding_dim = self._embedding_dims[modality] - nbytes = total_tokens * embedding_dim * self._element_size - - event = None - if self.rank == 0: - mm_data = EmbeddingData( - req_id, - num_parts, - part_idx, - grid_thw, - modality, - embedding=None, - embedding_shape=[total_tokens, embedding_dim], - **aux_data, - ) - self.embedding_to_send[req_id] = mm_data - event = asyncio.Event() - self._forward_ready_events[req_id] = event - self._forward_results[req_id] = {} - - return nbytes, total_tokens, embedding_dim, event - - def _handle_mooncake_encode_error( - self, req_id, num_parts, part_idx, modality, error_msg, error_code - ): - """Handle outer exception for mooncake async encode methods.""" - if self.rank == 0: - if req_id in self._forward_ready_events: - self._forward_results[req_id] = {"error": error_msg} - self._forward_ready_events[req_id].set() - mm_data = EmbeddingData( - req_id, - num_parts, - part_idx, - None, - modality, - error_msg=error_msg, - error_code=error_code, - ) - self.embedding_to_send[req_id] = mm_data - return 0, 0, 0, error_msg, error_code - - def _launch_mooncake_background_task(self, coro): - """Launch an async background task and track it.""" - task = asyncio.create_task(coro) - self.background_tasks.add(task) - task.add_done_callback(self.background_tasks.discard) - return task - - async def _cleanup_inflight_encode_state(self, req_id: str): - if not hasattr(self, "_inflight_encode_events"): - return - mooncake_send_done_count.pop(req_id, None) - 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) - # 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}" - ) - forward_state.pop("embedding", None) - # Release the embedding only after the MR is deregistered. - if mm_data is not None: - mm_data.embedding = None - mm_data.cached_embedding = None - 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, self.model_type) - - # 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 model-facing items on all ranks. Owner-deferred processor - # outputs stay per-item until encoder-DP assigns them. - mm_feature = _convert(_get_mm_feature(mm_inputs, modality)) - model_mm_items = self._build_mm_data_items( - mm_feature, - mm_inputs, - list(range(len(grid_thw))), - modality, - grid_thw, - ) - - async def _run_forward(): - try: - with torch.inference_mode(): - emb = get_feature_fn(model_mm_items) - 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( - self, requests: List[dict], modality: Modality - ) -> List[Tuple[int, int, int, Optional[str], Optional[int]]]: - """Cross-request encoder fusion (image/audio). No cache path.""" - # items_per_req counts grid entries (post-expansion) so per-request - # slicing of grid_dim/final_slices stays aligned for processors that - # expand one leaf into multiple grids (e.g. Kimi-VL/K2.5/K3 dict-of-images). - flat_items, items_per_req = [], [] - for req in requests: - leaves = MMEncoder._flatten_nested_items(req["mm_items"]) - flat_items.extend(leaves) - items_per_req.append(sum(self._grid_count_per_leaf(leaves, modality))) - total = sum(items_per_req) - - if encoder_metrics_collector is not None: - modality_str = modality.name.lower() - for n in items_per_req: - encoder_metrics_collector.observe_mm_items_per_request( - n, modality=modality_str - ) - encoder_metrics_collector.observe_mm_items_per_batch( - total, modality=modality_str - ) - - try: - mm_inputs, get_feat = await self._process_mm_items(flat_items, modality) - except NotImplementedError as e: - return self._batch_set_error( - requests, modality, InternalError(f"Not implemented error: {e}") - ) - except Exception as e: - return self._batch_set_error( - requests, modality, BadRequestError(f"Failed to process mm items: {e}") - ) - - try: - mm_feature = _convert(_get_mm_feature(mm_inputs, modality)) - grid_dim = _get_mm_grid_dim(mm_inputs, modality, self.model_type) - if len(grid_dim) != total: - return self._batch_set_error( - requests, - modality, - InternalError( - f"Grid count mismatch for {self.model_type}/" - f"{modality.name}: {len(flat_items)} leaves across " - f"{len(requests)} requests → expected {total} grids " - f"(per-req {items_per_req}), but processor produced " - f"{len(grid_dim)}. Add tile-expansion handling in " - f"_grid_count_per_leaf." - ), - ) - - final_slices = self._encode_missing( - mm_feature, - mm_inputs, - list(range(total)), - modality, - get_feat, - ) - - if self.profiler is not None: - for _ in requests: - self.profiler.step() - aux_data = _build_mm_aux_data(mm_inputs, self.model_type) - results = [] - offset = 0 - for req, n in zip(requests, items_per_req): - slices = final_slices[offset : offset + n] - emb = slices[0] if n == 1 else torch.cat(slices, dim=0) - req_aux_data = {} - if aux_data.get("original_image_sizes") is not None: - req_aux_data["original_image_sizes"] = aux_data[ - "original_image_sizes" - ][offset : offset + n] - if self.rank == 0: - self.embedding_to_send[req["req_id"]] = EmbeddingData( - req["req_id"], - req["num_parts"], - req["part_idx"], - grid_dim[offset : offset + n], - modality, - emb, - **req_aux_data, - ) - results.append((emb.nbytes, emb.shape[0], emb.shape[1], None, None)) - offset += n - return results - except Exception as e: - return self._batch_set_error( - requests, modality, InternalError(f"Internal encoding error: {e}") - ) - - def _batch_set_error( - self, requests: List[dict], modality: Modality, exc: Exception - ) -> List[Tuple[int, int, int, str, int]]: - code = getattr(exc, "code", HTTPStatus.INTERNAL_SERVER_ERROR) - msg = str(exc) - logger.error(f"Rank {self.rank} batch_encode failed: {msg} {code = }") - if self.rank == 0: - for req in requests: - self.embedding_to_send[req["req_id"]] = EmbeddingData( - req["req_id"], - req["num_parts"], - req["part_idx"], - None, - modality, - error_msg=msg, - error_code=code, - ) - return [(0, 0, 0, msg, code)] * len(requests) - - # For zmq_to_tokenizer zmq_to_scheduler and mooncake - async def send( - self, req_id, prefill_host, embedding_port, session_id=None, buffer_address=None - ): - mm_data: EmbeddingData = self.embedding_to_send[req_id] - await self._send( - mm_data.embedding, - mm_data, - session_id=session_id, - buffer_address=buffer_address, - prefill_host=prefill_host, - embedding_port=embedding_port, - ) - - # For zmq_to_scheduler - async def send_with_url( - self, - req_id, - ): - mm_data = self.embedding_to_send.get(req_id) - if not mm_data: - return - sent_urls: Set[str] = set() - all_tasks: List[Tuple[asyncio.Task, str]] = [] - start_time = asyncio.get_running_loop().time() - timeout = self.send_timeout - cond = await get_condition(req_id) - - try: - while True: - async with rid_lock: - current_targets = rid_to_receive_endpoint.get(req_id, set()).copy() - expected_count = rid_to_receive_count.get(req_id) - - new_targets = current_targets - sent_urls - - if new_targets: - logger.info( - f"Found {len(new_targets)} new endpoints for {req_id}. Starting tasks..." - ) - for url in new_targets: - task = asyncio.create_task( - self._send( - mm_data.embedding, - mm_data, - url=url, - ) - ) - all_tasks.append((task, url)) - sent_urls.add(url) # Mark as handled immediately - if expected_count is not None and len(sent_urls) >= expected_count: - logger.info( - f"All {expected_count} endpoints initiated for {req_id}. Breaking loop." - ) - break - remaining = timeout - (asyncio.get_running_loop().time() - start_time) - if remaining <= 0: - logger.error( - f"[{req_id}] Timeout! Sent {len(sent_urls)}/{expected_count}" - ) - break - - async with cond: - try: - await asyncio.wait_for(cond.wait(), timeout=remaining) - except asyncio.TimeoutError: - continue - - if all_tasks: - logger.info( - f"Loop finished. Awaiting completion of {len(all_tasks)} sending tasks..." - ) - tasks_only = [t[0] for t in all_tasks] - results = await asyncio.gather(*tasks_only, return_exceptions=True) - - # Process results and log errors - for i, result in enumerate(results): - url = all_tasks[i][1] # Retrieve URL associated with the task - if isinstance(result, Exception): - logger.error(f"Failed to send to {url}: {result}") - else: - logger.debug(f"Successfully sent to {url}") - - logger.info(f"All tasks completed for req_id: {req_id}") - - finally: - logger.info(f"Cleaning up resources for req_id {req_id}") - async with rid_lock: - rid_to_receive_endpoint.pop(req_id, None) - rid_to_receive_count.pop(req_id, None) - async with cond_dict_lock: - rid_to_cond.pop(req_id, None) - self.embedding_to_send.pop(req_id, None) - - async def get_embedding_port(self, prefill_url): - async with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=1800) - ) as session: - response = await session.post( - f"{prefill_url}/embedding_bootstrap", - json={"embedding_port": None}, - ) - response_json = await response.json() - return response_json["embedding_port"] - - -class EncoderProfiler: - def __init__(self, rank: int): - self.rank = rank - self.profiler = None - self.steps_left = None - self.output_dir = None - self.prefix = None - self.profile_id = None - - def start(self, obj: ProfileReq): - if self.profiler is not None: - return False, "profiling already running" - - output_dir = obj.output_dir or os.getenv("SGLANG_TORCH_PROFILER_DIR", "/tmp") - os.makedirs(output_dir, exist_ok=True) - self.output_dir = output_dir - self.prefix = obj.profile_prefix or "encoder" - self.profile_id = str(time.time()) - - activities = obj.activities or ["CPU", "GPU"] - torch_activities = [] - if "CPU" in activities: - torch_activities.append(torch.profiler.ProfilerActivity.CPU) - if "GPU" in activities: - torch_activities.append(torch.profiler.ProfilerActivity.CUDA) - - profile_memory = "MEM" in activities - if not torch_activities and not profile_memory: - return False, "no supported activities" - - self.profiler = torch.profiler.profile( - activities=torch_activities, - with_stack=True if obj.with_stack is None else obj.with_stack, - record_shapes=False if obj.record_shapes is None else obj.record_shapes, - profile_memory=profile_memory, - ) - self.profiler.start() - self.steps_left = obj.num_steps - logger.info( - f"Encoder profiling started. output_dir={self.output_dir} profile_id={self.profile_id}" - ) - return True, None - - def step(self): - if self.profiler is None: - return - self.profiler.step() - if self.steps_left is not None: - self.steps_left -= 1 - if self.steps_left <= 0: - self.stop() - - def stop(self): - if self.profiler is None: - return False, "profiling not running" - self.profiler.stop() - filename = f"{self.prefix}-rank{self.rank}-{self.profile_id}.trace.json" - trace_path = os.path.join(self.output_dir, filename) - self.profiler.export_chrome_trace(trace_path) - logger.info("Encoder profiling saved to: %s", trace_path) - self.profiler = None - self.steps_left = None - return True, None - - -class PendingRequest: - __slots__ = ("request", "future", "submit_time") - - def __init__(self, request: dict, loop: asyncio.AbstractEventLoop): - self.request = request - self.future: asyncio.Future = loop.create_future() - self.submit_time = time.time() - - -# VIDEO excluded: per-video preprocess kwargs (do_sample_frames, video_metadata) -# vary per request and can't merge into one HF processor call. -_BATCHABLE_MODALITIES = {Modality.IMAGE, Modality.AUDIO} -_KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE = 2 - - -def _resolve_encoder_batch_policy( - model_type: str, - configured_max_batch_size: int, - max_batch_size_is_explicit: bool, -) -> Tuple[int, bool]: - """Return effective batch size and same-turn coalescing policy.""" - max_batch_size = max(1, int(configured_max_batch_size)) - coalesce_same_turn = model_type == "kimi_k3" - if coalesce_same_turn and not max_batch_size_is_explicit: - max_batch_size = min(max_batch_size, _KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE) - return max_batch_size, coalesce_same_turn - - -class EncoderScheduler: - """Aggregate concurrent /encode requests into bounded image/audio batches.""" - - def __init__( - self, - encoder: "MMEncoder", - send_sockets: List[zmq.Socket], - max_batch_size: int, - coalesce_same_turn: bool = False, - request_timeout: float = ENCODER_REQ_TIMEOUT, - ): - self.encoder = encoder - self.send_sockets = send_sockets - self.max_batch_size = max(1, int(max_batch_size)) - self.coalesce_same_turn = bool(coalesce_same_turn) - self.request_timeout = max(1.0, float(request_timeout)) - self.pending_queue: asyncio.Queue[PendingRequest] = asyncio.Queue() - self._worker_task: Optional[asyncio.Task] = None - - def start(self) -> None: - if self._worker_task is None: - self._worker_task = asyncio.create_task(self._batch_worker()) - logger.info( - "EncoderScheduler started with " - f"max_batch_size={self.max_batch_size}, " - f"coalesce_same_turn={self.coalesce_same_turn}" - ) - - async def stop(self) -> None: - if self._worker_task is not None: - self._worker_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await self._worker_task - self._worker_task = None - # Reject any requests still queued so their HTTP handlers don't hang. - while True: - try: - pending = self.pending_queue.get_nowait() - except asyncio.QueueEmpty: - break - if not pending.future.done(): - pending.future.set_exception(RuntimeError("EncoderScheduler stopped")) - - async def submit(self, request: dict) -> Tuple: - pending = PendingRequest(request, asyncio.get_running_loop()) - await self.pending_queue.put(pending) - try: - return await asyncio.wait_for(pending.future, timeout=self.request_timeout) - except asyncio.TimeoutError: - if not pending.future.done(): - pending.future.cancel() - req_id = request.get("req_id") - logger.error( - f"EncoderScheduler.submit timed out after {self.request_timeout}s " - f"for req_id={req_id}" - ) - raise - - async def _collect_batch(self) -> List[PendingRequest]: - batch = [await self.pending_queue.get()] - first_modality = Modality.from_str(batch[0].request.get("modality", "image")) - should_yield = ( - self.coalesce_same_turn - and self.max_batch_size > 1 - and first_modality in _BATCHABLE_MODALITIES - ) - if should_yield: - # Let HTTP handlers that arrived in the same event-loop turn enqueue - # before dispatch. Unlike a fixed sleep, this adds no millisecond-scale - # tax to an isolated request. - await asyncio.sleep(0) - while len(batch) < self.max_batch_size: - try: - batch.append(self.pending_queue.get_nowait()) - except asyncio.QueueEmpty: - break - return batch - - async def _batch_worker(self) -> None: - while True: - batch: List[PendingRequest] = [] - try: - batch = await self._collect_batch() - groups: Dict[Modality, List[PendingRequest]] = defaultdict(list) - for p in batch: - groups[ - Modality.from_str(p.request.get("modality", "image")) - ].append(p) - for modality, group in groups.items(): - await self._dispatch_group(group, modality) - except asyncio.CancelledError: - for p in batch: - if not p.future.done(): - p.future.set_exception(RuntimeError("EncoderScheduler stopped")) - raise - except Exception as e: - logger.error( - f"Error in EncoderScheduler batch worker: {e}", exc_info=True - ) - for p in batch: - if not p.future.done(): - p.future.set_exception(e) - - @staticmethod - def _validate_request_shape(req: dict) -> Optional[str]: - # Cheap pre-broadcast checks: shape errors that don't require running - # the HF processor. Once a request reaches TP workers they enter - # batch_encode and expect to join its collectives — a malformed batch - # that makes rank-0 bail mid-flight would deadlock the workers. - if not isinstance(req, dict): - return f"request is not a dict: {type(req).__name__}" - if not req.get("req_id"): - return "missing req_id" - if not req.get("mm_items"): - return "missing or empty mm_items" - if "num_parts" not in req or "part_idx" not in req: - return "missing num_parts / part_idx" - h = req.get("hashes") - if h is not None and not isinstance(h, (list, tuple, str, int, bytes)): - return f"hashes must be list/scalar, got {type(h).__name__}" - return None - - async def _dispatch_group( - self, group: List[PendingRequest], modality: Modality - ) -> None: - # Video can't fuse (per-video preprocess kwargs vary). - if modality not in _BATCHABLE_MODALITIES: - await self._dispatch_per_request(group, modality) - return - - # Drop structurally-bad requests before broadcasting; otherwise TP - # workers would join batch_encode collectives that rank-0 has already - # abandoned. - valid: List[PendingRequest] = [] - for p in group: - err = self._validate_request_shape(p.request) - if err is None: - valid.append(p) - continue - logger.error(f"Dropping req_id={p.request.get('req_id')} from batch: {err}") - if not p.future.done(): - p.future.set_exception(BadRequestError(err)) - if not valid: - return - group = valid - - requests = [p.request for p in group] - start = time.time() - modality_str = modality.name.lower() - if encoder_metrics_collector is not None: - for p in group: - encoder_metrics_collector.observe_queue_wait( - max(0.0, start - p.submit_time), modality=modality_str - ) - try: - # The scheduler is the sole owner of batched dispatch order. Keep - # the collective broadcast and rank-0 execution under the same - # lock, while allowing concurrent HTTP handlers to enqueue before - # waiting on their individual futures. - async with self.encoder.encode_dispatch_lock: - for sock in self.send_sockets: - sock_send( - sock, - wrap_as_pickle( - { - "type": "batch_encode", - "modality": modality.name, - "requests": requests, - "enter_time": start, - } - ), - ) - - logger.info( - f"Dispatching batch of {len(group)} {modality.name} requests" - ) - results = await self.encoder.batch_encode(requests, modality) - if len(group) > 1: - logger.info( - f"Batch of {len(group)} {modality.name} requests completed in " - f"{(time.time() - start) * 1000:.1f}ms" - ) - except Exception as e: - # batch_encode normally catches and returns errors via _batch_set_error. - # If it raised, rank-0 may have skipped a collective broadcast, leaving - # TP workers stuck. Don't try to recover — fail every pending future - # and let the client retry. Re-broadcasting would risk a deadlock. - logger.error(f"batch_encode raised: {e}", exc_info=True) - for p in group: - if not p.future.done(): - p.future.set_exception(e) - return - - if len(results) != len(group): - err = RuntimeError( - f"batch_encode returned {len(results)} results for {len(group)} requests" - ) - logger.error(str(err)) - for p in group: - if not p.future.done(): - p.future.set_exception(err) - return - - for p, result in zip(group, results): - if not p.future.done(): - p.future.set_result(result) - - async def _dispatch_per_request( - self, - group: List[PendingRequest], - modality: Modality, - ) -> None: - modality_str = modality.name.lower() - for p in group: - req = p.request - try: - start = time.time() - if encoder_metrics_collector is not None: - encoder_metrics_collector.observe_queue_wait( - max(0.0, start - p.submit_time), modality=modality_str - ) - # Count like batch_encode: flatten nested items and expand - # per-leaf grids so {"type": "image", "image": [p1, p2, ...]} - # counts as N, not 1. - leaves = MMEncoder._flatten_nested_items(req.get("mm_items", [])) - mm_count = sum(self.encoder._grid_count_per_leaf(leaves, modality)) - encoder_metrics_collector.observe_mm_items_per_request( - mm_count, modality=modality_str - ) - encoder_metrics_collector.observe_mm_items_per_batch( - mm_count, modality=modality_str - ) - for sock in self.send_sockets: - sock_send(sock, wrap_as_pickle(req)) - result = await self.encoder.encode_request(req, modality) - if not p.future.done(): - p.future.set_result(result) - except Exception as e: - logger.error( - f"Per-request encode failed for req_id={req.get('req_id')}: {e}" - ) - if not p.future.done(): - p.future.set_exception(e) - - -encoder: Optional[MMEncoder] = None -send_sockets: List[zmq.Socket] = [] -encoder_scheduler: Optional[EncoderScheduler] = None - -# Per-process encoder metrics collector. Set in launch_server (non-DP) and in -# run_dp_worker (DP mode, with the worker's dp_rank). None when metrics disabled. -encoder_metrics_collector: Optional[EncoderMetricsCollector] = None - -# DP mode (--dp-size > 1): each rank runs as a subprocess with its own -# MMEncoder on its own GPU; the main process only routes via ZMQ so the -# asyncio event loop is never blocked by GPU work. -dp_dispatcher: Optional["DPDispatcher"] = None - - -async def _push_embedding_to_prefill(enc: MMEncoder, request: dict) -> None: - # No-op for mooncake (its /send is separate). embedding_port=None is - # rejected upfront, so ports is always a concrete list here. - req_id = request["req_id"] - backend = get_disagg().encoder_transfer_backend - - if backend == "zmq_to_tokenizer": - await enc.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - ) - enc.embedding_to_send.pop(req_id, None) - return - - if backend == "zmq_to_scheduler": - ports = request["embedding_port"] - assert isinstance(ports, list) - await asyncio.gather( - *( - enc.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=p, - ) - for p in ports - ) - ) - enc.embedding_to_send.pop(req_id, None) - - -async def _dp_worker_encode_and_send( - enc: MMEncoder, - sched: Optional[EncoderScheduler], - request: dict, -) -> Optional[dict]: - # Mooncake returns metadata for main to forward; zmq inlines the send. - # Soft errors raise MMError so the dispatcher route maps them to HTTP. - req_id = request["req_id"] - time_stats_json = request.pop("time_stats_json", None) - time_stats = EncoderReqTimeStats() - if time_stats_json: - time_stats.decode_json(time_stats_json) - request["enter_time"] = time.time() - modality = Modality.from_str(request["modality"]) - time_stats.modality = modality.name.lower() - time_stats.set_metrics_collector(encoder_metrics_collector) - backend = get_disagg().encoder_transfer_backend - - # URL state lives in main process module globals; workers don't see it. - if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: - raise MMError( - "Encoder DP mode does not support zmq_to_scheduler with " - "embedding_port=None (URL state isn't synchronised to workers). " - "Provide an explicit embedding_port list, switch to mooncake / " - "zmq_to_tokenizer, or run without --dp-size.", - code=HTTPStatus.BAD_REQUEST, - ) - - time_stats.set_mm_encode_start_time() - encode_coro = ( - sched.submit(request) - if sched is not None and modality in _BATCHABLE_MODALITIES - else enc.encode_request(request, modality) - ) - try: - nbytes, embedding_len, embedding_dim, error_msg, error_code = await encode_coro - except asyncio.TimeoutError: - time_stats.trace_ctx.abort(abort_info={"reason": "encoder batch timed out"}) - raise - - if error_msg: - time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) - # zmq backends still forward an error EmbeddingData to P so it - # doesn't block; send failures here are swallowed. - try: - await _push_embedding_to_prefill(enc, request) - except Exception as e: - logger.error( - f"DP error-send failed for req_id={req_id}: {e}", exc_info=True - ) - # Free the error EmbeddingData stored during encode, or it leaks in - # embedding_to_send and pins /health into "busy" (a non-empty - # embedding_to_send reads as busy, skipping the probe). Neither path - # guarantees cleanup on its own: mooncake's _push_embedding_to_prefill - # is a no-op, and a swallowed zmq send failure above skips its own pop. - # zmq lacks the inflight attrs so _cleanup_inflight_encode_state would - # early-return on it — pop directly. Mirrors the non-DP error path. - if backend == "mooncake": - await enc._cleanup_inflight_encode_state(req_id) - else: - enc.embedding_to_send.pop(req_id, None) - raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) - - time_stats.set_mm_encode_end_time() - - if backend == "mooncake": - request.pop("mm_items", None) - request.update( - embedding_size=nbytes, - embedding_len=embedding_len, - embedding_dim=embedding_dim, - ) - # Free the held embedding if the follow-up /send never arrives (same - # send_timeout cleanup the non-DP path uses). - enc._schedule_inflight_encode_cleanup(req_id) - return request - - await _push_embedding_to_prefill(enc, request) - return None - - -async def _dp_worker_health_encode(enc: MMEncoder) -> None: - """Functional health probe run on a DP worker. - - Process-liveness (proc.sentinel) can't see a worker that's alive but - wedged — hung GPU, NCCL deadlock, stalled ZMQ, or a blocked event loop. - When idle, run a tiny dummy encode to exercise the VIT forward and surface - those stalls. No prefill destination: the embedding is discarded, mirroring - the non-DP /health path. Raises on encode failure so the worker envelope - carries ``_error`` back to the dispatcher. - """ - # Busy worker: in-flight traffic already proves liveness, so skip the probe - # and report healthy — same `embedding_to_send` signal the non-DP /health - # path uses. A wedged-but-busy worker never reaches here (it can't service - # the recv), so the dispatcher's broadcast still times out → 503. - if enc.embedding_to_send: - return None - - if enc.image_processor is not None: - mm_items = [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"] - modality = Modality.IMAGE - elif enc.audio_processor is not None: - mm_items = [f"data:audio/wav;base64,{MINIMUM_WAV_SILENCE_BASE64}"] - modality = Modality.AUDIO - else: - # No processor → can't functionally probe; liveness alone is healthy. - return None - - # uuid keeps rids unique across workers; a bare time.time() can collide. - req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" - try: - _, _, _, error_msg, error_code = await enc.encode( - mm_items=mm_items, - modality=modality, - req_id=req_id, - num_parts=1, - part_idx=0, - ) - finally: - # Never leave the dummy embedding sitting in the send map. - enc.embedding_to_send.pop(req_id, None) - - if error_msg: - raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) - - -class DPDispatcher: - """Routes encode requests across DP ranks by least-pending count.""" - - def __init__( - self, - dp_size: int, - dispatch_sockets: List, - result_socket, - worker_processes: List[mp.Process], - enable_metrics: bool = False, - labels: Optional[Dict[str, str]] = None, - ): - self.dp_size = dp_size - self.dispatch_sockets = dispatch_sockets - self.result_socket = result_socket - self.worker_processes = worker_processes - # Key = req_id for encode/broadcast, req_id + "_send" for mooncake /send. - self.pending_futures: List[Dict[str, asyncio.Future]] = [ - {} for _ in range(dp_size) - ] - self.req_id_to_rank: Dict[str, int] = {} - self._rr_counter = 0 - self._broadcast_counter = 0 - self._dead_ranks: Set[int] = set() - # req_id -> monotonic ts a mooncake mapping has waited for its /send. - self._pending_send_at: Dict[str, float] = {} - # Set when _result_listener gives up; makes alive_ranks report empty. - self._listener_failed = False - - # Prometheus gauge: pending requests per DP rank. Lives in the main - # process (the dispatcher), unlike the per-worker EncoderMetricsCollector. - self.labels = dict(labels or {}) - self.pending_gauge = None - if enable_metrics: - from prometheus_client import Gauge - - self.pending_gauge = Gauge( - name="sglang:encoder_dp_pending_requests", - documentation="Number of pending requests per encoder DP rank.", - labelnames=list(self.labels.keys()) + ["dp_rank"], - multiprocess_mode="mostrecent", - ) - - @property - def pending_counts(self) -> List[int]: - return [len(d) for d in self.pending_futures] - - def _update_pending_gauge(self) -> None: - """Push current pending counts to the Prometheus gauge (absolute set).""" - if self.pending_gauge is not None: - for i, c in enumerate(self.pending_counts): - self.pending_gauge.labels(**self.labels, dp_rank=str(i)).set(c) - - @property - def alive_ranks(self) -> List[int]: - # Empty if the result listener died; else ranks not marked dead. - if self._listener_failed: - return [] - return [r for r in range(self.dp_size) if r not in self._dead_ranks] - - @property - def all_ranks_alive(self) -> bool: - # Strict health (only /health uses this); routing still degrades. - return len(self.alive_ranks) == self.dp_size - - def start(self) -> None: - logger.info(f"DP dispatcher started: {self.dp_size} ranks (all remote)") - asyncio.create_task(self._result_listener()) - asyncio.create_task(self._worker_watchdog()) - asyncio.create_task(self._cleanup_stale_mappings()) - - def _drop_pending_and_mapping(self, rank: int, req_id: str) -> None: - # dispatch / broadcast failure: no follow-up /send expected. - self.pending_futures[rank].pop(req_id, None) - self.req_id_to_rank.pop(req_id, None) - self._update_pending_gauge() - - def _fail_pending_for_rank(self, rank: int, reason: str, error_type: str) -> None: - # Resolve a rank's outstanding futures with 503 so awaiters don't hang. - pending = self.pending_futures[rank] - for key, future in list(pending.items()): - if not future.done(): - future.set_result( - { - "req_id": key.removesuffix("_send"), - "_dp_type": "send" if key.endswith("_send") else "encode", - "content": None, - "_error": reason, - "_error_type": error_type, - "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), - } - ) - pending.pop(key, None) - self._update_pending_gauge() - - def _fail_all_pending(self, reason: str, error_type: str) -> None: - for rank in range(self.dp_size): - self._fail_pending_for_rank(rank, reason, error_type) - self.req_id_to_rank.clear() - self._pending_send_at.clear() - - @staticmethod - def _timeout_envelope(req_id: str, dp_type: str, reason: str) -> dict: - return { - "req_id": req_id, - "_dp_type": dp_type, - "content": None, - "_error": reason, - "_error_type": "TimeoutError", - "_error_code": int(HTTPStatus.GATEWAY_TIMEOUT), - } - - async def dispatch(self, request: dict) -> dict: - counts = self.pending_counts - # Skip ranks whose worker process has died. - alive_ranks = self.alive_ranks - if not alive_ranks: - raise MMError( - "All encoder DP workers are dead.", - code=HTTPStatus.SERVICE_UNAVAILABLE, - ) - min_p = min(counts[r] for r in alive_ranks) - candidates = [r for r in alive_ranks if counts[r] == min_p] - rank = candidates[self._rr_counter % len(candidates)] - self._rr_counter += 1 - req_id = request["req_id"] - self.req_id_to_rank[req_id] = rank - future = asyncio.get_running_loop().create_future() - self.pending_futures[rank][req_id] = future - self._update_pending_gauge() - logger.info( - f"MM-Encoder DP dispatch: req_id={req_id}, " - f"modality={request.get('modality', 'image')}, " - f"dp_rank={rank}, pending={self.pending_counts}" - ) - - try: - await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request)) - # An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip - # the watchdog, so bound the wait explicitly. - return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) - except asyncio.TimeoutError: - self._drop_pending_and_mapping(rank, req_id) - return self._timeout_envelope( - req_id, - "encode", - f"Encoder DP rank={rank} timed out after {ENCODER_REQ_TIMEOUT}s", - ) - except BaseException: - self._drop_pending_and_mapping(rank, req_id) - raise - - async def dispatch_send(self, request: dict) -> dict: - req_id = request["req_id"] - # /send arrived → stop tracking it for stale-mapping GC. - self._pending_send_at.pop(req_id, None) - if self._listener_failed: - return { - "req_id": req_id, - "_error": "encoder DP result listener stopped; cannot route /send", - "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), - } - rank = self.req_id_to_rank.get(req_id) - if rank is None: - logger.warning( - f"MM-Encoder dispatch_send: unknown req_id={req_id}, " - f"cannot route to worker" - ) - return {"req_id": req_id, "_error": f"Unknown req_id: {req_id}"} - if rank in self._dead_ranks: - # Worker died between encode and /send; embedding is gone. - self.req_id_to_rank.pop(req_id, None) - return { - "req_id": req_id, - "_error": f"DP worker rank={rank} died before /send for req_id={req_id}", - "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), - } - key = req_id + "_send" - future = asyncio.get_running_loop().create_future() - self.pending_futures[rank][key] = future - request["_dp_type"] = "send" - logger.info( - f"MM-Encoder DP dispatch_send: req_id={req_id}, " - f"dp_rank={rank}, pending={self.pending_counts}" - ) - try: - await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request)) - return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) - except asyncio.TimeoutError: - self.pending_futures[rank].pop(key, None) - self.req_id_to_rank.pop(req_id, None) - return self._timeout_envelope( - req_id, - "send", - f"Encoder DP rank={rank} /send timed out after {ENCODER_REQ_TIMEOUT}s", - ) - except BaseException: - self.pending_futures[rank].pop(key, None) - self.req_id_to_rank.pop(req_id, None) - raise - - async def broadcast( - self, request: dict, timeout: Optional[float] = None - ) -> List[dict]: - # Skip dead ranks: a PUSH to a gone worker would just buffer and then - # surface as a spurious per-rank timeout. All dead → 503 (same as - # dispatch), which the profile endpoints turn into an HTTP error. - eff_timeout = timeout if timeout is not None else ENCODER_REQ_TIMEOUT - alive_ranks = self.alive_ranks - if not alive_ranks: - raise MMError( - "All encoder DP workers are dead.", - code=HTTPStatus.SERVICE_UNAVAILABLE, - ) - batch_id = self._broadcast_counter - self._broadcast_counter += 1 - rank_keys: List[Tuple[int, str]] = [] - futures: List[asyncio.Future] = [] - dp_type = request.get("_dp_type", "unknown") - try: - for rank in alive_ranks: - req_id = f"_broadcast_{batch_id}_{rank}" - future = asyncio.get_running_loop().create_future() - self.pending_futures[rank][req_id] = future - self.req_id_to_rank[req_id] = rank - rank_keys.append((rank, req_id)) - request_copy = {**request, "req_id": req_id} - await async_sock_send( - self.dispatch_sockets[rank], wrap_as_pickle(request_copy) - ) - futures.append(future) - # Concurrent wait → total bounded by eff_timeout, not - # dp_size × eff_timeout. - outcomes = await asyncio.gather( - *(asyncio.wait_for(fut, timeout=eff_timeout) for fut in futures), - return_exceptions=True, - ) - results: List[dict] = [] - for (rank, req_id), outcome in zip(rank_keys, outcomes): - if isinstance(outcome, asyncio.TimeoutError): - self._drop_pending_and_mapping(rank, req_id) - results.append( - self._timeout_envelope( - req_id, - dp_type, - f"Encoder DP rank={rank} broadcast timed out " - f"after {eff_timeout}s", - ) - ) - elif isinstance(outcome, BaseException): - self._drop_pending_and_mapping(rank, req_id) - raise outcome - else: - results.append(outcome) - return results - except BaseException: - for rank, req_id in rank_keys: - self._drop_pending_and_mapping(rank, req_id) - raise - - async def _worker_watchdog(self) -> None: - # proc.sentinel becomes readable on process exit; fail this rank's - # pending futures so awaiters don't hang on a dead worker. - loop = asyncio.get_running_loop() - watch: Dict[int, asyncio.Future] = {} - for rank, proc in enumerate(self.worker_processes): - fut: asyncio.Future = loop.create_future() - - # add_reader is level-triggered, so remove_reader inside the - # callback to avoid spinning every loop iteration. - def _on_exit(r=rank, f=fut, p=proc, lp=loop): - try: - lp.remove_reader(p.sentinel) - except (ValueError, OSError): - pass - if not f.done(): - f.set_result(r) - - try: - loop.add_reader(proc.sentinel, _on_exit) - except (ValueError, OSError): - continue - watch[rank] = fut - - while watch: - done, _ = await asyncio.wait( - watch.values(), return_when=asyncio.FIRST_COMPLETED - ) - for fut in done: - rank = fut.result() - proc = self.worker_processes[rank] - logger.error( - f"DP worker rank={rank} (pid={proc.pid}) exited " - f"with code={proc.exitcode}; failing pending requests" - ) - self._dead_ranks.add(rank) - reason = f"DP worker rank={rank} died (exitcode={proc.exitcode})" - self._fail_pending_for_rank(rank, reason, "WorkerDied") - self.req_id_to_rank = { - r: rk for r, rk in self.req_id_to_rank.items() if rk != rank - } - watch.pop(rank, None) - - async def _result_listener(self) -> None: - # Bounded back-off + give-up so a torn-down context exits in ~3s - # rather than spinning forever on recv errors. - consecutive_errors = 0 - while True: - try: - msg = await async_sock_recv(self.result_socket) - consecutive_errors = 0 - except asyncio.CancelledError: - raise - except Exception: - consecutive_errors += 1 - logger.error("_result_listener recv error", exc_info=True) - if consecutive_errors >= 30: - logger.error( - "_result_listener giving up after 30 consecutive errors" - ) - self._listener_failed = True - self._fail_all_pending( - "encoder DP result listener stopped after repeated " - "recv errors", - "ResultListenerStopped", - ) - return - await asyncio.sleep(min(0.1 * consecutive_errors, 1.0)) - continue - req_id = msg.get("req_id", "") - dp_type = msg.get("_dp_type", "encode") - key = (req_id + "_send") if dp_type == "send" else req_id - rank = self.req_id_to_rank.get(req_id) - if rank is None or key not in self.pending_futures[rank]: - logger.warning( - f"_result_listener: no pending future for " - f"req_id={req_id}, dp_type={dp_type}, dropping" - ) - continue - future = self.pending_futures[rank].pop(key) - self._update_pending_gauge() - # Only mooncake encode (content=request dict) needs the mapping - # kept for the follow-up /send. - keep_mapping = dp_type == "encode" and msg.get("content") is not None - if keep_mapping: - self._pending_send_at[req_id] = time.monotonic() - else: - self.req_id_to_rank.pop(req_id, None) - try: - future.set_result(msg) - - except asyncio.InvalidStateError: - logger.warning( - f"_result_listener: future already done for " - f"req_id={req_id}, dp_type={dp_type}" - ) - - async def _cleanup_stale_mappings(self) -> None: - # Evict req_id->rank mappings whose /send never came. The worker frees - # its own embedding via the send_timeout cleanup scheduled at encode, - # so both sides key off the same timeout. - ttl = envs.SGLANG_ENCODER_SEND_TIMEOUT.get() - interval = max(ttl / 4, 30) - while True: - await asyncio.sleep(interval) - now = time.monotonic() - stale = [rid for rid, ts in self._pending_send_at.items() if now - ts > ttl] - for rid in stale: - self._pending_send_at.pop(rid, None) - self.req_id_to_rank.pop(rid, None) - if stale: - logger.warning( - f"Evicted {len(stale)} stale encoder DP /send mapping(s) " - f"with no /send within {ttl}s" - ) - - -async def _dp_worker_handle_profile( - enc: MMEncoder, dp_rank: int, dp_type: str, request: dict -) -> dict: - prefix = f"dp_rank={dp_rank}: " - if dp_type == "start_profile": - req = request.get("profile_req") or ProfileReq() - req.req_type = ProfileReqType.START_PROFILE - if enc.profiler is None: - enc.profiler = EncoderProfiler(dp_rank) - ok, msg = enc.profiler.start(req) - detail = ( - f"started profiling, output_dir={enc.profiler.output_dir}" if ok else msg - ) - else: # stop_profile - if enc.profiler is None: - return {"ok": False, "msg": prefix + "profiling not initialized"} - ok, msg = enc.profiler.stop() - detail = "stopped profiling" if ok else msg - return {"ok": ok, "msg": prefix + detail} - - -async def _dp_worker_handle_request( - enc: MMEncoder, - sched: EncoderScheduler, - send_sock, - send_lock: asyncio.Lock, - dp_rank: int, - request: dict, - dp_type: str, -) -> None: - t0 = time.time() - modality_str = str(request.get("modality", "image")).lower() - is_encode = dp_type not in ( - "start_profile", - "stop_profile", - "health_encode", - "send", - ) - if is_encode and encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_received(modality=modality_str) - try: - if dp_type in ("start_profile", "stop_profile"): - content = await _dp_worker_handle_profile(enc, dp_rank, dp_type, request) - elif dp_type == "health_encode": - content = await _dp_worker_health_encode(enc) - elif dp_type == "send": - req_id = request["req_id"] - await enc.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - session_id=request["session_id"], - buffer_address=request["buffer_address"], - ) - # cancels the scheduled cleanup + frees embedding/forward state - await enc._cleanup_inflight_encode_state(req_id) - content = None - else: - content = await _dp_worker_encode_and_send(enc, sched, request) - - logger.info( - f"MM-Encoder [dp_rank={dp_rank}] {dp_type} done: " - f"req_id={request.get('req_id', '?')}, " - f"modality={request.get('modality', 'image')}, " - f"cost={(time.time() - t0) * 1000:.1f}ms" - ) - if is_encode and encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="success" - ) - envelope = { - "req_id": request.get("req_id", ""), - "_dp_type": dp_type, - "content": content, - } - except Exception as e: - logger.error( - f"DP worker {dp_rank} error on {dp_type} " - f"req_id={request.get('req_id', '?')}: {e}", - exc_info=True, - ) - if is_encode and encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="error" - ) - err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR) - envelope = { - "req_id": request.get("req_id", ""), - "_dp_type": dp_type, - "content": None, - "_error": str(e), - "_error_type": type(e).__name__, - "_error_code": err_code, - } - - # pyzmq async send isn't safe for concurrent senders. - try: - async with send_lock: - await async_sock_send(send_sock, wrap_as_pickle(envelope)) - except Exception: - logger.error( - f"DP worker {dp_rank} failed to send envelope for " - f"req_id={request.get('req_id', '?')}", - exc_info=True, - ) - - -async def run_dp_worker( - server_args: ServerArgs, - dp_rank: int, - gpu_id: int, - dispatch_path: str, - result_path: str, -): - logger.info( - f"DP worker {dp_rank} starting on gpu_id={gpu_id} " - f"(CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES', 'unset')})" - ) - - # gpu_id is the device chosen by maybe_reindex_device_id in the parent: - # 0 when CVD is pinned to one GPU, else the absolute id. - enc = MMEncoder( - server_args, - dist_init_method=f"tcp://127.0.0.1:{get_free_port()}", - rank=0, - gpu_id=gpu_id, - ) - - global encoder_metrics_collector - if get_observability().enable_metrics: - set_prometheus_multiproc_dir() - labels = { - "model_name": get_serving().served_model_name, - "dp_rank": str(dp_rank), - } - if get_observability().extra_metric_labels: - labels.update(get_observability().extra_metric_labels) - encoder_metrics_collector = EncoderMetricsCollector(labels) - enc.dp_rank = dp_rank - - max_batch_size, coalesce_same_turn = _resolve_encoder_batch_policy( - enc.model_type, - ENCODER_MAX_BATCH_SIZE, - ENCODER_MAX_BATCH_SIZE_EXPLICIT, - ) - sched = EncoderScheduler( - encoder=enc, - send_sockets=[], - max_batch_size=max_batch_size, - coalesce_same_turn=coalesce_same_turn, - ) - - ctx = zmq.asyncio.Context(2) - recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False) - send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False) - send_lock = asyncio.Lock() - inflight: Set[asyncio.Task] = set() - # Acquire-before-recv → back-pressure propagates to the dispatcher - # PUSH buffer. Must be at least max_batch_size or batching degrades. - max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get() - if max_inflight < max_batch_size: - logger.warning( - f"SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT={max_inflight} is below " - f"the effective encoder max_batch_size={max_batch_size}; the encoder " - f"will never assemble a full batch." - ) - inflight_sem = asyncio.Semaphore(max_inflight) - sched.start() - logger.info(f"DP worker {dp_rank} ready") - - # Task-per-request so EncoderScheduler.pending_queue accumulates and - # actual cross-request batching can happen. - try: - while True: - await inflight_sem.acquire() - # Released by _run on success or the outer finally if not spawned. - spawned = False - try: - try: - request = await async_sock_recv(recv_sock) - except asyncio.CancelledError: - raise - except Exception: - logger.error(f"DP worker {dp_rank} recv error", exc_info=True) - continue - if not isinstance(request, dict): - logger.error( - f"DP worker {dp_rank} received non-dict request " - f"({type(request).__name__}); dropping" - ) - continue - dp_type = request.pop("_dp_type", "encode") - - async def _run(req=request, t=dp_type): - try: - await _dp_worker_handle_request( - enc, sched, send_sock, send_lock, dp_rank, req, t - ) - finally: - inflight_sem.release() - - task = asyncio.create_task(_run()) - # Ownership transferred to _run; mark before any op that could - # raise (theoretical: set.add / add_done_callback) and cause a - # double-release. - spawned = True - inflight.add(task) - task.add_done_callback(inflight.discard) - finally: - if not spawned: - inflight_sem.release() - finally: - # Close zmq on exception/cancellation (normal stop is parent SIGKILL). - for task in inflight: - task.cancel() - ctx.destroy(linger=0) - - -def launch_dp_worker( - server_args: ServerArgs, - dp_rank: int, - gpu_id: int, - dispatch_path: str, - result_path: str, -): - try: - configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") - asyncio.run( - run_dp_worker(server_args, dp_rank, gpu_id, dispatch_path, result_path) - ) - except KeyboardInterrupt: - logger.info(f"DP worker {dp_rank} exiting") - except Exception: - traceback.print_exc() - - -@contextlib.asynccontextmanager -async def _lifespan(app: FastAPI): - global encoder_scheduler - if dp_dispatcher is not None: - dp_dispatcher.start() - yield - return - if encoder is not None: - max_batch_size, coalesce_same_turn = _resolve_encoder_batch_policy( - encoder.model_type, - ENCODER_MAX_BATCH_SIZE, - ENCODER_MAX_BATCH_SIZE_EXPLICIT, - ) - encoder_scheduler = EncoderScheduler( - encoder, - send_sockets, - max_batch_size=max_batch_size, - coalesce_same_turn=coalesce_same_turn, - ) - encoder_scheduler.start() - try: - yield - finally: - if encoder_scheduler is not None: - await encoder_scheduler.stop() - - -app = FastAPI(lifespan=_lifespan) - - -async def run_encoder( - server_args: ServerArgs, schedule_path, dist_init_method, rank: int -): - encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank) - while True: - request = await async_sock_recv(encoder.schedule_socket) - await _handle_encoder_worker_request(encoder, request) - - -async def _handle_encoder_worker_request(encoder: MMEncoder, request): - if isinstance(request, ProfileReq): - if request.req_type == ProfileReqType.START_PROFILE: - if encoder.profiler is None: - encoder.profiler = EncoderProfiler(encoder.rank) - encoder.profiler.start(request) - else: - encoder.profiler.stop() - elif isinstance(request, dict) and request.get("type") == "batch_encode": - await encoder.batch_encode( - request["requests"], - Modality.from_str(request["modality"]), - ) - elif ( - isinstance(request, dict) - and isinstance(request.get("req_id"), str) - and request["req_id"].startswith(HEALTH_CHECK_RID_PREFIX) - ): - await encoder.encode( - mm_items=request["mm_items"], - modality=Modality.from_str(request["modality"]), - req_id=request["req_id"], - num_parts=request["num_parts"], - part_idx=request["part_idx"], - hashes=request.get("hashes"), - ) - else: - await encoder.encode_request(request, Modality.from_str(request["modality"])) - - -def launch_encoder(server_args, schedule_path, dist_init_method, rank): - try: - asyncio.run(run_encoder(server_args, schedule_path, dist_init_method, rank)) - except KeyboardInterrupt: - logger.info(f"Exit rank {rank}") - except Exception: - traceback.print_exc() - - -def _register_encoder_url_with_bootstrap(server_args: ServerArgs): - """Asynchronously register this encoder with each bootstrap URL. - - Spawns a daemon thread that retries each URL independently with bounded - backoff. The encoder's own startup is not blocked: if some bootstrap - server is slow or unreachable, only the background worker waits. - - Inspired by ``_ensure_prefill_info`` in disaggregation/decode.py: each - target keeps its own retry count and is retried at a fixed interval - instead of serialising sleeps in a single thread. - """ - - host = server_args.host - if not host or host in ("0.0.0.0", "::"): - host = get_local_ip_auto(server_args.host) - scheme = "https" if server_args.ssl_certfile else "http" - encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) - payload = {"url": encoder_url} - bootstrap_urls = list(server_args.encoder_register_urls) - if not bootstrap_urls: - return - - max_retries = 30 - retry_interval = 5.0 - request_timeout = 5.0 - - def _try_register_once(bootstrap_url: str) -> bool: - try: - resp = http_requests.post( - f"{bootstrap_url}/register_encoder_url", - json=payload, - timeout=request_timeout, - ) - if resp.status_code == 200: - logger.info( - f"Registered encoder URL '{encoder_url}' with bootstrap " - f"at {bootstrap_url}" - ) - return True - logger.warning( - f"Bootstrap {bootstrap_url} returned {resp.status_code}: {resp.text}" - ) - except Exception as e: - logger.debug(f"Register attempt to {bootstrap_url} failed: {e}") - return False - - def _worker(): - pending = list(bootstrap_urls) - retry_count = {url: 0 for url in pending} - while pending: - still_pending = [] - for bootstrap_url in pending: - if _try_register_once(bootstrap_url): - continue - retry_count[bootstrap_url] += 1 - if retry_count[bootstrap_url] >= max_retries: - logger.error( - f"Giving up on bootstrap {bootstrap_url} after " - f"{max_retries} attempts. Encoder discovery via this " - f"bootstrap will be incomplete." - ) - continue - still_pending.append(bootstrap_url) - pending = still_pending - if pending: - time.sleep(retry_interval) - - threading.Thread( - target=_worker, daemon=True, name="encoder-bootstrap-register" - ).start() - - -def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs): - host = server_args.host - if not host or host in ("0.0.0.0", "::"): - host = get_local_ip_auto(server_args.host) - scheme = "https" if server_args.ssl_certfile else "http" - encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) - payload = {"url": encoder_url} - - for bootstrap_url in server_args.encoder_register_urls: - try: - resp = http_requests.delete( - f"{bootstrap_url}/unregister_encoder_url", - json=payload, - timeout=2.0, - ) - if resp.status_code == 200: - logger.info( - f"Unregistered encoder URL '{encoder_url}' from " - f"bootstrap at {bootstrap_url}" - ) - else: - logger.warning( - f"Bootstrap {bootstrap_url} returned " - f"{resp.status_code} on unregister: {resp.text}" - ) - except Exception as e: - logger.debug(f"Unregister from {bootstrap_url} failed: {e}") - - -def launch_server(server_args: ServerArgs): - configure_logger(server_args, prefix=" encode_server") - # Publish before the launch path reads configuration; the encoder built - # below re-projects the same object. - publish(server_args, role="encoder") - if get_parallel().dp_size > 1: - _launch_server_dp(server_args) - return - - global encoder, encoder_metrics_collector - - # Set up prometheus metrics. - if get_observability().enable_metrics: - set_prometheus_multiproc_dir() - labels = { - "model_name": get_serving().served_model_name, - "dp_rank": "0", - } - if get_observability().extra_metric_labels: - labels.update(get_observability().extra_metric_labels) - encoder_metrics_collector = EncoderMetricsCollector(labels) - add_prometheus_middleware(app) - - ctx = mp.get_context("spawn") - zmq_ctx = zmq.Context(10) - ipc_path_prefix = random_uuid() - port_args = PortArgs.init_new(server_args) - if get_parallel().dist_init_addr: - na = NetworkAddress.parse(get_parallel().dist_init_addr) - dist_init_method = na.to_tcp() - else: - dist_init_method = NetworkAddress( - get_serving().host or "127.0.0.1", port_args.nccl_port - ).to_tcp() - if get_observability().enable_trace: - process_tracing_init( - get_observability().otlp_traces_endpoint, - "sglang", - trace_modules=get_observability().trace_modules, - ) - trace_set_thread_info("Encoder") - for rank in range(1, configured_tp_size()): - schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" - send_sockets.append( - get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) - ) - ctx.Process( - target=launch_encoder, - args=(server_args, schedule_path, dist_init_method, rank), - daemon=True, - ).start() - encoder = MMEncoder(server_args, dist_init_method=dist_init_method) - - # Register this encoder's URL with prefill server(s) if configured. - if get_disagg().encoder_register_urls: - import atexit - - _register_encoder_url_with_bootstrap(server_args) - atexit.register(_unregister_encoder_url_from_bootstrap, server_args) - - uvicorn.run(app, host=get_serving().host, port=get_serving().port) - - -def _launch_server_dp(server_args: ServerArgs): - global dp_dispatcher - - if get_parallel().dp_size <= 1 or server_args.tp_size != 1: - raise ValueError( - "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " - f"dp_size={get_parallel().dp_size}, tp_size={server_args.tp_size}." - ) - dp_size = get_parallel().dp_size - logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") - - # DP mode: workers (subprocesses) write metrics to the shared multiproc dir; - # the main process exposes the aggregated /metrics endpoint. - if server_args.enable_metrics: - set_prometheus_multiproc_dir() - add_prometheus_middleware(app) - - ctx = mp.get_context("spawn") - ipc_prefix = random_uuid() - async_zmq_ctx = zmq.asyncio.Context(dp_size + 1) - - result_path = f"ipc:///tmp/{ipc_prefix}_dp_result" - result_socket = get_zmq_socket(async_zmq_ctx, zmq.PULL, result_path, True) - - dispatch_sockets: List[zmq.asyncio.Socket] = [ - get_zmq_socket( - async_zmq_ctx, zmq.PUSH, f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{r}", True - ) - for r in range(dp_size) - ] - - # Register atexit BEFORE spawn loop so partial spawns get reaped on - # exception (atexit holds the list ref and reads it at exit time). - import atexit - - worker_processes: List[mp.Process] = [] - - def _kill_workers(): - for p in worker_processes: - if p.is_alive(): - p.kill() - for p in worker_processes: - p.join(timeout=5) - - atexit.register(_kill_workers) - - for dp_rank in range(dp_size): - gpu_id = server_args.base_gpu_id + dp_rank - # Pin the device parent-side around spawn (same convention as the - # scheduler launcher and DP controller) so the child inherits - # CUDA_VISIBLE_DEVICES from its first instruction, before any import - # can enumerate CUDA. No-op unless SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS - # is set, in which case gpu_id is reindexed to 0 and CVD is pinned. - with maybe_reindex_device_id(gpu_id) as gpu_id: - proc = ctx.Process( - target=launch_dp_worker, - args=( - server_args, - dp_rank, - gpu_id, - f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}", - result_path, - ), - daemon=False, - ) - proc.start() - worker_processes.append(proc) - - labels = {"model_name": get_serving().served_model_name} - if server_args.extra_metric_labels: - labels.update(server_args.extra_metric_labels) - dp_dispatcher = DPDispatcher( - dp_size, - dispatch_sockets, - result_socket, - worker_processes, - enable_metrics=server_args.enable_metrics, - labels=labels, - ) - - # Register this encoder's URL with prefill server(s) if configured. - if server_args.encoder_register_urls: - import atexit - - _register_encoder_url_with_bootstrap(server_args) - atexit.register(_unregister_encoder_url_from_bootstrap, server_args) - - uvicorn.run(app, host=server_args.host, port=server_args.port) - - -def _summarise_dp_broadcast(results: List[dict]) -> Response: - # Treat missing/None content as failure so a stuck rank doesn't hide - # behind the others' "ok". Status = the most severe per-rank error code - # (5xx beats 4xx) rather than a blanket 400, so a worker's 500/503/504 - # isn't misreported as a client error. - msgs: List[str] = [] - error_codes: List[int] = [] - for r in results: - content = r.get("content") - if isinstance(content, dict): - msgs.append(content.get("msg", "")) - if not content.get("ok"): - # Worker ran but reported a logical failure; no transport code, - # so treat as a bad request (matches the non-DP profile path). - error_codes.append(int(r.get("_error_code") or HTTPStatus.BAD_REQUEST)) - else: - msgs.append(r.get("_error", "unknown error")) - error_codes.append( - int(r.get("_error_code") or HTTPStatus.INTERNAL_SERVER_ERROR) - ) - status_code = 200 if not error_codes else max(error_codes) - return Response( - content="\n".join(msgs) + "\n", - status_code=status_code, - ) - - -async def get_condition(rid): - async with cond_dict_lock: - if rid not in rid_to_cond: - rid_to_cond[rid] = asyncio.Condition() - return rid_to_cond[rid] - - -@app.post("/encode") -async def handle_encode_request(request: dict): - req_id = request["req_id"] - start_time = time.monotonic() - time_stats_json = request.pop("time_stats_json", None) - time_stats = EncoderReqTimeStats() - if dp_dispatcher is not None: - if time_stats_json: - request = dict(request) - request["time_stats_json"] = time_stats_json - try: - result = await dp_dispatcher.dispatch(request) - except MMError as e: - # Surface MMError.code (503 when all workers dead) instead of - # FastAPI's default 500. - logger.error(f"DP dispatch refused req_id={req_id}: {e}") - return ORJSONResponse( - status_code=int(e.code), - content={"status": "error", "message": str(e), "req_id": req_id}, - ) - if result.get("_error"): - error_type = result.get("_error_type", "") - # `or` (not `dict.get(key, default)`) so explicit None falls back too. - status_code = result.get("_error_code") or ( - HTTPStatus.BAD_REQUEST - if error_type == "ValueError" - else HTTPStatus.INTERNAL_SERVER_ERROR - ) - logger.error(f"DP worker error for req_id={req_id}: {result['_error']}") - return ORJSONResponse( - status_code=status_code, - content={ - "status": "error", - "message": result["_error"], - "req_id": req_id, - }, - ) - elapsed = time.monotonic() - start_time - logger.info( - f"[{req_id}] /encode completed in {elapsed:.3f}s, " - f"modality={request.get('modality', 'image')}" - ) - return ORJSONResponse(content=result.get("content")) - - modality_str = str(request.get("modality", "image")).lower() - 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 get_disagg().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 time_stats_json: - time_stats.decode_json(time_stats_json) - - modality_str = modality.name.lower() - time_stats.modality = modality_str - time_stats.set_metrics_collector(encoder_metrics_collector) - time_stats.set_mm_encode_start_time() - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_received(modality=modality_str) - 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: - time_stats.trace_ctx.abort( - abort_info={"reason": "encoder batch timed out"} - ) - return ORJSONResponse( - status_code=HTTPStatus.GATEWAY_TIMEOUT, - content={ - "status": "error", - "message": "encoder batch timed out", - "req_id": req_id, - }, - ) - else: - # Non-batched requests still own their collective dispatch order - # directly; batched requests take this lock in _dispatch_group. - # Locking direct dispatch together with the rank0 await keeps its - # NCCL launch order matching the ZMQ dispatch order rank>0 sees. - async with encoder.encode_dispatch_lock: - for socket in send_sockets: - sock_send(socket, wrap_as_pickle(request)) - nbytes, embedding_len, embedding_dim, error_msg, error_code = ( - await encoder.encode_request(request, modality) - ) - - if error_msg: - time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) - else: - time_stats.set_mm_encode_end_time() - - if error_msg: - if get_disagg().encoder_transfer_backend == "zmq_to_scheduler": - if request["embedding_port"] is None: - start_background_send(req_id) - else: - for port in request["embedding_port"]: - await encoder.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=port, - ) - # Signal waiters on failure for mooncake - if get_disagg().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) - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="error" - ) - return ORJSONResponse( - status_code=error_code, - content={"status": "error", "message": error_msg, "req_id": req_id}, - ) - if get_disagg().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( - { - "embedding_size": nbytes, - "embedding_len": embedding_len, - "embedding_dim": embedding_dim, - } - ) - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="success" - ) - return ORJSONResponse(content=request) - elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler": - logger.info(f"{request['embedding_port'] = }") - if request["embedding_port"] is None: - await encoder.send_with_url( - req_id=request["req_id"], - ) - else: - assert type(request["embedding_port"]) == list - tasks = [] - for embedding_port in request["embedding_port"]: - tasks.append( - encoder.send( - req_id=request["req_id"], - prefill_host=request["prefill_host"], - embedding_port=embedding_port, - ) - ) - await asyncio.gather(*tasks) - encoder.embedding_to_send.pop(request["req_id"], None) - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="success" - ) - return ORJSONResponse(content=None) - elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer": - await encoder.send( - req_id=request["req_id"], - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - ) - encoder.embedding_to_send.pop(request["req_id"], None) - elapsed = time.monotonic() - start_time - logger.info( - f"[{req_id}] /encode completed in {elapsed:.3f}s, " - f"modality={request['modality']}, tokens={embedding_len}" - ) - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="success" - ) - return ORJSONResponse(content=None) - except Exception as e: - 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 get_disagg().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) - if encoder_metrics_collector is not None: - encoder_metrics_collector.inc_requests_total( - modality=modality_str, status="error" - ) - return ORJSONResponse( - status_code=HTTPStatus.INTERNAL_SERVER_ERROR, - content={ - "status": "error", - "message": error_msg, - "req_id": req_id, - }, - ) - - -@app.post("/send") -async def handle_send_request(request: dict): - # mooncake backend - if dp_dispatcher is not None: - try: - result = await dp_dispatcher.dispatch_send(request) - except MMError as e: - req_id = request.get("req_id", "?") - logger.error(f"DP dispatch_send refused req_id={req_id}: {e}") - return Response( - content=f"Encoder DP worker send error: {e}", - status_code=int(e.code), - ) - if result.get("_error"): - req_id = request.get("req_id", "?") - status_code = result.get("_error_code") or int( - HTTPStatus.INTERNAL_SERVER_ERROR - ) - logger.error( - f"DP worker send error for req_id={req_id}: {result['_error']}" - ) - return Response( - content=f"Encoder DP worker send error: {result['_error']}", - status_code=status_code, - ) - return ORJSONResponse(content=result.get("content")) - await encoder.send( - req_id=request["req_id"], - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - session_id=request["session_id"], - buffer_address=request["buffer_address"], - ) - req_id = request["req_id"] - # Keep embedding until all ranks have /send'd; release early when receive_count is met. - expected_sends = request.get("receive_count") - if expected_sends: - done = mooncake_send_done_count.get(req_id, 0) + 1 - if done >= expected_sends: - mooncake_send_done_count.pop(req_id, None) - await encoder._cleanup_inflight_encode_state(req_id) - else: - mooncake_send_done_count[req_id] = done - return ORJSONResponse(content=None) - - -@app.post("/scheduler_receive_url") -async def handle_scheduler_receive_url_request(request: dict): - rid = request["req_id"] - async with rid_lock: - global rid_to_receive_endpoint - if rid not in rid_to_receive_endpoint: - rid_to_receive_endpoint[rid] = set() - rid_to_receive_count[rid] = request["receive_count"] - assert rid_to_receive_count[rid] == request["receive_count"] - rid_to_receive_endpoint[rid].add(request["receive_url"]) - cond = await get_condition(rid) - async with cond: - cond.notify_all() - - -@app.get("/health") -@app.get("/health_generate") -async def health_generate(): - """ - Health check endpoint for the encoder server. - Performs a dummy encode to verify the encoder is functional. - Returns 200 if the encoder is healthy, 503 otherwise. - """ - if dp_dispatcher is not None: - # Strict: any dead (exited) rank fails health → orchestrator restarts. - if not dp_dispatcher.all_ranks_alive: - return Response(status_code=503) - # Process-liveness (proc.sentinel) can't see a worker that's alive but - # wedged (hung GPU / NCCL deadlock / stalled ZMQ). Probe every rank with - # a tiny dummy encode; each worker runs it only when idle and otherwise - # reports healthy at once, keeping the probe off the GPU under load. - try: - results = await dp_dispatcher.broadcast( - {"_dp_type": "health_encode"}, - timeout=HEALTH_CHECK_TIMEOUT, - ) - except MMError: - return Response(status_code=503) - if any(r.get("_error") for r in results): - return Response(status_code=503) - return Response(status_code=200) - if encoder is None: - return Response(status_code=503) - - # Pick the first available modality for the dummy encode - if encoder.image_processor is not None: - mm_items = [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"] - modality = Modality.IMAGE - elif encoder.audio_processor is not None: - mm_items = [f"data:audio/wav;base64,{MINIMUM_WAV_SILENCE_BASE64}"] - modality = Modality.AUDIO - else: - # No processor available, fall back to liveness check only - return Response(status_code=200) - - try: - # uuid keeps rids unique across workers; a bare time.time() can collide. - req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" - - dummy_request = { - "mm_items": mm_items, - "modality": modality.name, - "req_id": req_id, - "num_parts": 1, - "part_idx": 0, - } - - # A health encode participates in the same TP collectives as a real - # request. Serialize its broadcast and rank-0 forward with every other - # collective dispatch, then recheck whether traffic made the probe - # unnecessary while it waited for the lock. - async with encoder.encode_dispatch_lock: - if encoder.embedding_to_send: - return Response(status_code=200) - for socket in send_sockets: - sock_send(socket, wrap_as_pickle(dummy_request)) - - _, _, _, error_msg, _ = await asyncio.wait_for( - encoder.encode( - mm_items=mm_items, - modality=modality, - req_id=req_id, - num_parts=1, - part_idx=0, - ), - timeout=HEALTH_CHECK_TIMEOUT, - ) - - # Clean up stored embedding - encoder.embedding_to_send.pop(req_id, None) - - if error_msg: - logger.error(f"Encoder health check failed: {error_msg}") - return Response(status_code=503) - - return Response(status_code=200) - - except asyncio.TimeoutError: - logger.error(f"Encoder health check timed out after {HEALTH_CHECK_TIMEOUT}s") - return Response(status_code=503) - except Exception as e: - logger.error(f"Encoder health check failed: {e}") - return Response(status_code=503) - - -@app.api_route("/start_profile", methods=["GET", "POST"]) -async def start_profile_async(obj: Annotated[Optional[ProfileReq], Body()] = None): - if dp_dispatcher is not None: - if obj is not None: - obj.req_type = ProfileReqType.START_PROFILE - try: - results = await dp_dispatcher.broadcast( - {"_dp_type": "start_profile", "profile_req": obj} - ) - except MMError as e: - return Response(content=f"{e}\n", status_code=int(e.code)) - return _summarise_dp_broadcast(results) - if encoder is None: - return Response(content="encoder not ready\n", status_code=503) - req = obj or ProfileReq() - req.req_type = ProfileReqType.START_PROFILE - for socket in send_sockets: - sock_send(socket, req) - if encoder.profiler is None: - encoder.profiler = EncoderProfiler(encoder.rank) - ok, msg = encoder.profiler.start(req) - if ok: - detail = ( - f"Start profiling. output_dir={encoder.profiler.output_dir} " - f"profile_id={encoder.profiler.profile_id}\n" - ) - return Response(content=detail, status_code=200) - return Response( - content=(msg or "Start profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST - ) - - -@app.api_route("/stop_profile", methods=["GET", "POST"]) -async def stop_profile_async(): - if dp_dispatcher is not None: - try: - results = await dp_dispatcher.broadcast({"_dp_type": "stop_profile"}) - except MMError as e: - return Response(content=f"{e}\n", status_code=int(e.code)) - return _summarise_dp_broadcast(results) - if encoder is None: - return Response(content="encoder not ready\n", status_code=503) - if encoder.profiler is None: - return Response( - content="profiling not initialized\n", status_code=HTTPStatus.BAD_REQUEST - ) - req = ProfileReq(req_type=ProfileReqType.STOP_PROFILE) - for socket in send_sockets: - sock_send(socket, req) - ok, msg = encoder.profiler.stop() - if ok: - return Response(content="Stop profiling.\n", status_code=200) - return Response( - content=(msg or "Stop profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST - ) diff --git a/python/sglang/srt/disaggregation/encoder/__init__.py b/python/sglang/srt/disaggregation/encoder/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/disaggregation/encode_grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py similarity index 90% rename from python/sglang/srt/disaggregation/encode_grpc_server.py rename to python/sglang/srt/disaggregation/encoder/grpc_server.py index 3516f56d4..163dc276d 100644 --- a/python/sglang/srt/disaggregation/encode_grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -21,11 +21,7 @@ from grpc_health.v1 import health_pb2, health_pb2_grpc from grpc_reflection.v1alpha import reflection from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc -from sglang.srt.disaggregation.encode_server import ( - MMEncoder, - handle_scheduler_receive_url_request, - launch_encoder, -) +from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality from sglang.srt.runtime_context import get_disagg @@ -92,6 +88,7 @@ class SGLangEncoderServer(SGLangEncoderServicer): try: request_dict = { "mm_items": list(request.mm_items), + "modality": Modality.IMAGE.name, "req_id": request.req_id, "num_parts": request.num_parts, "part_idx": request.part_idx, @@ -99,21 +96,17 @@ class SGLangEncoderServer(SGLangEncoderServicer): for socket in self.send_sockets: await async_sock_send(socket, wrap_as_pickle(request_dict)) - # gRPC encode is image-only; encoder.encode() requires modality + # gRPC encode is image-only; the request follows the configured + # cache and transfer backend. ( nbytes, embedding_len, embedding_dim, error_msg, error_code, - ) = await self.encoder.encode( - mm_items=list(request.mm_items), - modality=Modality.IMAGE, - req_id=request.req_id, - num_parts=request.num_parts, - part_idx=request.part_idx, - ) + ) = await self.encoder.encode_request(request_dict, Modality.IMAGE) if error_msg is not None: + await self.encoder.release_request(request.req_id) context.set_code(grpc.StatusCode.INTERNAL) context.set_details(error_msg) return sglang_encoder_pb2.EncodeResponse() @@ -140,7 +133,7 @@ class SGLangEncoderServer(SGLangEncoderServicer): ) ) await asyncio.gather(*tasks) - self.encoder.embedding_to_send.pop(request.req_id, None) + await self.encoder.release_request(request.req_id) return sglang_encoder_pb2.EncodeResponse() elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer": embedding_port = ( @@ -151,7 +144,7 @@ class SGLangEncoderServer(SGLangEncoderServicer): prefill_host=request.prefill_host, embedding_port=embedding_port, ) - self.encoder.embedding_to_send.pop(request.req_id, None) + await self.encoder.release_request(request.req_id) return sglang_encoder_pb2.EncodeResponse() return sglang_encoder_pb2.EncodeResponse() @@ -159,6 +152,7 @@ class SGLangEncoderServer(SGLangEncoderServicer): except Exception as e: logger.error(f"Encode error: {e}") traceback.print_exc() + await self.encoder.release_request(request.req_id) context.set_code(grpc.StatusCode.INTERNAL) context.set_details(str(e)) return sglang_encoder_pb2.EncodeResponse() @@ -176,12 +170,13 @@ class SGLangEncoderServer(SGLangEncoderServicer): request.buffer_address if request.buffer_address else None ), ) - self.encoder.embedding_to_send.pop(request.req_id, None) + await self.encoder.release_request(request.req_id) return sglang_encoder_pb2.SendResponse() except Exception as e: logger.error(f"Send error: {e}") traceback.print_exc() + await self.encoder.release_request(request.req_id) context.set_code(grpc.StatusCode.INTERNAL) context.set_details(str(e)) return sglang_encoder_pb2.SendResponse() @@ -190,12 +185,10 @@ class SGLangEncoderServer(SGLangEncoderServicer): self, request: sglang_encoder_pb2.SchedulerReceiveUrlRequest, context ) -> sglang_encoder_pb2.SchedulerReceiveUrlResponse: try: - await handle_scheduler_receive_url_request( - { - "req_id": request.req_id, - "receive_count": request.receive_count, - "receive_url": request.receive_url, - } + await self.encoder.register_embedding_destinations( + request.req_id, + request.receive_count, + [request.receive_url], ) return sglang_encoder_pb2.SchedulerReceiveUrlResponse() diff --git a/python/sglang/srt/disaggregation/encoder/http_server.py b/python/sglang/srt/disaggregation/encoder/http_server.py new file mode 100644 index 000000000..d87a020d7 --- /dev/null +++ b/python/sglang/srt/disaggregation/encoder/http_server.py @@ -0,0 +1,653 @@ +"""HTTP API layer for the EPD encoder server. + +This module is designed to be replaceable by a Rust implementation. +It contains the FastAPI application, HTTP route handlers, HTTP lifecycle, and +response conversion. Backend scheduling and process management are provided by +the protocol-neutral :mod:`runtime` module. + +GPU tensor operations remain in :mod:`server.MMEncoder`. +""" + +import asyncio +import contextlib +import logging +import threading +import time +import uuid +from http import HTTPStatus +from typing import Annotated, List, Optional + +import requests as http_requests +import uvicorn +import zmq +from fastapi import Body, FastAPI +from fastapi.responses import ORJSONResponse, Response + +import sglang.srt.disaggregation.encoder.server as server_module +from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX +from sglang.srt.disaggregation.encoder.runtime import ( + DPDispatcher, + EncoderRuntime, + EncoderScheduler, + execute_encode_pipeline, + launch_dp_runtime, + launch_local_runtime, +) +from sglang.srt.disaggregation.encoder.server import ( + EncoderProfiler, + MMEncoder, + MMError, +) +from sglang.srt.managers.io_struct import ( + ProfileReq, + ProfileReqType, + sock_send, + wrap_as_pickle, +) +from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.runtime_context import ( + get_disagg, + get_observability, + get_parallel, + get_serving, + publish, +) +from sglang.srt.server_args import ServerArgs +from sglang.srt.utils import ( + add_prometheus_middleware, + configure_logger, +) +from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto + +logger = logging.getLogger(__name__) + +HEALTH_CHECK_TIMEOUT = 30 + +# Minimal 32x32 black PNG for health check dummy encode +MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" + +# Minimal WAV: 16kHz mono 16-bit PCM, 160 samples (0.01s) of silence +MINIMUM_WAV_SILENCE_BASE64 = "UklGRmQBAABXQVZFZm10IBAAAAABAAEAgD4AAAB9AAACABAAZGF0YUABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==" + +encoder: Optional[MMEncoder] = None +send_sockets: List[zmq.Socket] = [] +encoder_scheduler: Optional[EncoderScheduler] = None +local_runtime: Optional[EncoderRuntime] = None + +# DP mode (--dp-size > 1): the protocol-neutral runtime owns worker processes +# and ZMQ; HTTP only keeps the dispatcher handle used by route handlers. +dp_dispatcher: Optional["DPDispatcher"] = None + + +def is_health_check_request(rid: Optional[str]) -> bool: + return isinstance(rid, str) and rid.startswith(HEALTH_CHECK_RID_PREFIX) + + +@contextlib.asynccontextmanager +async def _lifespan(app: FastAPI): + if dp_dispatcher is not None: + dp_dispatcher.start() + yield + return + if local_runtime is not None: + local_runtime.start() + try: + yield + finally: + if local_runtime is not None: + await local_runtime.stop() + + +app = FastAPI(lifespan=_lifespan) + + +def _register_encoder_url_with_bootstrap(server_args: ServerArgs): + """Asynchronously register this encoder with each bootstrap URL. + + Spawns a daemon thread that retries each URL independently with bounded + backoff. The encoder's own startup is not blocked: if some bootstrap + server is slow or unreachable, only the background worker waits. + + Inspired by ``_ensure_prefill_info`` in disaggregation/decode.py: each + target keeps its own retry count and is retried at a fixed interval + instead of serialising sleeps in a single thread. + """ + + host = server_args.host + if not host or host in ("0.0.0.0", "::"): + host = get_local_ip_auto(server_args.host) + scheme = "https" if server_args.ssl_certfile else "http" + encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) + payload = {"url": encoder_url} + bootstrap_urls = list(server_args.encoder_register_urls) + if not bootstrap_urls: + return + + max_retries = 30 + retry_interval = 5.0 + request_timeout = 5.0 + + def _try_register_once(bootstrap_url: str) -> bool: + try: + resp = http_requests.post( + f"{bootstrap_url}/register_encoder_url", + json=payload, + timeout=request_timeout, + ) + if resp.status_code == 200: + logger.info( + f"Registered encoder URL '{encoder_url}' with bootstrap " + f"at {bootstrap_url}" + ) + return True + logger.warning( + f"Bootstrap {bootstrap_url} returned {resp.status_code}: {resp.text}" + ) + except Exception as e: + logger.debug(f"Register attempt to {bootstrap_url} failed: {e}") + return False + + def _worker(): + pending = list(bootstrap_urls) + retry_count = {url: 0 for url in pending} + while pending: + still_pending = [] + for bootstrap_url in pending: + if _try_register_once(bootstrap_url): + continue + retry_count[bootstrap_url] += 1 + if retry_count[bootstrap_url] >= max_retries: + logger.error( + f"Giving up on bootstrap {bootstrap_url} after " + f"{max_retries} attempts. Encoder discovery via this " + f"bootstrap will be incomplete." + ) + continue + still_pending.append(bootstrap_url) + pending = still_pending + if pending: + time.sleep(retry_interval) + + threading.Thread( + target=_worker, daemon=True, name="encoder-bootstrap-register" + ).start() + + +def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs): + host = server_args.host + if not host or host in ("0.0.0.0", "::"): + host = get_local_ip_auto(server_args.host) + scheme = "https" if server_args.ssl_certfile else "http" + encoder_url = NetworkAddress(host, server_args.port).to_url(scheme) + payload = {"url": encoder_url} + + for bootstrap_url in server_args.encoder_register_urls: + try: + resp = http_requests.delete( + f"{bootstrap_url}/unregister_encoder_url", + json=payload, + timeout=2.0, + ) + if resp.status_code == 200: + logger.info( + f"Unregistered encoder URL '{encoder_url}' from " + f"bootstrap at {bootstrap_url}" + ) + else: + logger.warning( + f"Bootstrap {bootstrap_url} returned " + f"{resp.status_code} on unregister: {resp.text}" + ) + except Exception as e: + logger.debug(f"Unregister from {bootstrap_url} failed: {e}") + + +def launch_server(server_args: ServerArgs): + global dp_dispatcher, encoder, encoder_scheduler, local_runtime, send_sockets + + configure_logger(server_args, prefix=" encode_server") + # Publish before the launch path reads configuration; each encoder built + # below re-projects the same object in its process. + publish(server_args, role="encoder") + if get_parallel().dp_size > 1: + dp_dispatcher = launch_dp_runtime(server_args) + # runtime initializes multiprocess metrics before spawning; + # HTTP only exposes their endpoint. + if get_observability().enable_metrics: + add_prometheus_middleware(app) + else: + local_runtime = launch_local_runtime(server_args) + # Compatibility aliases for the existing HTTP request path. Runtime is + # now the sole constructor and lifecycle owner of these objects. + encoder = local_runtime.encoder + encoder_scheduler = local_runtime.scheduler + send_sockets = local_runtime.send_sockets + if get_observability().enable_metrics: + add_prometheus_middleware(app) + + # Register this encoder's URL with prefill server(s) if configured. + if get_disagg().encoder_register_urls: + import atexit + + _register_encoder_url_with_bootstrap(server_args) + atexit.register(_unregister_encoder_url_from_bootstrap, server_args) + + uvicorn.run(app, host=get_serving().host, port=get_serving().port) + + +def _summarise_dp_broadcast(results: List[dict]) -> Response: + # Treat missing/None content as failure so a stuck rank doesn't hide + # behind the others' "ok". Status = the most severe per-rank error code + # (5xx beats 4xx) rather than a blanket 400, so a worker's 500/503/504 + # isn't misreported as a client error. + msgs: List[str] = [] + error_codes: List[int] = [] + for r in results: + content = r.get("content") + if isinstance(content, dict): + msgs.append(content.get("msg", "")) + if not content.get("ok"): + # Worker ran but reported a logical failure; no transport code, + # so treat as a bad request (matches the non-DP profile path). + error_codes.append(int(r.get("_error_code") or HTTPStatus.BAD_REQUEST)) + else: + msgs.append(r.get("_error", "unknown error")) + error_codes.append( + int(r.get("_error_code") or HTTPStatus.INTERNAL_SERVER_ERROR) + ) + status_code = 200 if not error_codes else max(error_codes) + return Response( + content="\n".join(msgs) + "\n", + status_code=status_code, + ) + + +@app.post("/encode") +async def handle_encode_request(request: dict): + req_id = request["req_id"] + start_time = time.monotonic() + time_stats_json = request.pop("time_stats_json", None) + if dp_dispatcher is not None: + if time_stats_json: + request = dict(request) + request["time_stats_json"] = time_stats_json + try: + result = await dp_dispatcher.dispatch(request) + except MMError as e: + # Surface MMError.code (503 when all workers dead) instead of + # FastAPI's default 500. + logger.error(f"DP dispatch refused req_id={req_id}: {e}") + return ORJSONResponse( + status_code=int(e.code), + content={"status": "error", "message": str(e), "req_id": req_id}, + ) + if result.get("_error"): + error_type = result.get("_error_type", "") + # `or` (not `dict.get(key, default)`) so explicit None falls back too. + status_code = result.get("_error_code") or ( + HTTPStatus.BAD_REQUEST + if error_type == "ValueError" + else HTTPStatus.INTERNAL_SERVER_ERROR + ) + logger.error(f"DP worker error for req_id={req_id}: {result['_error']}") + return ORJSONResponse( + status_code=status_code, + content={ + "status": "error", + "message": result["_error"], + "req_id": req_id, + }, + ) + elapsed = time.monotonic() - start_time + logger.info( + f"[{req_id}] /encode completed in {elapsed:.3f}s, " + f"modality={request.get('modality', 'image')}" + ) + content = result.get("content") + return ORJSONResponse(content=content) + + try: + if time_stats_json: + request["time_stats_json"] = time_stats_json + content = await execute_encode_pipeline( + encoder, + encoder_scheduler, + request, + send_sockets=send_sockets, + ) + elapsed = time.monotonic() - start_time + logger.info( + f"[{req_id}] /encode completed in {elapsed:.3f}s, " + f"modality={request.get('modality', 'image')}" + ) + return ORJSONResponse(content=content) + except asyncio.TimeoutError: + return ORJSONResponse( + status_code=HTTPStatus.GATEWAY_TIMEOUT, + content={ + "status": "error", + "message": "encoder batch timed out", + "req_id": req_id, + }, + ) + except MMError as e: + return ORJSONResponse( + status_code=int(e.code), + content={"status": "error", "message": str(e), "req_id": req_id}, + ) + except Exception as e: + error_msg = str(e) + logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}") + return ORJSONResponse( + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + content={ + "status": "error", + "message": error_msg, + "req_id": req_id, + }, + ) + + +@app.post("/send") +async def handle_send_request(request: dict): + """Mooncake-only: drive the RDMA push of a staged embedding. The zmq + backends deliver embeddings inline during /encode and never call /send.""" + req_id = request["req_id"] + receive_count = request.get("receive_count") + if dp_dispatcher is not None: + try: + result = await dp_dispatcher.dispatch_send(request) + except MMError as e: + logger.error(f"DP dispatch_send refused req_id={req_id}: {e}") + return Response( + content=f"Encoder DP worker send error: {e}", + status_code=int(e.code), + ) + if result.get("_error"): + status_code = result.get("_error_code") or int( + HTTPStatus.INTERNAL_SERVER_ERROR + ) + logger.error( + f"DP worker send error for req_id={req_id}: {result['_error']}" + ) + return Response( + content=f"Encoder DP worker send error: {result['_error']}", + status_code=status_code, + ) + return ORJSONResponse(content=result.get("content")) + sent = await encoder.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=request["embedding_port"], + session_id=request["session_id"], + buffer_address=request["buffer_address"], + ) + if not sent: + # No transfer happened: fail fast rather than 200 + a phantom count. + return ORJSONResponse( + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + content={ + "status": "error", + "message": f"no staged embedding for req_id={req_id} (already released)", + "req_id": req_id, + }, + ) + # Sibling ranks share this embedding, so free it only once all have sent. + # No count means a pre-refcount decoder: leave it to the sweep, as when + # some rank never sends at all. + if receive_count: + await server_module.meta_registry.note_send_done(req_id, receive_count) + return ORJSONResponse(content=None) + + +@app.post("/scheduler_receive_meta_data") +async def handle_scheduler_receive_meta_data(request: dict): + """Decoder pull endpoint for the per-part encode metadata. Blocks until the + encode publishes its sizes, so a pull that beats the encode simply waits.""" + req_id = request["req_id"] + if dp_dispatcher is not None: + try: + result = await dp_dispatcher.dispatch_wait_metadata(request) + except MMError as e: + return ORJSONResponse( + status_code=int(e.code), + content={"status": "error", "message": str(e), "req_id": req_id}, + ) + if result.get("_error"): + return ORJSONResponse( + status_code=result.get("_error_code") + or int(HTTPStatus.INTERNAL_SERVER_ERROR), + content={ + "status": "error", + "message": result["_error"], + "req_id": req_id, + }, + ) + meta = result.get("content") + else: + try: + meta = await server_module.meta_registry.wait(req_id) + except asyncio.TimeoutError: + logger.error(f"[{req_id}] /scheduler_receive_meta_data timed out") + return ORJSONResponse( + status_code=HTTPStatus.GATEWAY_TIMEOUT, + content={ + "status": "error", + "message": "encode metadata not ready", + "req_id": req_id, + }, + ) + if meta is None or meta.get("error") is not None: + message = meta["error"] if meta else "encode metadata missing" + return ORJSONResponse( + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + content={"status": "error", "message": message, "req_id": req_id}, + ) + return ORJSONResponse( + content={ + "req_id": req_id, + "part_idx": request["part_idx"], + "embedding_size": meta["embedding_size"], + "embedding_len": meta["embedding_len"], + "embedding_dim": meta["embedding_dim"], + } + ) + + +@app.post("/scheduler_receive_url") +async def handle_scheduler_receive_url_request(request: dict): + if dp_dispatcher is not None: + try: + result = await dp_dispatcher.dispatch_register_destinations(request) + except MMError as e: + return ORJSONResponse( + status_code=int(e.code), + content={ + "status": "error", + "message": str(e), + "req_id": request["req_id"], + }, + ) + if result.get("_error"): + return ORJSONResponse( + status_code=result.get("_error_code") + or int(HTTPStatus.INTERNAL_SERVER_ERROR), + content={ + "status": "error", + "message": result["_error"], + "req_id": request["req_id"], + }, + ) + return ORJSONResponse(content=None) + if encoder is None: + return ORJSONResponse( + status_code=HTTPStatus.SERVICE_UNAVAILABLE, + content={ + "status": "error", + "message": "encoder not ready", + "req_id": request["req_id"], + }, + ) + try: + await encoder.register_embedding_destinations( + request["req_id"], + request["receive_count"], + [request["receive_url"]], + ) + except MMError as e: + return ORJSONResponse( + status_code=int(e.code), + content={ + "status": "error", + "message": str(e), + "req_id": request["req_id"], + }, + ) + return ORJSONResponse(content=None) + + +@app.get("/health") +@app.get("/health_generate") +async def health_generate(): + """ + Health check endpoint for the encoder server. + Performs a dummy encode to verify the encoder is functional. + Returns 200 if the encoder is healthy, 503 otherwise. + """ + if dp_dispatcher is not None: + # Strict: any dead (exited) rank fails health → orchestrator restarts. + if not dp_dispatcher.all_ranks_alive: + return Response(status_code=503) + # Process-liveness (proc.sentinel) can't see a worker that's alive but + # wedged (hung GPU / NCCL deadlock / stalled ZMQ). Probe every rank with + # a tiny dummy encode; each worker runs it only when idle and otherwise + # reports healthy at once, keeping the probe off the GPU under load. + try: + results = await dp_dispatcher.broadcast( + {"_dp_type": "health_encode"}, + timeout=HEALTH_CHECK_TIMEOUT, + ) + except MMError: + return Response(status_code=503) + if any(r.get("_error") for r in results): + return Response(status_code=503) + return Response(status_code=200) + if encoder is None: + return Response(status_code=503) + + # Pick the first available modality for the dummy encode + if encoder.supports_modality(Modality.IMAGE): + mm_items = [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"] + modality = Modality.IMAGE + elif encoder.supports_modality(Modality.AUDIO): + mm_items = [f"data:audio/wav;base64,{MINIMUM_WAV_SILENCE_BASE64}"] + modality = Modality.AUDIO + else: + # No processor available, fall back to liveness check only + return Response(status_code=200) + + try: + # uuid keeps rids unique across workers; a bare time.time() can collide. + req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" + + dummy_request = { + "mm_items": mm_items, + "modality": modality.name, + "req_id": req_id, + "num_parts": 1, + "part_idx": 0, + } + + # A health encode participates in the same TP collectives as a real + # request. Serialize its broadcast and rank-0 forward with every other + # collective dispatch, then recheck whether traffic made the probe + # unnecessary while it waited for the lock. + async with encoder.encode_dispatch_lock: + if encoder.has_pending_embeddings(): + return Response(status_code=200) + for socket in send_sockets: + sock_send(socket, wrap_as_pickle(dummy_request)) + + _, _, _, error_msg, _ = await asyncio.wait_for( + encoder.encode( + mm_items=mm_items, + modality=modality, + req_id=req_id, + num_parts=1, + part_idx=0, + ), + timeout=HEALTH_CHECK_TIMEOUT, + ) + + # Clean up stored embedding + await encoder.release_request(req_id) + + if error_msg: + logger.error(f"Encoder health check failed: {error_msg}") + return Response(status_code=503) + + return Response(status_code=200) + + except asyncio.TimeoutError: + logger.error(f"Encoder health check timed out after {HEALTH_CHECK_TIMEOUT}s") + return Response(status_code=503) + except Exception as e: + logger.error(f"Encoder health check failed: {e}") + return Response(status_code=503) + + +@app.api_route("/start_profile", methods=["GET", "POST"]) +async def start_profile_async(obj: Annotated[Optional[ProfileReq], Body()] = None): + if dp_dispatcher is not None: + if obj is not None: + obj.req_type = ProfileReqType.START_PROFILE + try: + results = await dp_dispatcher.broadcast( + {"_dp_type": "start_profile", "profile_req": obj} + ) + except MMError as e: + return Response(content=f"{e}\n", status_code=int(e.code)) + return _summarise_dp_broadcast(results) + if encoder is None: + return Response(content="encoder not ready\n", status_code=503) + req = obj or ProfileReq() + req.req_type = ProfileReqType.START_PROFILE + for socket in send_sockets: + sock_send(socket, req) + if encoder.profiler is None: + encoder.profiler = EncoderProfiler(encoder.rank) + ok, msg = encoder.profiler.start(req) + if ok: + detail = ( + f"Start profiling. output_dir={encoder.profiler.output_dir} " + f"profile_id={encoder.profiler.profile_id}\n" + ) + return Response(content=detail, status_code=200) + return Response( + content=(msg or "Start profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST + ) + + +@app.api_route("/stop_profile", methods=["GET", "POST"]) +async def stop_profile_async(): + if dp_dispatcher is not None: + try: + results = await dp_dispatcher.broadcast({"_dp_type": "stop_profile"}) + except MMError as e: + return Response(content=f"{e}\n", status_code=int(e.code)) + return _summarise_dp_broadcast(results) + if encoder is None: + return Response(content="encoder not ready\n", status_code=503) + if encoder.profiler is None: + return Response( + content="profiling not initialized\n", status_code=HTTPStatus.BAD_REQUEST + ) + req = ProfileReq(req_type=ProfileReqType.STOP_PROFILE) + for socket in send_sockets: + sock_send(socket, req) + ok, msg = encoder.profiler.stop() + if ok: + return Response(content="Stop profiling.\n", status_code=200) + return Response( + content=(msg or "Stop profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST + ) diff --git a/python/sglang/srt/disaggregation/encoder/preprocessor.py b/python/sglang/srt/disaggregation/encoder/preprocessor.py new file mode 100644 index 000000000..ca5d1f62b --- /dev/null +++ b/python/sglang/srt/disaggregation/encoder/preprocessor.py @@ -0,0 +1,827 @@ +"""CPU-bound multimodal preprocessing for the EPD encoder. + +This module is designed to be replaceable by a Rust implementation. +It handles all CPU-bound work: media I/O (image/video/audio loading), +HF processor calls, config validation, and related helper computations. +GPU tensor operations remain in :mod:`server.MMEncoder`. +""" + +import asyncio +import concurrent.futures +import functools +import logging +import os +from dataclasses import dataclass +from typing import Callable, List, Optional, Tuple, Union + +import numpy as np +import torch +from transformers import AutoProcessor + +from sglang.srt.configs.model_config import ModelConfig +from sglang.srt.environ import envs +from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media +from sglang.srt.multimodal.encoder_preprocessing import ( + EncoderMediaProcessorConfig, + EncoderPreprocessOutput, + invoke_encoder_preprocessor, +) +from sglang.srt.multimodal.processors.qwen_vl import preprocess_video +from sglang.srt.runtime_context import ( + get_device, + get_mm, + get_model, + get_parallel, + get_serving, +) +from sglang.srt.server_args import ServerArgs +from sglang.srt.utils import ( + CLIENT_MEDIA_EXCEPTIONS, + load_audio, + load_image, + load_video, +) +from sglang.srt.utils.hf_transformers_utils import resolve_image_processor_backend + +logger = logging.getLogger(__name__) + + +_mm_grid_attrs = { + # Kimi K2.5/K3 HF processors use grid_thws (see base_processor.ATTR_NAME_TO_MODALITY). + Modality.IMAGE: ("image_grid_thw", "image_grid_hws", "grid_thws"), + Modality.VIDEO: ("video_grid_thw",), + Modality.AUDIO: ("audio_feature_lens_raw",), +} + + +def _convert(data): + if isinstance(data, torch.Tensor): + return data + elif isinstance(data, np.ndarray): + return torch.tensor(data) + elif isinstance(data, list) and isinstance(data[0], np.ndarray): + return torch.tensor(np.array(data)) + elif isinstance(data, list) and isinstance(data[0], (int, float)): + return torch.tensor(data) + else: + return data + + +def _get_original_image_size(image): + """Return an image's original (width, height) before encoder preprocessing.""" + if isinstance(image, dict): + image = image.get("image") + if isinstance(image, torch.Tensor): + if image.ndim < 2: + raise ValueError(f"Invalid image tensor shape: {tuple(image.shape)}") + return [int(image.shape[-1]), int(image.shape[-2])] + if hasattr(image, "size"): + width, height = image.size + return [int(width), int(height)] + raise TypeError(f"Cannot determine original image size from {type(image)}") + + +@dataclass +class EncoderPreprocessResult: + mm_inputs: dict + grid_thw: Union[torch.Tensor, List] + token_counts: List[int] + + +class EncoderPreprocessor: + """CPU-bound multimodal preprocessing pipeline. + + Takes raw media URLs / base64 data and produces HF processor output dicts + (CPU tensors). The GPU model is never touched here — only the HF + image/video/audio processors are invoked. + + Parameters + ---------- + server_args : ServerArgs + Server configuration (model path, processor flags, etc.). + model_config : ModelConfig + Model configuration (hf_config, hidden_size, etc.). + model_preprocessor : callable, optional + Optional model-specific preprocessor (``model.preprocess_mm_for_encoder``). + When provided, overrides the default HF processor path for the given + modality. + """ + + def __init__( + self, + server_args: ServerArgs, + model_config: ModelConfig, + encoder_media_processor_config: EncoderMediaProcessorConfig, + model_preprocessor: Optional[Callable] = None, + ): + self.server_args = server_args + self.model_config = model_config + self._model_preprocessor = model_preprocessor + self.encoder_media_processor_config = encoder_media_processor_config + self.model_type = getattr( + model_config.hf_config, "model_type", "unknown" + ).lower() + + self.device = get_device().device + + use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get() + self.use_image_processor_gpu = ( + use_image_processor_gpu + and resolve_image_processor_backend(server_args) != "pil" + ) + + self._load_mm_processor(server_args) + self._supported_modalities = frozenset( + modality + for modality, processor in ( + (Modality.IMAGE, self.image_processor), + (Modality.VIDEO, self.video_processor), + (Modality.AUDIO, self.audio_processor), + ) + if processor is not None or self._model_preprocessor is not None + ) + self._build_vision_config(get_mm().mm_process_config) + self.model_audio_sr = self._resolve_audio_sr() + logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz") + + self.preproc_executor = concurrent.futures.ThreadPoolExecutor( + max_workers=envs.SGLANG_ENCODER_PREPROC_WORKERS.get() + ) + self.io_executor = concurrent.futures.ThreadPoolExecutor( + max_workers=int(os.environ.get("SGLANG_ENCODER_MM_LOAD_WORKERS", 4)) + ) + + # ------------------------------------------------------------------ + # HF Processor Loading + # ------------------------------------------------------------------ + + def _load_mm_processor(self, server_args: ServerArgs): + from transformers import AutoImageProcessor, AutoVideoProcessor + + image_processor_backend = resolve_image_processor_backend(server_args) + image_processor_kwargs = ( + {} + if image_processor_backend == "auto" + else {"backend": image_processor_backend} + ) + try: + self.image_processor = AutoImageProcessor.from_pretrained( + get_serving().tokenizer_path or get_model().model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + **image_processor_kwargs, + ) + except Exception as e: + logger.warning(f"Failed to load image processor: {e}") + self.image_processor = None + + try: + self.video_processor = AutoVideoProcessor.from_pretrained( + get_serving().tokenizer_path or get_model().model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + ) + except Exception as e: + logger.warning(f"Failed to load video processor: {e}") + self.video_processor = None + + try: + _audio_proc = AutoProcessor.from_pretrained( + get_serving().tokenizer_path or get_model().model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + ) + if not hasattr(_audio_proc, "feature_extractor"): + logger.warning( + "Loaded AutoProcessor has no feature_extractor attribute, " + "audio processing will be unavailable." + ) + self.audio_processor = None + else: + self.audio_processor = _audio_proc + except Exception as e: + logger.warning(f"Failed to load audio processor: {e}") + self.audio_processor = None + + # ------------------------------------------------------------------ + # Config Validation + # ------------------------------------------------------------------ + + def _build_vision_config(self, mm_process_config): + self.vision_config = ( + mm_process_config.get("vision_config", {}) + if mm_process_config is not None + else {} + ) + for modality_str in ["image", "video", "audio"]: + if not self.vision_config.get(modality_str, None): + self.vision_config[modality_str] = {} + if self.use_image_processor_gpu: + self.vision_config[modality_str]["device"] = self.device + + if modality_str == "video": + video_defaults = {"fps": 2.0, "max_frames": 768, "min_frames": 4} + for k, v in video_defaults.items(): + self.vision_config["video"].setdefault(k, v) + + if modality_str == "audio": + if "return_attention_mask" not in self.vision_config["audio"]: + self.vision_config["audio"]["return_attention_mask"] = True + if "padding" not in self.vision_config["audio"]: + if self.model_type == "qwen2_audio": + self.vision_config["audio"]["padding"] = "max_length" + else: + self.vision_config["audio"]["padding"] = True + if "truncation" not in self.vision_config["audio"]: + if ( + hasattr(self, "audio_processor") + and self.audio_processor is not None + ): + if self.audio_processor.__class__.__name__ in { + "Gemma3nProcessor", + "GlmAsrProcessor", + "Qwen2AudioProcessor", + "Qwen3OmniMoeProcessor", + }: + self.vision_config["audio"]["truncation"] = False + + def _resolve_audio_sr(self) -> int: + def _read(obj, attr): + if obj is None: + return None + if isinstance(obj, dict): + return obj.get(attr) + return getattr(obj, attr, None) + + audio_cfg = self.vision_config.get("audio", {}) + sr = audio_cfg.get("audio_sampling_rate") + if sr: + return int(sr) + + hf_cfg = self.model_config.hf_config + thinker_cfg = _read(hf_cfg, "thinker_config") + pc = _read(thinker_cfg, "processor_config") or _read(hf_cfg, "processor_config") + sr = _read(pc, "audio_sampling_rate") + if sr: + return int(sr) + ac = _read(thinker_cfg, "audio_config") or _read(hf_cfg, "audio_config") + for attr in ("sampling_rate", "sample_rate"): + sr = _read(ac, attr) + if sr: + return int(sr) + + sr = audio_cfg.get("sampling_rate") + if sr: + return int(sr) + logger.warning( + "No audio sampling rate found in mm_config or hf_config; " + "falling back to 16000 Hz. If the model expects a different SR " + "(e.g. MiMo-V2 defaults to 24000), audio will be warped." + ) + return 16000 + + # ------------------------------------------------------------------ + # Media I/O + # ------------------------------------------------------------------ + + def _load_single_item( + self, + data, + modality: Modality, + frame_count_limit=None, + discard_alpha_channel=True, + ): + from sglang.srt.disaggregation.encoder.server import BadRequestError, MMError + + media_metadata = {} + content_hash = None + if isinstance(data, dict): + if "url" not in data: + return data + media_metadata = {key: value for key, value in data.items() if key != "url"} + content_hash = parse_content_hash(data.get("content_hash")) + data = data["url"] + try: + if modality == Modality.IMAGE: + if content_hash is not None: + snapshot = snapshot_media(data) + if snapshot.content_digest != content_hash: + raise BadRequestError( + "Encoder media content hash mismatch: " + f"expected {content_hash}, got {snapshot.content_digest}" + ) + data = snapshot.data + gpu_image_decode = ( + self.encoder_media_processor_config.image_decode_mode + if self.use_image_processor_gpu + else False + ) + img, _ = load_image(data, gpu_image_decode) + if ( + discard_alpha_channel + and not isinstance(img, torch.Tensor) + and img.mode != "RGB" + ): + img = img.convert("RGB") + if ( + media_metadata + and self.encoder_media_processor_config.preserve_media_metadata + ): + return { + "type": "image", + "image": img, + **media_metadata, + } + return img + elif modality == Modality.VIDEO: + return load_video(data, frame_count_limit) + elif modality == Modality.AUDIO: + return load_audio(data, self.model_audio_sr) + + except MMError: + raise + except CLIENT_MEDIA_EXCEPTIONS as e: + # Not ValueError: the DP envelope classifies by `.code`, which only + # MMError carries. + raise BadRequestError(f"Error while loading data {data}: {e}") from e + except Exception as e: + raise RuntimeError(f"Error while loading data {data}: {e}") + + def _submit_data_loading_tasks(self, items, modalities): + futures = [] + task_info = [] + + for data, modality in zip(items, modalities): + if modality is not None: + futures.append( + self.io_executor.submit( + self._load_single_item, + data, + modality, + ) + ) + task_info.append((modality, data)) + return futures, task_info + + async def _flatten_and_load_data_by_modality(self, mm_items, modality): + if not isinstance(mm_items, (list, tuple)): + futures, _ = self._submit_data_loading_tasks([mm_items], [modality]) + return await asyncio.wrap_future(futures[0]) + + if len(mm_items) > 0 and isinstance(mm_items[0], (list, tuple)): + flat_data = [] + flat_indices = [] + for group_idx, item_group in enumerate(mm_items): + for item in item_group: + flat_data.append(item) + flat_indices.append(group_idx) + + futures, _ = self._submit_data_loading_tasks( + flat_data, [modality] * len(flat_data) + ) + + async_futures = [asyncio.wrap_future(f) for f in futures] + results = await asyncio.gather(*async_futures) + + nested_results = [[] for _ in range(len(mm_items))] + for idx, result in zip(flat_indices, results): + nested_results[idx].append(result) + + return nested_results + + else: + futures, _ = self._submit_data_loading_tasks( + mm_items, [modality] * len(mm_items) + ) + async_futures = [asyncio.wrap_future(f) for f in futures] + return await asyncio.gather(*async_futures) + + async def _flatten_and_load_images(self, mm_items): + return await self._flatten_and_load_data_by_modality(mm_items, Modality.IMAGE) + + async def _flatten_and_load_videos(self, mm_items): + if not isinstance(mm_items, (list, tuple)): + mm_items = [mm_items] + + futures, _ = self._submit_data_loading_tasks( + mm_items, [Modality.VIDEO] * len(mm_items) + ) + async_futures = [asyncio.wrap_future(f) for f in futures] + video_items = await asyncio.gather(*async_futures) + + video_processor_kwargs = {} + if "qwen" in self.model_type: + video_processed = [ + await preprocess_video( + video, video_config=self.vision_config.get("video", {}) + ) + for video in video_items + ] + videos, video_metadata = map(list, zip(*video_processed)) + video_processor_kwargs["do_sample_frames"] = False + if video_metadata: + video_processor_kwargs["video_metadata"] = video_metadata + return videos, video_processor_kwargs + else: + raise NotImplementedError( + f"Video processing is not supported for {self.model_type} model." + ) + + async def _flatten_and_load_audios(self, mm_items): + return await self._flatten_and_load_data_by_modality(mm_items, Modality.AUDIO) + + # ------------------------------------------------------------------ + # HF Processor Calls + # ------------------------------------------------------------------ + + async def process_mm_items( + self, mm_items, modality: Modality + ) -> EncoderPreprocessResult: + """Process multimodal items through the HF processor pipeline. + + Returns the ``mm_inputs`` dict produced by the HF image/video/audio + processor, its normalized grid metadata, and one output token count per + grid entry. Does not look up ``get_feature_fn``; that stays in + :class:`MMEncoder`. + """ + if modality == Modality.IMAGE: + mm_inputs = await self._process_image_items( + mm_items, self._model_preprocessor + ) + elif modality == Modality.VIDEO: + mm_inputs = await self._process_video_items( + mm_items, self._model_preprocessor + ) + elif modality == Modality.AUDIO: + mm_inputs = await self._process_audio_items( + mm_items, self._model_preprocessor + ) + else: + raise ValueError(f"Unsupported modality: {modality}") + grid_thw = self._get_mm_grid_dim(mm_inputs, modality) + token_counts = [self.get_num_tokens(grid, modality) for grid in grid_thw] + return EncoderPreprocessResult( + mm_inputs=mm_inputs, + grid_thw=grid_thw, + token_counts=token_counts, + ) + + def supports_modality(self, modality: Modality) -> bool: + return modality in self._supported_modalities + + async def process_batch_mm_items( + self, requests: List[dict], modality: Modality + ) -> tuple[EncoderPreprocessResult, List[int]]: + """Flatten requests, run the processor once, and return batch layout.""" + flat_items, items_per_req = self._flatten_batch_requests(requests, modality) + result = await self.process_mm_items(flat_items, modality) + return result, items_per_req + + def _flatten_batch_requests( + self, requests: List[dict], modality: Modality + ) -> tuple[List, List[int]]: + # items_per_req counts grid entries (post-expansion) so per-request + # slicing of grid_dim/final_slices stays aligned for processors that + # expand one leaf into multiple grids (e.g. Kimi-VL/K2.5/K3 dict-of-images). + flat_items = [] + items_per_req = [] + for req in requests: + leaves = self._flatten_nested_items(req["mm_items"]) + flat_items.extend(leaves) + items_per_req.append(sum(self._grid_count_per_leaf(leaves, modality))) + return flat_items, items_per_req + + async def _process_image_items(self, mm_items, model_preprocessor): + if not (self.image_processor or model_preprocessor): + raise ValueError("No image processor available") + images = await self._flatten_and_load_images(mm_items) + if self.model_type in ["kimi_k25", "kimi_k3", "kimi_vl"]: + images = self._normalize_kimi_encoder_images(images) + original_image_sizes = [_get_original_image_size(item) for item in images] + if model_preprocessor: + processor_output = invoke_encoder_preprocessor( + model_preprocessor, + images, + Modality.IMAGE, + self.vision_config, + image_processor=self.image_processor, + use_gpu_preprocessing=self.use_image_processor_gpu, + ) + if ( + isinstance(processor_output, EncoderPreprocessOutput) + and processor_output.materialize_local_items is not None + ): + parallel = get_parallel() + await asyncio.get_running_loop().run_in_executor( + self.preproc_executor, + processor_output.materialize_for_rank, + parallel.attn_tp_rank, + parallel.attn_tp_size, + ) + return processor_output + image_config = self.vision_config.get("image", {}) + processor_input = await asyncio.get_running_loop().run_in_executor( + self.preproc_executor, + functools.partial(self.image_processor, images=images, **image_config), + ) + if self.model_type == "kimi_k3": + processor_input["original_image_sizes"] = original_image_sizes + return processor_input + + async def _process_video_items(self, mm_items, model_preprocessor): + if model_preprocessor: + return model_preprocessor(mm_items, Modality.VIDEO, self.vision_config) + if not self.video_processor: + raise ValueError("No video processor available") + + videos, video_processor_kwargs = await self._flatten_and_load_videos(mm_items) + processor_input = await asyncio.get_running_loop().run_in_executor( + self.preproc_executor, + functools.partial( + self.video_processor, videos=videos, **video_processor_kwargs + ), + ) + + if ( + self.model_type + in [ + "qwen3_vl", + "qwen3_vl_moe", + "qwen3_5", + "qwen3_5_moe", + "intern_s2_preview", + ] + and video_processor_kwargs.get("video_metadata", None) is not None + ): + video_metadata = video_processor_kwargs["video_metadata"] + try: + merge_size = ( + self.model_config.hf_config.vision_config.spatial_merge_size + ) + except (AttributeError, KeyError): + merge_size = 2 + video_timestamps = [] + for metadata in video_metadata: + video_fps = metadata.get("fps", None) or 24 + frames_indices = metadata.get("frames_indices", None) + timestamps = self._calculate_timestamps( + frames_indices, video_fps, merge_size + ) + video_timestamps.append(timestamps) + processor_input["video_timestamps"] = video_timestamps + elif ( + self.model_type in ["qwen2_5_vl", "qwen2_5_omni", "qwen3_omni_moe"] + and processor_input.get("video_grid_thw", None) is not None + ): + video_grid_thw = processor_input["video_grid_thw"] + try: + temporal_patch_size = self.video_processor.temporal_patch_size + except AttributeError: + temporal_patch_size = 2 + fps_list = [ + self.vision_config.get("video", {}).get("fps", None) or 2 + ] * len(video_grid_thw) + second_per_grid_ts = [(temporal_patch_size / fps) for fps in fps_list] + second_per_grid_ts_tensor = torch.tensor( + second_per_grid_ts, dtype=torch.float32 + ) + processor_input["second_per_grid_ts"] = second_per_grid_ts_tensor + + return processor_input + + async def _process_audio_items(self, mm_items, model_preprocessor): + audios = await self._flatten_and_load_audios(mm_items) + + if model_preprocessor: + return model_preprocessor(audios, Modality.AUDIO, self.vision_config) + + if not self.audio_processor: + raise ValueError("No audio processor available") + + audio_config = self.vision_config.get("audio", {}) + processor_input = await asyncio.get_running_loop().run_in_executor( + self.preproc_executor, + functools.partial( + self.audio_processor.feature_extractor, audios, **audio_config + ), + ) + processor_input["feature_attention_mask"] = processor_input.pop( + "attention_mask" + ) + input_lengths = torch.tensor( + processor_input["feature_attention_mask"].sum(-1), dtype=torch.long + ) + processor_input["audio_feature_lens_raw"] = input_lengths + output_lengths = self._get_feat_extract_output_lengths(input_lengths) + processor_input["audio_feature_lens"] = output_lengths + return processor_input + + # ------------------------------------------------------------------ + # Audio Feature Length Computation + # ------------------------------------------------------------------ + + def _get_feat_extract_output_lengths(self, feature_lens): + if self.model_type in ["qwen2_audio", "qwen2_5_omni"]: + input_length = (feature_lens - 1) // 2 + 1 + return (input_length - 2) // 2 + 1 + elif self.model_type in ["qwen3_asr", "qwen3_omni_moe"]: + input_lengths_leave = feature_lens % 100 + feat_lengths = (input_lengths_leave - 1) // 2 + 1 + output_lengths = ( + ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (feature_lens // 100) * 13 + ) + return output_lengths + elif self.model_type == "mimo_v2": + return feature_lens + else: + logger.warning( + f"Fallback to original HF audio sample logic for {self.model_type}" + ) + input_length = (feature_lens - 1) // 2 + 1 + return (input_length - 2) // 2 + 1 + + def _get_mm_grid_dim(self, mm_inputs: dict, modality: Modality): + # Kimi K2.5/K3 vision processors only emit `grid_thws`; prefer it over generic keys + # so we never pick a mis-typed or stale `image_grid_hws` field from kwargs. + attrs = _mm_grid_attrs[modality] + model_type = (self.model_type or "").lower() + if modality == Modality.IMAGE: + # Kimi K2.5/K3 emit grid_thws, while Kimi-VL emits image_grid_hws. + # Other model types keep the generic attr order above. + if model_type in ("kimi_k25", "kimi_k3"): + attrs = ("grid_thws", "image_grid_thw", "image_grid_hws") + elif model_type == "kimi_vl": + attrs = ("image_grid_hws", "image_grid_thw", "grid_thws") + + for attr in attrs: + if attr in mm_inputs and mm_inputs[attr] is not None: + return _convert(mm_inputs[attr]) + raise ValueError( + f"Grid dim ({_mm_grid_attrs[modality]}) not found in {mm_inputs}" + ) + + def get_num_patches( + self, grid: Union[torch.Tensor, List[int]], modality: Modality + ) -> int: + """Calculate number of raw patches (before merge/sampling). Used for pixel_values slicing.""" + if modality == Modality.AUDIO: + return int(grid.item()) + if self.model_type == "kimi_vl" and modality == Modality.IMAGE: + h, w = self._kimi_hw_from_patch_grid(grid) + return h * w + return int(grid[0] * grid[1] * grid[2]) + + @staticmethod + def _kimi_hw_from_patch_grid( + grid: Union[torch.Tensor, np.ndarray, List[int], Tuple[int, ...]], + ) -> Tuple[int, int]: + """Extract (height, width) from Kimi 2D or 3D patch-grid metadata.""" + if isinstance(grid, torch.Tensor): + values = grid.flatten().tolist() + elif isinstance(grid, np.ndarray): + values = grid.reshape(-1).tolist() + else: + values = np.asarray(grid).reshape(-1).tolist() + + if len(values) not in (2, 3): + raise ValueError( + f"Invalid Kimi image grid metadata: {values}; " + "expected [h, w] or [t, h, w]" + ) + return int(values[-2]), int(values[-1]) + + def _kimi_tokens_from_patch_grid(self, grid: Union[torch.Tensor, List[int]]) -> int: + """Calculate Kimi image tokens from either 2D or 3D patch metadata.""" + h, w = self._kimi_hw_from_patch_grid(grid) + merge_h, merge_w = self.model_config.hf_config.vision_config.merge_kernel_size + return (h * w) // (merge_h * merge_w) + + def get_num_tokens( + self, grid: Union[torch.Tensor, List[int]], modality: Modality + ) -> int: + """Compatibility helper for callers that still provide patch grids.""" + if modality == Modality.AUDIO: + input_length = self.get_num_patches(grid, modality) + return self._get_feat_extract_output_lengths(input_length) + else: + if ( + self.model_type in ["kimi_k25", "kimi_k3", "kimi_vl"] + and modality == Modality.IMAGE + ): + return self._kimi_tokens_from_patch_grid(grid) + merge_size = getattr(self.image_processor, "merge_size", 2) + return self.get_num_patches(grid, modality) // (merge_size**2) + + # ------------------------------------------------------------------ + # Video Timestamp Computation + # ------------------------------------------------------------------ + + def _calculate_timestamps(self, indices, video_fps: float, merge_size: int = 2): + if not isinstance(indices, list): + indices = indices.tolist() + if len(indices) % merge_size != 0: + indices.extend( + indices[-1] for _ in range(merge_size - len(indices) % merge_size) + ) + timestamps = [idx / video_fps for idx in indices] + timestamps = [ + (timestamps[i] + timestamps[i + merge_size - 1]) / 2 + for i in range(0, len(timestamps), merge_size) + ] + return timestamps + + # ------------------------------------------------------------------ + # Kimi Normalization + # ------------------------------------------------------------------ + + def _normalize_kimi_encoder_images(self, images): + """Normalize Kimi image inputs for the image processor call.""" + from PIL import Image as PILImage + + def wrap_one(img): + if isinstance(img, dict) and img.get("type") in ("image", "video_chunk"): + return [img] + if isinstance(img, PILImage.Image): + return [{"type": "image", "image": img}] + return [img] + + if not images: + return images + + # Disagg may supply nested lists from grouped routing. + images = self._flatten_nested_items(images) + + if self.model_type == "kimi_vl": + normalized = [] + for img in images: + if ( + isinstance(img, dict) + and img.get("type") == "image" + and "image" in img + ): + inner = img["image"] + if isinstance(inner, (list, tuple)): + normalized.extend(self._flatten_nested_items(inner)) + else: + normalized.append(inner) + else: + normalized.append(img) + return normalized + + # Kimi-K2.5/K3 vision processors expect media dicts. + normalized = [] + for img in images: + wrapped = wrap_one(img) + for media in wrapped: + if ( + isinstance(media, dict) + and media.get("type") == "image" + and isinstance(media.get("image"), (list, tuple)) + ): + for inner in self._flatten_nested_items(media["image"]): + normalized.append({**media, "image": inner}) + else: + normalized.append(media) + + return normalized + + # ------------------------------------------------------------------ + # Utility Helpers + # ------------------------------------------------------------------ + + @staticmethod + def _flatten_nested_items(items): + if not isinstance(items, (list, tuple)): + return [items] + + flat = [] + for item in items: + if isinstance(item, (list, tuple)): + flat.extend(EncoderPreprocessor._flatten_nested_items(item)) + else: + flat.append(item) + return flat + + def _grid_count_per_leaf(self, leaves: List, modality: Modality) -> List[int]: + """Number of grid entries each leaf produces under the model's processor. + + Most processors map 1 leaf -> 1 grid. Kimi-VL/K2.5/K3 image processors expand + a leaf shaped {"type": "image", "image": [pil1, pil2, ...]} into N grids. + """ + if ( + self.model_type not in ("kimi_k25", "kimi_k3", "kimi_vl") + or modality != Modality.IMAGE + ): + return [1] * len(leaves) + + def count(leaf): + if ( + isinstance(leaf, dict) + and leaf.get("type") == "image" + and isinstance(leaf.get("image"), (list, tuple)) + ): + return len(self._flatten_nested_items(leaf["image"])) + return 1 + + return [count(leaf) for leaf in leaves] diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py similarity index 75% rename from python/sglang/srt/disaggregation/encode_receiver.py rename to python/sglang/srt/disaggregation/encoder/receiver.py index 041353aa3..b44cdf08a 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -54,6 +54,14 @@ if TYPE_CHECKING: from sglang.srt.managers.scheduler import Scheduler +def _mark_keep_device_embedding(mm_inputs) -> None: + """Tell general_mm_embed_routine not to copy embeddings back to CPU.""" + if mm_inputs is None: + return + for item in mm_inputs.mm_items: + item.keep_device_embedding = True + + class EncoderBootstrapServer: """Lightweight bootstrap server for dynamic encoder discovery. @@ -386,28 +394,6 @@ def _grpc_encode_request(target, encode_request): channel.close() -def _grpc_send_request(target, request_json): - import grpc - from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc - - timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get() - channel = grpc.insecure_channel(target) - stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) - try: - stub.Send( - sglang_encoder_pb2.SendRequest( - req_id=request_json["req_id"], - prefill_host=request_json["prefill_host"], - embedding_port=request_json["embedding_port"], - session_id=request_json["session_id"], - buffer_address=request_json["buffer_address"], - ), - timeout=timeout_secs, - ) - finally: - channel.close() - - class EmbeddingData: def __init__( self, @@ -434,7 +420,9 @@ class EmbeddingData: self.shape = embedding_shape else: self.shape = list(embedding.shape) if embedding is not None else None - self.cached_embedding = None + # Encoder-side mooncake MR for `embedding`. Underscored so + # copy_without_embedding drops this process-local address. + self._mr_ptr: Optional[int] = None self.error_msg = error_msg # Coerce to plain int: this object crosses process boundaries via # safe_pickle_loads, whose allowlist blocks http.HTTPStatus. @@ -466,8 +454,7 @@ class EmbeddingData: error_code=self.error_code, ) for key, value in self.__dict__.items(): - # cached_embedding is a GPU tensor used only by mooncake's in-process - if key.startswith("_") or key in ("embedding", "cached_embedding"): + if key.startswith("_") or key == "embedding": continue setattr(new_data, key, value) return new_data @@ -701,7 +688,17 @@ class MultiModalEmbeddingData(EmbeddingData): self._set_image_meta_for_part(pid, embedding_data) -class WaitingImageRequestStatus(IntEnum): +def _aggregate_embedding_part(current, recv_obj, model_type): + """Fold one received part into the aggregate (the first part creates it).""" + if current is None: + return MultiModalEmbeddingData.from_embedding_data( + recv_obj, model_type=model_type + ) + current.add(recv_obj) + return current + + +class WaitingMMRequestStatus(IntEnum): FAIL = -1 PENDING = 0 SUCCESS = 1 @@ -760,8 +757,13 @@ def calculate_modality_num_parts(modalities, num_items_assigned): return total_num_parts, modality_num_parts -# For zmq_to_scheduler -class WaitingImageRequest: +class WaitingMMRequestBase(ABC): + """One in-flight multimodal request on a scheduler rank, waiting for + encoder embeddings. Owns the shared machinery: the ZMQ receive loop, + failure handling (_fail_and_release), pool-slot lifetime, and the + TP-consistent status. Subclasses bind the transport. + """ + def __init__( self, rid: str, @@ -771,6 +773,7 @@ class WaitingImageRequest: model_type, host_name, receive_count, + embedding_pool: Optional["EmbeddingPool"] = None, zmq_context=None, embedding_port=None, ): @@ -785,10 +788,8 @@ class WaitingImageRequest: self.host_name = host_name self.receive_count = receive_count self.num_items_assigned = recv_req.num_items_assigned - self.zmq_context = zmq_context + self.zmq_context = zmq_context or zmq.Context() if embedding_port is None: - if self.zmq_context is None: - raise ValueError("zmq_context is required for a per-request socket") self.embedding_port, self.recv_socket = get_zmq_socket_on_host( self.zmq_context, zmq.PULL, host=host_name ) @@ -798,11 +799,213 @@ class WaitingImageRequest: logger.info(f"Waiting for input {self.embedding_port = }") self.recv_embedding_data = None # ok=1 pending=0 fail=-1 - self.status = WaitingImageRequestStatus.PENDING + self.status = WaitingMMRequestStatus.PENDING self.error_msg = None self.error_code = None self.start_time = time.time() + # Optional GPU pool bounding received embeddings (zmq_to_scheduler): + # _try_recv_mm_data stages parts into one slot, staying PENDING while + # the pool is full. + self.embedding_pool = embedding_pool + self.embeddings_buffer = None + self._pool_slot_id: Optional[int] = None + # Success-path finalizer handle so abort can release the slot early. + self._mm_finalizer: Optional[weakref.finalize] = None + self._pool_full_warned = False + @abstractmethod + def send_encode_request(self) -> None: + """Kick off the transport-specific encode / receive flow.""" + + def _try_recv_mm_data(self): + if self.status != WaitingMMRequestStatus.PENDING: + return + + # A complete request can remain pending while the GPU pool is full. + # Retry assembly on every scheduler tick, including shared-socket mode. + if self.recv_embedding_data is not None and self.recv_embedding_data.ready: + if self._assemble_mm_inputs_from_embeddings(): + self.close_recv_socket() + return + + if self.recv_socket is None: + 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: + # No data available yet, wait a bit and retry + return + except zmq.ZMQError: + # Socket closed by another path (e.g. the RDMA receive thread + # after an encoder error); status is already terminal. + return + self.consume_parts(parts) + if self.status != WaitingMMRequestStatus.PENDING: + return + + def consume_parts(self, parts) -> None: + """Consume one message from either a per-request or shared ZMQ socket.""" + if self.status != WaitingMMRequestStatus.PENDING: + return + + try: + 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}: " + f"{recv_obj.error_msg} {recv_obj.error_code = }" + ) + self._fail_and_release(recv_obj.error_msg, recv_obj.error_code) + return + if not self._is_valid_embedding_part(recv_obj): + return + # ZMQ materializes frame 1; RDMA already wrote the registered buffer. + self._extract_embedding_from_buffer(recv_obj, parts) + self.recv_embedding_data = _aggregate_embedding_part( + self.recv_embedding_data, recv_obj, self.model_type + ) + except Exception as e: + # A malformed message must fail this request, not the scheduler loop. + logger.exception("Failed to decode embedding message for rid=%s", self.rid) + self._fail_and_release(f"Failed to decode embedding message: {e}") + return + + if ( + self.recv_embedding_data.ready + and self._assemble_mm_inputs_from_embeddings() + ): + self.close_recv_socket() + + def close_recv_socket(self) -> None: + if self.recv_socket is not None: + self.recv_socket.close() + self.recv_socket = None + + def _fail_and_release(self, error_msg, error_code=None) -> None: + """Terminal failure: record the error, free buffers, close the socket.""" + self.error_msg = error_msg + self.error_code = error_code + self.status = WaitingMMRequestStatus.FAIL + self._cleanup_gpu_buffer() + self.close_recv_socket() + + async def _check_encoder_responses(self, responses, endpoint: str) -> bool: + """Validate gathered encoder responses; on the first error, FAIL the + request and release its resources. Returns True if all succeeded.""" + msg = await _extract_encoder_error(responses, endpoint, f"rid={self.rid}") + if msg is None: + return True + self._fail_and_release(msg) + return False + + def _is_valid_embedding_part(self, recv_obj) -> bool: + """Check for and drop stale or out-of-sync payloads; normalize the part req_id to the original rid.""" + original_req_id = extract_original_req_id(recv_obj.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)" + ) + return False + recv_obj.req_id = original_req_id + return True + + @abstractmethod + def _extract_embedding_from_buffer(self, recv_obj, parts) -> None: + """Materialize ``recv_obj.embedding`` from one received part message.""" + + @abstractmethod + def _prepare_embedding_buffer(self) -> bool: + """Make ``embeddings_buffer`` ready for assembly, or leave it None + for the CPU-concat path. False = not ready yet, stay PENDING.""" + + def _view_dtype(self): + """dtype of the bytes in ``embeddings_buffer``.""" + return self.recv_embedding_data.dtype + + def _assemble_mm_inputs_from_embeddings(self) -> bool: + """Assemble mm_inputs from the received embeddings and mark + SUCCESS/FAIL. Failures are caught so they still reach the TP-wide + status all-reduce. + + Returns True when done so the caller closes the recv socket; False + when the buffer is not ready yet (stay PENDING, retry next tick) or + the request can never fit the pool (already FAILed, socket closed). + """ + try: + if not self._prepare_embedding_buffer(): + return False + if self.embeddings_buffer is not None: + # Zero-copy per-modality views into the GPU buffer; slot + # lifetime is bound to mm_inputs GC in _finish_assemble. + recv_embedding = _view_pool_buffer_by_modality( + self.embeddings_buffer, + self.recv_embedding_data, + self._view_dtype(), + ) + else: + recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True) + self._finish_assemble(recv_embedding) + # Releases whatever is still attached: no-op once the slot was + # detached; RDMA's override also deregisters non-pool buffers. + self._cleanup_gpu_buffer() + except Exception as e: + self._fail_assemble(e) + return True + + def _finish_assemble(self, recv_embedding) -> None: + """get_mm_data → bind pool slot → publish onto recv_req → SUCCESS.""" + mm_inputs = self.mm_processor.get_mm_data( + _select_mm_processor_prompt(self.recv_req, self.mm_processor), + recv_embedding, + **self.recv_embedding_data.get_mm_extra_meta(), + ) + self._bind_pool_slot_to_mm_inputs(mm_inputs) + self.recv_req.mm_inputs = mm_inputs + self.recv_req.input_ids = array("q", mm_inputs.input_ids) + self.status = WaitingMMRequestStatus.SUCCESS + + def _fail_assemble(self, e: Exception) -> None: + logger.exception("Failed to assemble multimodal inputs for rid=%s", self.rid) + self._fail_and_release(f"Failed to assemble multimodal inputs: {e}") + + def _bind_pool_slot_to_mm_inputs(self, mm_inputs) -> bool: + """Bind pool-slot release to mm_inputs GC. Returns True if bound.""" + if ( + mm_inputs is None + or self._pool_slot_id is None + or self.embedding_pool is None + ): + return False + # Keep the handle so abort can release the slot immediately. + self._mm_finalizer = self.embedding_pool.release_on_gc( + mm_inputs, self._pool_slot_id + ) + _mark_keep_device_embedding(mm_inputs) + # Detach so _cleanup_gpu_buffer no-ops; finalize now owns release. + self._pool_slot_id = None + self.embeddings_buffer = None + return True + + def _cleanup_gpu_buffer(self): + 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 + + def release_resources(self): + """Free pool/GPU resources on abort/fail/timeout. Idempotent.""" + self._cleanup_gpu_buffer() + finalizer, self._mm_finalizer = self._mm_finalizer, None + if finalizer is not None: + finalizer() # at-most-once; a later GC call becomes a no-op + + +# For zmq_to_scheduler: embedding parts arrive as ZMQ payload frames and +# are optionally staged into the GPU EmbeddingPool. +class WaitingZmqRequest(WaitingMMRequestBase): def send_encode_request(self): async def _send_single_request(session, url, payload): @@ -878,6 +1081,15 @@ class WaitingImageRequest: ) else: logger.debug(f"Request {i} succeeded.") + failed = [r for r in results if isinstance(r, BaseException)] + if failed: + # A rank without a registered receive URL can never be + # pushed to; fail via the normal completion path now + # instead of pending until the embedding wait times out. + self._fail_and_release( + f"Failed to register receive URL with encoder: {failed[0]!r}", + int(HTTPStatus.BAD_GATEWAY), + ) asyncio.run( send_embedding_port( @@ -888,106 +1100,63 @@ class WaitingImageRequest: ) ) - def _try_recv_mm_data(self): - if self.status != WaitingImageRequestStatus.PENDING: - return - if self.recv_socket is None: - return - while self.status == WaitingImageRequestStatus.PENDING: - try: - parts = self.recv_socket.recv_multipart(flags=zmq.NOBLOCK, copy=False) - except zmq.Again: - # No data available yet, wait a bit and retry - return - self.consume_parts(parts) + def _extract_embedding_from_buffer(self, recv_obj, parts) -> None: + """ZMQ transport carries the embedding bytes as frame 1. Clone so we + don't depend on the ZMQ buffer after the next recv.""" + buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] + recv_obj.embedding = ( + torch.frombuffer(buffer, dtype=recv_obj.dtype) + .reshape(recv_obj.shape) + .clone() + ) - def consume_parts(self, parts): - if self.status != WaitingImageRequestStatus.PENDING: - return + def _prepare_embedding_buffer(self) -> bool: + """Stage the CPU parts into the GPU pool when one is configured; + without a pool the CPU-concat path is used (buffer stays None).""" + if self.embedding_pool is None: + return True + return self._try_stage_into_pool() - try: - recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) - if getattr(recv_obj, "error_msg", None) is not None: + def _try_stage_into_pool(self) -> bool: + """Copy the received parts into one pooled GPU slot, packed in part + order (modality-contiguous by construction, see _extract_url_data — + _view_pool_buffer_by_modality asserts this). + + Returns True once ``self.embeddings_buffer`` views the slot; False + when the pool is currently full (retry next tick) or the request can + never fit (marked FAIL, socket closed). + """ + if self.embeddings_buffer is not None: + return True + parts = self.recv_embedding_data.embedding_list + total_bytes = sum(p.nbytes for p in parts if p is not None) + if total_bytes > self.embedding_pool.size_bytes: + error_msg = ( + f"EmbeddingPool cannot fit {total_bytes // (1024 * 1024)}MB " + f"(pool is {self.embedding_pool.size_bytes // (1024 * 1024)}MB). " + f"Raise SGLANG_EMBEDDING_POOL_SIZE_MB." + ) + logger.error(f"{error_msg} rid={self.rid}") + self._fail_and_release(error_msg) + return False + staged = self.embedding_pool.try_stage([p for p in parts if p is not None]) + if staged is None: + if not self._pool_full_warned: logger.warning( - f"Received error signal from encoder for {self.rid}: {recv_obj.error_msg} {recv_obj.error_code = }" + f"EmbeddingPool full; rid={self.rid} pending for " + f"{total_bytes // (1024 * 1024)}MB. Raise " + f"SGLANG_EMBEDDING_POOL_SIZE_MB if this is frequent." ) - self.error_msg = recv_obj.error_msg - self.error_code = recv_obj.error_code - self.status = WaitingImageRequestStatus.FAIL - self.close_recv_socket() - return - - # Extract original req_id from part_req_id and drop stale payloads - # that may arrive on a reused ZMQ port after a prior request aborted. - original_req_id = extract_original_req_id(recv_obj.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)" - ) - return - recv_obj.req_id = original_req_id - - buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] - recv_obj.embedding = ( - torch.frombuffer(buffer, dtype=recv_obj.dtype) - .reshape(recv_obj.shape) - .clone() - ) - - if self.recv_embedding_data is None: - self.recv_embedding_data = MultiModalEmbeddingData.from_embedding_data( - recv_obj, model_type=self.model_type - ) - else: - self.recv_embedding_data.add(recv_obj) - except Exception as e: - # A message the scheduler cannot decode (blocked unpickle, - # bad shape/dtype, ...) must fail this request, not crash the - # scheduler event loop; FAIL still reaches the TP-wide status - # all-reduce in _process_waiting_requests. - logger.exception("Failed to decode embedding message for rid=%s", self.rid) - self.error_msg = f"Failed to decode embedding message: {e}" - self.status = WaitingImageRequestStatus.FAIL - self._cleanup_gpu_buffer() - self.close_recv_socket() - return - - if not self.recv_embedding_data.ready: - return - - # Assemble mm_inputs. Wrapped so an assembly failure still reaches the - # TP-wide status all-reduce in _process_waiting_requests instead of - # raising past it. - try: - recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True) - mm_inputs = self.mm_processor.get_mm_data( - _select_mm_processor_prompt(self.recv_req, self.mm_processor), - recv_embedding, - **self.recv_embedding_data.get_mm_extra_meta(), - ) - self.recv_req.mm_inputs = mm_inputs - self.recv_req.input_ids = array("q", mm_inputs.input_ids) - self.status = WaitingImageRequestStatus.SUCCESS - except Exception as e: - logger.exception( - "Failed to assemble multimodal inputs for rid=%s", self.rid - ) - self.status = WaitingImageRequestStatus.FAIL - self.error_msg = f"Failed to assemble multimodal inputs: {e}" - self._cleanup_gpu_buffer() - self.close_recv_socket() - - def close_recv_socket(self): - if self.recv_socket is not None: - self.recv_socket.close() - self.recv_socket = None - - def _cleanup_gpu_buffer(self): - pass + self._pool_full_warned = True + return False + self.embeddings_buffer, self._pool_slot_id = staged + # Drop the CPU clones now that they live in the pool. + for i in range(len(parts)): + parts[i] = None + return True -class WaitingImageRequestGrpc(WaitingImageRequest): +class WaitingZmqRequestGrpc(WaitingZmqRequest): def send_encode_request(self): async def send_embedding_port(req_id, receive_count, host_name, embedding_port): tasks = [] @@ -1034,7 +1203,7 @@ class WaitingImageRequestGrpc(WaitingImageRequest): ) -class WaitingImageRDMARequest(WaitingImageRequest): +class WaitingRDMARequest(WaitingMMRequestBase): def __init__( self, rid, @@ -1059,93 +1228,84 @@ class WaitingImageRDMARequest(WaitingImageRequest): model_type=model_type, host_name=host_name, receive_count=receive_count, + embedding_pool=embedding_pool, zmq_context=zmq_context, embedding_port=embedding_port, ) 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 + # The receive thread owns the buffer while _receive_running; once + # _terminal latches, it releases the buffer itself on exit so the + # scheduler thread never has to wait on it. + self._buffer_lock = threading.Lock() + self._terminal = False + self._receive_running = False def send_encode_request(self): - self._encode_thread = threading.Thread( - target=self._run_encode_in_thread, daemon=True - ) - self._encode_thread.start() + # Base-class hook. The tokenizer owns /encode, so this rank only pulls + # sizes and drives the RDMA receive. + self._receive_running = True + threading.Thread(target=self._run_receive_in_thread, daemon=True).start() - def _run_encode_in_thread(self): + def _run_receive_in_thread(self): try: - asyncio.run(self._send_encode_and_rdma_request()) + asyncio.run(self._pull_meta_and_receive_embedding()) 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() + logger.error(f"RDMA receive failed for rid={self.rid}: {e}") + self._fail_and_release(str(e)) + finally: + with self._buffer_lock: + self._receive_running = False + if self._terminal: + self._release_buffer_locked() - async def _send_encode_and_rdma_request(self): + async def _pull_meta_and_receive_embedding(self): + """Pull per-part sizes, allocate the landing buffer, then drive /send. + + The tokenizer owns /encode; part_idx numbering matches it because both + derive it from the num_items_assigned frozen onto the request. + """ 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": [ - _encoder_media_item(d) - 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, - # Echoed via /send so encoder can release GPU embedding early. - "receive_count": self.receive_count, + "req_id": create_part_req_id(self.recv_req.rid, part_idx), } ) cum_idx += 1 - cum_num_items += assigned_num - part_idx_offset += num_parts + part_idx_offset += modality_num_parts[modality] 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. + # Phase 1: pull per-part sizes, blocking until the encode publishes. tasks = [ session.post( - f"{self.encoder_urls[r['encoder_idx']]}/encode", - json=r, + f"{self.encoder_urls[r['encoder_idx']]}/scheduler_receive_meta_data", + json={"req_id": r["req_id"], "part_idx": r["part_idx"]}, ) for r in encode_requests ] responses = await asyncio.gather(*tasks, return_exceptions=True) - if not await self._check_encoder_responses(responses, "/encode"): + if not await self._check_encoder_responses( + responses, "/scheduler_receive_meta_data" + ): return response_json_list = [await r.json() for r in responses] @@ -1165,21 +1325,17 @@ class WaitingImageRDMARequest(WaitingImageRequest): 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 " + # Oversize or alloc timeout — fatal for this request. + self._fail_and_release( + f"EmbeddingPool 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 + with self._buffer_lock: + self.embeddings_buffer = pool_view + self._pool_slot_id = slot_id logger.info( f"Pool-allocated Mooncake GPU landing buffer: " f"rid={self.rid}, size={total_bytes}, " @@ -1192,9 +1348,9 @@ class WaitingImageRDMARequest(WaitingImageRequest): 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 + with self._buffer_lock: + self.embeddings_buffer = gpu_buffer logger.info( f"Per-request registered Mooncake GPU landing buffer " f"(pool disabled): rid={self.rid}, size={total_bytes}, " @@ -1204,21 +1360,34 @@ class WaitingImageRDMARequest(WaitingImageRequest): self.embeddings_buffer = None buffer_address = 0 - # Phase 2 cont: POST /send with RDMA info. + # Abort/timeout may have latched _terminal; don't start RDMA into + # a buffer that will be released when this thread exits. + with self._buffer_lock: + if self._terminal: + return + + # Phase 2 cont: POST /send. Metadata carries no routing, so the + # shard comes from our own part map. + encoder_idx_by_part = { + r["part_idx"]: r["encoder_idx"] for r in encode_requests + } offset = 0 send_tasks = [] for idx in range(total_num_parts): rj = response_sorted[idx] - encoder_idx = rj.pop("encoder_idx", None) rj.update( { + "prefill_host": self.host_name, + "embedding_port": self.embedding_port, "session_id": self.embeddings_engine.session_id, "buffer_address": offset + buffer_address, + # Frees the embedding once all of us have taken it. + "receive_count": self.receive_count, } ) send_tasks.append( session.post( - f"{self.encoder_urls[encoder_idx]}/send", + f"{self.encoder_urls[encoder_idx_by_part[idx]]}/send", json=rj, ) ) @@ -1226,146 +1395,87 @@ class WaitingImageRDMARequest(WaitingImageRequest): # 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 - ): + if not await self._check_encoder_responses(send_responses, "/send"): 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. + def _extract_embedding_from_buffer(self, recv_obj, parts) -> None: + # The embedding already landed in the pre-registered GPU buffer via + # RDMA; the completion message carries no payload, so + # recv_obj.embedding stays None. + pass - 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 + def _prepare_embedding_buffer(self) -> bool: + # The receive thread already landed the embedding via RDMA (or left + # the buffer None for the zero-byte case). 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 - except zmq.ZMQError: - # The RDMA pipeline thread closed the socket after an encoder - # error (e.g. OOM). It already set status=FAIL; just bail. - 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( - _select_mm_processor_prompt(self.recv_req, self.mm_processor), - 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 = array("q", mm_inputs.input_ids) - self.status = WaitingImageRequestStatus.SUCCESS - self._cleanup_gpu_buffer() - self.recv_socket.close() + def _view_dtype(self): + # Parts carry no payload (aggregate dtype is None); use the model + # dtype the receiver was constructed with. + return self.dtype 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 + """Latch _terminal and release the GPU buffer. While the receive + thread runs it owns the buffer (RDMA may be in flight), so release + is deferred to its exit hook instead of blocking here. Idempotent.""" + with self._buffer_lock: + self._terminal = True + if not self._receive_running: + self._release_buffer_locked() + + def _release_buffer_locked(self): + """Caller must hold _buffer_lock.""" + if self.embeddings_buffer is None: + return + if self._pool_slot_id is not None: + # Pool-backed: the backing tensor stays registered; just free the slot. + self.embedding_pool.release(self._pool_slot_id) + self._pool_slot_id = None + else: 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 + self.embeddings_buffer = None + + +async def _extract_encoder_error(responses, endpoint, context, encode_requests=None): + """Return the first error among gathered encoder responses, or None. + + Pure check — logs each error but has no other side effects; the caller + decides how to react. ``encode_requests`` optionally enriches each log + line with the matching request's encoder label. + """ + for i, resp in enumerate(responses): + ctx = context + if encode_requests is not None: + label = encode_requests[i].get( + "encoder_url", f"idx={encode_requests[i].get('encoder_idx')}" + ) + ctx = f"{context}, encoder={label}" + if isinstance(resp, asyncio.TimeoutError): + timeout_val = envs.SGLANG_ENCODER_HTTP_TIMEOUT.get() + logger.error( + f"Encoder {endpoint} timeout ({timeout_val}s) for {ctx} " + f"(request {i})" + ) + return f"Encoder {endpoint} timeout ({timeout_val}s)" + if isinstance(resp, Exception): + logger.error( + f"Encoder {endpoint} failed for {ctx} (request {i}): {resp}", + exc_info=resp, + ) + return str(resp) + if 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}") + return msg + return None def _sort_responses_and_compute_total_bytes(response_json_list, total_num_parts): @@ -1380,27 +1490,30 @@ def _sort_responses_and_compute_total_bytes(response_json_list, total_num_parts) return embedding_sizes, response_sorted, total_bytes -class MooncakeEmbeddingPool: - """Persistent GPU buffer pool registered once with the Mooncake engine. +class EmbeddingPool: + """Persistent GPU buffer pool for received multimodal embeddings. 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). + a peer `release()`s a slot; `try_alloc()` is the non-blocking variant + for callers that re-poll (the zmq_to_scheduler tick). Each successful + alloc returns a slot_id that must be passed back to release() when the + consumer is done with the buffer. With `engine` set (mooncake), the + buffer is registered once so encoder RDMA writes land in pool slots. """ _ALIGN = 256 - def __init__(self, engine, gpu_id: int, size_bytes: int): - self.engine = engine + def __init__(self, gpu_id: int, size_bytes: int, engine=None): 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.engine = engine + if engine is not None: + 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 @@ -1408,10 +1521,43 @@ class MooncakeEmbeddingPool: 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}" + f"EmbeddingPool allocated: gpu={gpu_id}, " + f"size={size_bytes // (1024 * 1024)}MB, base=0x{self.base:x}, " + f"rdma_registered={engine is not None}" ) + def try_alloc(self, nbytes: int) -> Optional[Tuple[torch.Tensor, int, int]]: + """Non-blocking alloc: ``(tensor_view, gpu_addr, slot_id)``, or + ``None`` when the pool is currently full (oversize requests also get + ``None`` — callers detect those via ``size_bytes``).""" + aligned = (nbytes + self._ALIGN - 1) & ~(self._ALIGN - 1) + with self._lock: + return self._try_alloc_locked(nbytes, aligned) + + def try_stage( + self, parts: List[torch.Tensor] + ) -> Optional[Tuple[torch.Tensor, int]]: + """Copy CPU part tensors into one slot, packed in list order. + + Returns ``(slot_view, slot_id)``, or ``None`` when the pool is + currently full. Seam for future async staging (copy streams). + """ + alloc_result = self.try_alloc(sum(p.nbytes for p in parts)) + if alloc_result is None: + return None + slot_view, _, slot_id = alloc_result + offset = 0 + for part in parts: + nbytes = part.nbytes + slot_view[offset : offset + nbytes].copy_(part.flatten().view(torch.uint8)) + offset += nbytes + return slot_view, slot_id + + def release_on_gc(self, obj, slot_id: int) -> weakref.finalize: + """Release ``slot_id`` when ``obj`` is GC'd; the returned finalizer + can be called early to release now (at-most-once either way).""" + return weakref.finalize(obj, self.release, slot_id) + def alloc( self, nbytes: int, timeout: float = 60.0 ) -> Optional[Tuple[torch.Tensor, int, int]]: @@ -1430,7 +1576,7 @@ class MooncakeEmbeddingPool: """ if nbytes > self.size_bytes: logger.error( - f"MooncakeEmbeddingPool: requested {nbytes // (1024 * 1024)}MB " + f"EmbeddingPool: requested {nbytes // (1024 * 1024)}MB " f"exceeds pool capacity {self.size_bytes // (1024 * 1024)}MB. " f"Raise SGLANG_EMBEDDING_POOL_SIZE_MB." ) @@ -1447,7 +1593,7 @@ class MooncakeEmbeddingPool: inflight_mb = self._total_inflight // (1024 * 1024) cap_mb = self.size_bytes // (1024 * 1024) logger.warning( - f"MooncakeEmbeddingPool full: " + f"EmbeddingPool full: " f"{inflight_mb}/{cap_mb}MB in-flight across " f"{len(self._inflight)} requests; queueing a " f"{nbytes // (1024 * 1024)}MB request. Raise " @@ -1457,7 +1603,7 @@ class MooncakeEmbeddingPool: remaining = deadline - time.monotonic() if remaining <= 0: logger.error( - f"MooncakeEmbeddingPool alloc timed out after " + f"EmbeddingPool alloc timed out after " f"{timeout}s waiting for {nbytes // (1024 * 1024)}MB." ) return None @@ -1505,55 +1651,45 @@ class MooncakeEmbeddingPool: 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.""" +def _iter_part_ranges(embedding_data, dtype): + """Yield ``(part_idx, shape, byte_start, byte_end)`` for each non-None + part, packed in part order — the buffer layout shared by the encoder's + RDMA writes and EmbeddingPool.try_stage.""" 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 + offset = 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 + yield i, shape, offset, offset + nbytes + offset += nbytes + + +def _view_pool_buffer_by_modality(raw_buffer, embedding_data, dtype): + """Zero-copy view of raw_buffer as {modality: [total_tokens, hidden]}. + + Parts of the same modality are contiguous in raw_buffer (the encoder + writes them modality-outer), so each modality is one reshape of the byte + range — 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. + """ + # mod -> [byte_start, byte_end, total_tokens, hidden] + mod_info: Dict[Modality, List[int]] = {} + for i, shape, start, end in _iter_part_ranges(embedding_data, dtype): 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]] + mod_info[mod] = [start, end, 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 + assert info[1] == start, f"non-contiguous parts in modality {mod}" + info[1] = end 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() @@ -1604,8 +1740,8 @@ class MMReceiverBase(ABC): self.tp_group = tp_group self.nnodes = server_args.nnodes self.hostname = get_local_ip_auto() - self.waiting_list: List[WaitingImageRequest] = [] - self.waiting_by_rid: Dict[str, WaitingImageRequest] = {} + self.waiting_list: List[WaitingMMRequestBase] = [] + self.waiting_by_rid: Dict[str, WaitingMMRequestBase] = {} self.scheduler_embedding_port = None self.scheduler_recv_socket = None if ( @@ -1626,6 +1762,7 @@ class MMReceiverBase(ABC): self.scheduler = scheduler self.gpu_id = scheduler.ps.gpu_id if scheduler is not None else 0 self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() + self.embedding_pool = None self.model_type = ( getattr(hf_config, "model_type", "").lower() @@ -1647,23 +1784,39 @@ class MMReceiverBase(ABC): or get_exec().moe.mooncake_ib_device ), ) - self.embeddings_buffer = dict() - 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: try: - self.embedding_pool = MooncakeEmbeddingPool( - self.embeddings_engine, self.gpu_id, pool_mb * 1024 * 1024 + self.embedding_pool = EmbeddingPool( + self.gpu_id, + pool_mb * 1024 * 1024, + engine=self.embeddings_engine, ) except Exception: logger.exception( - "Failed to allocate MooncakeEmbeddingPool, " + "Failed to allocate EmbeddingPool, " "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": + # Unlike mooncake, do NOT apply the default pool size: explicitly + # set SGLANG_EMBEDDING_POOL_SIZE_MB= to bound received + # embeddings on GPU; unset/0 keeps the unpooled CPU receive. + if envs.SGLANG_EMBEDDING_POOL_SIZE_MB.is_set() and scheduler is not None: + pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get() + if pool_mb and pool_mb > 0: + try: + self.embedding_pool = EmbeddingPool( + self.gpu_id, pool_mb * 1024 * 1024 + ) + except Exception: + logger.exception( + "Failed to allocate EmbeddingPool, " + "falling back to unpooled receive" + ) + self.embedding_pool = None if hf_config is not None: self._init_mm_processor( server_args, @@ -1717,6 +1870,22 @@ class MMReceiverBase(ABC): def process_waiting_requests(self, recv_reqs): pass + def abort_waiting_requests(self, recv_req) -> None: + """Mark matching waiting requests FAIL and free their resources; the + next process_waiting_requests tick reports the abort through the + existing FAIL channel. AbortReq is broadcast, so every TP rank does + this and the status all-reduce stays consistent.""" + for waiting_req in self.waiting_list: + if not (recv_req.abort_all or waiting_req.rid.startswith(recv_req.rid)): + continue + if waiting_req.status in ( + WaitingMMRequestStatus.PENDING, + WaitingMMRequestStatus.SUCCESS, + ): + waiting_req._fail_and_release("Aborted by user", error_code=400) + waiting_req.release_resources() + logger.info(f"Abort waiting mm request. rid={waiting_req.rid}") + async def recv_mm_data( self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True ): @@ -1741,19 +1910,41 @@ class MMReceiverBase(ABC): f"modalities={modalities}, num_items={len(mm_data)}" ) send_time = time.monotonic() - asyncio.create_task( + encode_task = asyncio.create_task( self.encode( req_id, mm_data, embedding_port, "encode", - "send", encode_urls=encode_urls, ) ) - result = await asyncio.wait_for( - self._recv_mm_data(req_id, recv_socket, mm_processor, prompt), + # Parts stream onto the socket while other parts are still + # encoding, so receive concurrently with the dispatch; the dispatch + # result only matters for failing fast when nothing will arrive. + recv_task = asyncio.create_task( + self._recv_mm_data(req_id, recv_socket, mm_processor, prompt) + ) + done, _ = await asyncio.wait( + {encode_task, recv_task}, timeout=self.recv_timeout, + return_when=asyncio.FIRST_COMPLETED, + ) + if ( + done + and recv_task not in done + and ( + encode_task.exception() is not None or encode_task.result() is False + ) + ): + logger.warning( + f"[{req_id}] Encoder dispatch failed; skipping embedding wait" + ) + recv_task.cancel() + return None + result = await asyncio.wait_for( + recv_task, + timeout=self.recv_timeout - (time.monotonic() - send_time), ) elapsed = time.monotonic() - send_time logger.info(f"[{req_id}] Received embedding from E in {elapsed:.3f}s") @@ -1761,31 +1952,15 @@ class MMReceiverBase(ABC): except asyncio.TimeoutError: elapsed = time.monotonic() - send_time logger.warning(f"[{req_id}] Embedding recv timeout after {elapsed:.3f}s") - if req_id is not None: - self._cleanup_mooncake_buffer(req_id) return None - def _cleanup_mooncake_buffer(self, req_id): - if self.encoder_transfer_backend != "mooncake": - return - if not hasattr(self, "embeddings_buffer"): - return - embeddings = self.embeddings_buffer.pop(req_id, None) - if embeddings is None: - return - try: - self.embeddings_engine.deregister(embeddings.data_ptr()) - except Exception: - logger.exception( - "mooncake: failed to deregister buffer for req_id=%s", req_id - ) - async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt): + """zmq_to_tokenizer receive: embedding parts arrive as 2-frame ZMQ + messages on the tokenizer's PULL socket. (mooncake/zmq_to_scheduler lands on scheduler ranks via + WaitingMMRequest.)""" if req_id is None: return None - recv_embedding = None - recv_embedding_data: MultiModalEmbeddingData = None try: @@ -1799,55 +1974,33 @@ class MMReceiverBase(ABC): f"Encoder error for req_id={req_id}: {recv_obj.error_msg} " f"error_code={getattr(recv_obj, 'error_code', None)}" ) - self._cleanup_mooncake_buffer(req_id) return None logger.debug("recv_obj=%s", recv_obj) - # Extract original req_id from part_req_id - part_req_id = recv_obj.req_id - original_req_id = extract_original_req_id(part_req_id) - # Update recv_obj.req_id to original for aggregation - recv_obj.req_id = original_req_id - if self.encoder_transfer_backend == "zmq_to_tokenizer": - if len(parts) < 2: - logger.error( - "zmq_to_tokenizer expected 2-part message, got %d parts", - len(parts), - ) - return None - buffer = ( - parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] - ) - # Clone so we don't depend on ZMQ buffer after next recv. - recv_obj.embedding = ( - torch.frombuffer(buffer, dtype=recv_obj.dtype) - .reshape(recv_obj.shape) - .clone() - ) - if recv_embedding_data is None: - recv_embedding_data = MultiModalEmbeddingData.from_embedding_data( - recv_obj, model_type=self.model_type - ) - else: - recv_embedding_data.add(recv_obj) - - if self.encoder_transfer_backend == "mooncake": - if req_id not in self.embeddings_buffer: + # Normalize the part req_id to the original for aggregation. + recv_obj.req_id = extract_original_req_id(recv_obj.req_id) + if len(parts) < 2: logger.error( - "mooncake: embeddings_buffer missing req_id=%s", req_id + "zmq_to_tokenizer expected 2-part message, got %d parts", + len(parts), ) return None - raw_buffer = self.embeddings_buffer.pop(req_id) - self.embeddings_engine.deregister(raw_buffer.data_ptr()) - _slice_embedding_buffer(raw_buffer, recv_embedding_data, self.dtype) + buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] + # Clone so we don't depend on ZMQ buffer after next recv. + recv_obj.embedding = ( + torch.frombuffer(buffer, dtype=recv_obj.dtype) + .reshape(recv_obj.shape) + .clone() + ) + recv_embedding_data = _aggregate_embedding_part( + recv_embedding_data, recv_obj, self.model_type + ) recv_embedding = recv_embedding_data.get_embedding(is_concat=True) - - mm_inputs = mm_processor.get_mm_data( + return mm_processor.get_mm_data( prompt, recv_embedding, **recv_embedding_data.get_mm_extra_meta(), ) - return mm_inputs finally: recv_socket.close() @@ -1879,15 +2032,6 @@ class MMReceiverBase(ABC): # subprocess uses the same list when indexing encoder_idx. obj.encoder_urls = encode_urls - # 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=( @@ -1895,7 +2039,6 @@ class MMReceiverBase(ABC): mm_data, "encode", num_items_assigned, - None, encode_urls, time_stats_json, ), @@ -1914,7 +2057,7 @@ class MMReceiverBase(ABC): ) obj.need_wait_for_mm_inputs = False - def _sync_fail_info_across_tp(self, waiting_req: WaitingImageRequest) -> None: + def _sync_fail_info_across_tp(self, waiting_req: WaitingMMRequestBase) -> None: """Share encoder error fields across TP ranks before abort. The encoder sends ZMQ error signals to each TP rank's receive socket, @@ -2006,11 +2149,12 @@ class MMReceiverBase(ABC): current_time = time.time() local_status = [] for waiting_req in self.waiting_list: - if self.scheduler_recv_socket is None: - waiting_req._try_recv_mm_data() + # Per-request sockets receive here; shared-socket requests use this + # tick to retry pool staging after _drain_scheduler_embeddings(). + 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.status = WaitingMMRequestStatus.TIMEOUT + waiting_req.release_resources() waiting_req.close_recv_socket() local_status.append(waiting_req.status) @@ -2026,13 +2170,16 @@ class MMReceiverBase(ABC): abort_reqs = [] for i, waiting_req in enumerate(self.waiting_list): status_value = local_status[i].item() - if status_value == WaitingImageRequestStatus.SUCCESS: + if status_value == WaitingMMRequestStatus.SUCCESS: new_recv_reqs.append(waiting_req.recv_req) - elif status_value == WaitingImageRequestStatus.FAIL: + elif status_value == WaitingMMRequestStatus.FAIL: self._sync_fail_info_across_tp(waiting_req) logger.error( f"Waiting request {waiting_req.rid} failed: {waiting_req.error_msg} {waiting_req.error_code = }" ) + # A peer's FAIL can force-abort this locally PENDING/SUCCESS + # rank, so release any buffer/slot it still holds. + waiting_req.release_resources() abort_reqs.append( ( self.create_req(waiting_req.recv_req), @@ -2040,10 +2187,11 @@ class MMReceiverBase(ABC): waiting_req.error_code, ) ) - elif status_value == WaitingImageRequestStatus.TIMEOUT: + elif status_value == WaitingMMRequestStatus.TIMEOUT: logger.error( f"Timed out waiting for image embeddings for request {waiting_req.rid}" ) + waiting_req.release_resources() abort_reqs.append( ( self.create_req(waiting_req.recv_req), @@ -2051,7 +2199,7 @@ class MMReceiverBase(ABC): HTTPStatus.REQUEST_TIMEOUT, ) ) - else: # status_value == WaitingImageRequestStatus.PENDING + else: # status_value == WaitingMMRequestStatus.PENDING new_waiting.append(waiting_req) continue self.waiting_by_rid.pop(waiting_req.rid, None) @@ -2065,18 +2213,19 @@ class MMReceiverBase(ABC): mm_data, endpoint_encode, num_items_assigned, - embedding_port, encode_urls=None, time_stats_json=None, ): + # ``embedding_port`` is always None on this path: zmq_to_scheduler / + # mooncake ranks register their receive ports with the encoder later + # via /scheduler_receive_url, so the dispatch itself carries no port. try: asyncio.run( self.encode( req_id=req_id, mm_data=mm_data, - embedding_port=embedding_port, + embedding_port=None, endpoint_encode=endpoint_encode, - endpoint_send=None, num_items_assigned=num_items_assigned, encode_urls=encode_urls, time_stats_json=time_stats_json, @@ -2124,19 +2273,6 @@ class MMReceiverBase(ABC): req.tokenizer = self.scheduler.tokenizer return req - async def allocate_embedding_buffer(self, req_id, total_bytes): - logger.info( - f"Pre-allocating GPU buffer for mooncake RDMA: " - f"req_id={req_id}, size={total_bytes} bytes" - ) - embeddings = torch.empty(total_bytes, dtype=torch.uint8, device=self.gpu_id) - self.embeddings_engine.register( - embeddings.data_ptr(), - embeddings.nbytes, - ) - self.embeddings_buffer[req_id] = embeddings - return embeddings.data_ptr() - def _assign_items_by_modality( self, mm_data, encoder_num, random_shuffle=True ) -> Dict: @@ -2215,8 +2351,13 @@ class MMReceiverBase(ABC): if mm_items: mm_items = flatten_mm_items(mm_items) for mm_item in mm_items: + if mm_item is None: + continue + raw_url = to_raw_url(mm_item) + if raw_url is None: + continue entry = { - "url": to_raw_url(mm_item), + "url": raw_url, "modality": modality, } entry.update( @@ -2277,44 +2418,15 @@ class MMReceiverHTTP(MMReceiverBase): if self.encoder_transfer_backend == "mooncake": return self._process_waiting_requests( recv_reqs, - WaitingImageRDMARequest, + WaitingRDMARequest, embeddings_engine=self.embeddings_engine, dtype=self.dtype, gpu_id=self.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() - encoder_label = encode_requests[i].get( - "encoder_url", f"idx={encode_requests[i].get('encoder_idx')}" - ) - logger.error( - f"Encoder HTTP request timeout ({timeout_val}s) for req_id={req_id} " - f"(request {i}), " - f"encoder={encoder_label}" - ) - 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 + return self._process_waiting_requests( + recv_reqs, WaitingZmqRequest, embedding_pool=self.embedding_pool + ) async def encode( self, @@ -2322,7 +2434,6 @@ class MMReceiverHTTP(MMReceiverBase): mm_data, embedding_port, endpoint_encode, - endpoint_send, num_items_assigned=None, encode_urls=None, time_stats_json=None, @@ -2399,49 +2510,16 @@ class MMReceiverHTTP(MMReceiverBase): ] responses = await asyncio.gather(*tasks, return_exceptions=True) - - if not await self._check_encoder_responses( - responses, encode_requests, req_id - ): - return - response_json_list_unsort = [ - await response.json() for response in responses - ] - - # zmq backend: return is None - if None in response_json_list_unsort: - return - - # mooncake backend: send bootstrap info - - embedding_size_list_sort, response_json_list_sort, total_embedding_bytes = ( - _sort_responses_and_compute_total_bytes( - response_json_list_unsort, total_num_parts + # Dispatch only. The embedding never comes back through this call: + # zmq_to_tokenizer is pushed to our PULL socket during /encode, + # zmq_to_scheduler to the ports its ranks registered, and mooncake + # by RDMA once those ranks have pulled sizes and driven /send. + return ( + await _extract_encoder_error( + responses, "HTTP request", f"req_id={req_id}", encode_requests ) + is None ) - offset = 0 - metadata_tasks = [] - buffer_address = await self.allocate_embedding_buffer( - req_id, - total_embedding_bytes, - ) - for idx in range(len(tasks)): - response_json = response_json_list_sort[idx] - buffer_address_adjust = offset + buffer_address - response_json.update( - { - "session_id": self.embeddings_engine.session_id, - "buffer_address": buffer_address_adjust, - } - ) - metadata_tasks.append( - session.post( - f"{effective_urls[response_json['encoder_idx']]}/{endpoint_send}", - json=response_json, - ) - ) - offset += embedding_size_list_sort[idx] - await asyncio.gather(*metadata_tasks) class MMReceiverGrpc(MMReceiverBase): @@ -2456,6 +2534,13 @@ class MMReceiverGrpc(MMReceiverBase): scheduler: Optional["Scheduler"] = None, encode_urls: Optional[List[str]] = None, ): + if get_disagg().encoder_transfer_backend == "mooncake": + # The RDMA receive path (WaitingRDMARequest + /meta + /send) only + # exists for HTTP encoders; gRPC has no RDMA-capable receive. + raise NotImplementedError( + "mooncake encoder_transfer_backend requires HTTP encoders; " + "use zmq_to_scheduler / zmq_to_tokenizer with gRPC." + ) super().__init__( server_args, dtype=dtype, @@ -2475,9 +2560,9 @@ class MMReceiverGrpc(MMReceiverBase): self.send_encode_request(encode_req) return encode_req - # For zmq_to_scheduler and mooncake + # For zmq_to_scheduler def process_waiting_requests(self, recv_reqs): - return self._process_waiting_requests(recv_reqs, WaitingImageRequestGrpc) + return self._process_waiting_requests(recv_reqs, WaitingZmqRequestGrpc) async def encode( self, @@ -2485,7 +2570,6 @@ class MMReceiverGrpc(MMReceiverBase): mm_data, embedding_port, endpoint_encode, - endpoint_send, num_items_assigned=None, encode_urls=None, ): @@ -2548,55 +2632,7 @@ class MMReceiverGrpc(MMReceiverBase): ) for encode_request in encode_requests ] - grpc_responses = await asyncio.gather(*grpc_tasks) - response_json_unsorted = [] - for encode_request, response in zip(encode_requests, grpc_responses): - if self.encoder_transfer_backend == "zmq_to_scheduler": - response_json_unsorted.append(None) - continue - response_json_unsorted.append( - { - "req_id": encode_request["req_id"], - "prefill_host": encode_request["prefill_host"], - "embedding_port": encode_request["embedding_port"], - "encoder_idx": encode_request["encoder_idx"], - "part_idx": encode_request["part_idx"], - "embedding_size": response.embedding_size, - "embedding_len": response.embedding_len, - "embedding_dim": response.embedding_dim, - } - ) - - if None in response_json_unsorted: - return - - 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, - total_embedding_bytes, - ) - grpc_metadata_tasks = [] - for response_json in response_json_sorted: - response_json.update( - { - "session_id": self.embeddings_engine.session_id, - "buffer_address": offset + buffer_address, - } - ) - grpc_metadata_tasks.append( - asyncio.to_thread( - _grpc_send_request, - _grpc_target(effective_urls[response_json["encoder_idx"]]), - response_json, - ) - ) - offset += embedding_size_by_part[response_json["part_idx"]] - - if grpc_metadata_tasks: - await asyncio.gather(*grpc_metadata_tasks) + await asyncio.gather(*grpc_tasks) def _validate_transport_mode(transport_mode: str, encoder_urls): diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py new file mode 100644 index 000000000..ba47db066 --- /dev/null +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -0,0 +1,1643 @@ +"""Protocol-neutral runtime for the EPD encoder server. + +The current runtime keeps :class:`EncoderScheduler` and the rank-0 +:class:`server.MMEncoder` in the same process. It also owns HTTP's +existing DP replica processes and dispatch plumbing so another transport can +reuse that backend topology without importing the HTTP server. +""" + +import asyncio +import atexit +import contextlib +import logging +import multiprocessing as mp +import os +import time +import traceback +import uuid +from collections import defaultdict +from dataclasses import dataclass +from http import HTTPStatus +from typing import Dict, List, Optional, Set, Tuple + +import zmq +import zmq.asyncio + +import sglang.srt.disaggregation.encoder.server as server_module +from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX +from sglang.srt.disaggregation.encoder.server import ( + ENCODER_MAX_BATCH_SIZE, + ENCODER_MAX_BATCH_SIZE_EXPLICIT, + EncoderProfiler, + MMEncoder, + MMError, + launch_encoder, +) +from sglang.srt.environ import envs +from sglang.srt.managers.io_struct import ( + ProfileReq, + ProfileReqType, + async_sock_recv, + async_sock_send, + sock_send, + wrap_as_pickle, +) +from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.observability.metrics_collector import EncoderMetricsCollector +from sglang.srt.observability.req_time_stats import EncoderReqTimeStats +from sglang.srt.observability.trace import ( + process_tracing_init, + trace_set_thread_info, +) +from sglang.srt.runtime_context import ( + configured_tp_size, + get_observability, + get_parallel, + get_serving, +) +from sglang.srt.server_args import PortArgs, ServerArgs +from sglang.srt.utils import configure_logger, random_uuid, set_prometheus_multiproc_dir +from sglang.srt.utils.common import maybe_reindex_device_id +from sglang.srt.utils.network import NetworkAddress, get_free_port, get_zmq_socket + +logger = logging.getLogger(__name__) + + +class PendingRequest: + __slots__ = ("request", "future", "submit_time") + + def __init__(self, request: dict, loop: asyncio.AbstractEventLoop): + self.request = request + self.future: asyncio.Future = loop.create_future() + self.submit_time = time.time() + + +# VIDEO excluded: per-video preprocess kwargs (do_sample_frames, video_metadata) +# vary per request and can't merge into one HF processor call. +_BATCHABLE_MODALITIES = {Modality.IMAGE, Modality.AUDIO} +_KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE = 2 + + +def _resolve_encoder_batch_policy( + model_type: str, + configured_max_batch_size: int, + max_batch_size_is_explicit: bool, +) -> Tuple[int, bool]: + """Return effective batch size and same-turn coalescing policy.""" + max_batch_size = max(1, int(configured_max_batch_size)) + coalesce_same_turn = model_type == "kimi_k3" + if coalesce_same_turn and not max_batch_size_is_explicit: + max_batch_size = min(max_batch_size, _KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE) + return max_batch_size, coalesce_same_turn + + +# Minimal 32x32 black PNG for health check dummy encode +MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" + +# Minimal WAV: 16kHz mono 16-bit PCM, 160 samples (0.01s) of silence +MINIMUM_WAV_SILENCE_BASE64 = "UklGRmQBAABXQVZFZm10IBAAAAABAAEAgD4AAAB9AAACABAAZGF0YUABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==" + + +class EncoderScheduler: + """Aggregate concurrent /encode requests into bounded image/audio batches.""" + + def __init__( + self, + encoder: "MMEncoder", + send_sockets: List[zmq.Socket], + max_batch_size: int, + coalesce_same_turn: bool = False, + request_timeout: float = server_module.ENCODER_REQ_TIMEOUT, + ): + self.encoder = encoder + self.send_sockets = send_sockets + self.max_batch_size = max(1, int(max_batch_size)) + self.coalesce_same_turn = bool(coalesce_same_turn) + self.request_timeout = max(1.0, float(request_timeout)) + self.pending_queue: asyncio.Queue[PendingRequest] = asyncio.Queue() + self._worker_task: Optional[asyncio.Task] = None + + def start(self) -> None: + if self._worker_task is None: + self._worker_task = asyncio.create_task(self._batch_worker()) + logger.info( + "EncoderScheduler started with " + f"max_batch_size={self.max_batch_size}, " + f"coalesce_same_turn={self.coalesce_same_turn}" + ) + + async def stop(self) -> None: + if self._worker_task is not None: + self._worker_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._worker_task + self._worker_task = None + # Reject any requests still queued so their HTTP handlers don't hang. + while True: + try: + pending = self.pending_queue.get_nowait() + except asyncio.QueueEmpty: + break + if not pending.future.done(): + pending.future.set_exception(RuntimeError("EncoderScheduler stopped")) + + async def submit(self, request: dict) -> Tuple: + pending = PendingRequest(request, asyncio.get_running_loop()) + await self.pending_queue.put(pending) + try: + return await asyncio.wait_for(pending.future, timeout=self.request_timeout) + except asyncio.TimeoutError: + if not pending.future.done(): + pending.future.cancel() + req_id = request.get("req_id") + # Free anything the abandoned batch may still stage for this rid. + await self.encoder.release_request(req_id) + logger.error( + f"EncoderScheduler.submit timed out after {self.request_timeout}s " + f"for req_id={req_id}" + ) + raise + + async def _collect_batch(self) -> List[PendingRequest]: + batch = [await self.pending_queue.get()] + first_modality = Modality.from_str(batch[0].request.get("modality", "image")) + should_yield = ( + self.coalesce_same_turn + and self.max_batch_size > 1 + and first_modality in _BATCHABLE_MODALITIES + ) + if should_yield: + # Let HTTP handlers that arrived in the same event-loop turn enqueue + # before dispatch. Unlike a fixed sleep, this adds no millisecond-scale + # tax to an isolated request. + await asyncio.sleep(0) + while len(batch) < self.max_batch_size: + try: + batch.append(self.pending_queue.get_nowait()) + except asyncio.QueueEmpty: + break + return batch + + async def _batch_worker(self) -> None: + while True: + batch: List[PendingRequest] = [] + try: + batch = await self._collect_batch() + groups: Dict[Modality, List[PendingRequest]] = defaultdict(list) + for p in batch: + groups[ + Modality.from_str(p.request.get("modality", "image")) + ].append(p) + for modality, group in groups.items(): + await self._dispatch_group(group, modality) + except asyncio.CancelledError: + for p in batch: + if not p.future.done(): + p.future.set_exception(RuntimeError("EncoderScheduler stopped")) + raise + except Exception as e: + logger.error( + f"Error in EncoderScheduler batch worker: {e}", exc_info=True + ) + for p in batch: + if not p.future.done(): + p.future.set_exception(e) + + @staticmethod + def _validate_request_shape(req: dict) -> Optional[str]: + # Cheap pre-broadcast checks: shape errors that don't require running + # the HF processor. Once a request reaches TP workers they enter + # batch_encode and expect to join its collectives — a malformed batch + # that makes rank-0 bail mid-flight would deadlock the workers. + if not isinstance(req, dict): + return f"request is not a dict: {type(req).__name__}" + if not req.get("req_id"): + return "missing req_id" + if not req.get("mm_items"): + return "missing or empty mm_items" + if "num_parts" not in req or "part_idx" not in req: + return "missing num_parts / part_idx" + h = req.get("hashes") + if h is not None and not isinstance(h, (list, tuple, str, int, bytes)): + return f"hashes must be list/scalar, got {type(h).__name__}" + return None + + async def _dispatch_group( + self, group: List[PendingRequest], modality: Modality + ) -> None: + # A request may time out while queued. Never start work that no caller + # can observe, or its eventual staged embedding would have no owner. + group = [pending for pending in group if not pending.future.done()] + if not group: + return + + # Video can't fuse (per-video preprocess kwargs vary). + if modality not in _BATCHABLE_MODALITIES: + await self._dispatch_per_request(group, modality) + return + + # Drop structurally-bad requests before broadcasting; otherwise TP + # workers would join batch_encode collectives that rank-0 has already + # abandoned. + valid: List[PendingRequest] = [] + for p in group: + err = self._validate_request_shape(p.request) + if err is None: + valid.append(p) + continue + logger.error(f"Dropping req_id={p.request.get('req_id')} from batch: {err}") + if not p.future.done(): + p.future.set_exception(server_module.BadRequestError(err)) + if not valid: + return + group = valid + + requests = [p.request for p in group] + start = time.time() + modality_str = modality.name.lower() + if server_module.encoder_metrics_collector is not None: + for p in group: + server_module.encoder_metrics_collector.observe_queue_wait( + max(0.0, start - p.submit_time), modality=modality_str + ) + try: + # The scheduler is the sole owner of batched dispatch order. Keep + # the collective broadcast and rank-0 execution under the same + # lock, while allowing concurrent HTTP handlers to enqueue before + # waiting on their individual futures. + async with self.encoder.encode_dispatch_lock: + for sock in self.send_sockets: + sock_send( + sock, + wrap_as_pickle( + { + "type": "batch_encode", + "modality": modality.name, + "requests": requests, + "enter_time": start, + } + ), + ) + + logger.info( + f"Dispatching batch of {len(group)} {modality.name} requests" + ) + results = await self.encoder.batch_encode(requests, modality) + if len(group) > 1: + logger.info( + f"Batch of {len(group)} {modality.name} requests completed in " + f"{(time.time() - start) * 1000:.1f}ms" + ) + except Exception as e: + # batch_encode normally catches and returns errors via _stage_errors. + # If it raised, rank-0 may have skipped a collective broadcast, leaving + # TP workers stuck. Don't try to recover — fail every pending future + # and let the client retry. Re-broadcasting would risk a deadlock. + logger.error(f"batch_encode raised: {e}", exc_info=True) + for p in group: + if not p.future.done(): + p.future.set_exception(e) + return + + if len(results) != len(group): + err = RuntimeError( + f"batch_encode returned {len(results)} results for {len(group)} requests" + ) + logger.error(str(err)) + for p in group: + if not p.future.done(): + p.future.set_exception(err) + return + + for p, result in zip(group, results): + if not p.future.done(): + p.future.set_result(result) + + async def _dispatch_per_request( + self, + group: List[PendingRequest], + modality: Modality, + ) -> None: + modality_str = modality.name.lower() + for p in group: + if p.future.done(): + continue + req = p.request + try: + start = time.time() + if server_module.encoder_metrics_collector is not None: + server_module.encoder_metrics_collector.observe_queue_wait( + max(0.0, start - p.submit_time), modality=modality_str + ) + for sock in self.send_sockets: + sock_send(sock, wrap_as_pickle(req)) + result = await self.encoder.encode( + 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"), + ) + if not p.future.done(): + p.future.set_result(result) + except Exception as e: + logger.error( + f"Per-request encode failed for req_id={req.get('req_id')}: {e}" + ) + if not p.future.done(): + p.future.set_exception(e) + + +@dataclass +class EncoderRuntime: + """Current non-DP backend runtime. + + The Scheduler and rank-0 MMEncoder remain colocated. TP followers use the + existing ZMQ control path and are intentionally not split behind a new + Scheduler/Worker IPC contract in this phase. + """ + + encoder: MMEncoder + scheduler: EncoderScheduler + send_sockets: List[zmq.Socket] + zmq_context: zmq.Context + tp_processes: List[mp.Process] + + def start(self) -> None: + self.scheduler.start() + + async def stop(self) -> None: + # Preserve the existing lifecycle: Uvicorn stops the Scheduler, while + # daemon TP followers exit with their parent process. + await self.scheduler.stop() + + +class DPDispatcher: + """Routes encode requests across DP ranks by least-pending count.""" + + def __init__( + self, + dp_size: int, + dispatch_sockets: List, + result_socket, + worker_processes: List[mp.Process], + enable_metrics: bool = False, + labels: Optional[Dict[str, str]] = None, + ): + self.dp_size = dp_size + self.dispatch_sockets = dispatch_sockets + self.result_socket = result_socket + self.worker_processes = worker_processes + # Key = req_id for encode/broadcast, or a per-control-request key for + # Mooncake metadata waits, sends, and destination registrations. + self.pending_futures: List[Dict[str, asyncio.Future]] = [ + {} for _ in range(dp_size) + ] + self.req_id_to_rank: Dict[str, int] = {} + self._mapping_condition = asyncio.Condition() + self._rr_counter = 0 + self._broadcast_counter = 0 + self._metadata_counter = 0 + self._dead_ranks: Set[int] = set() + # req_id -> monotonic ts a mooncake mapping has waited for its /send. + self._pending_send_at: Dict[str, float] = {} + # Set when _result_listener gives up; makes alive_ranks report empty. + self._listener_failed = False + + # Prometheus gauge: pending requests per DP rank. Lives in the main + # process (the dispatcher), unlike the per-worker EncoderMetricsCollector. + self.labels = dict(labels or {}) + self.pending_gauge = None + if enable_metrics: + from prometheus_client import Gauge + + self.pending_gauge = Gauge( + name="sglang:encoder_dp_pending_requests", + documentation="Number of pending requests per encoder DP rank.", + labelnames=list(self.labels.keys()) + ["dp_rank"], + multiprocess_mode="mostrecent", + ) + + @property + def pending_counts(self) -> List[int]: + return [len(d) for d in self.pending_futures] + + def _update_pending_gauge(self) -> None: + """Push current pending counts to the Prometheus gauge (absolute set).""" + if self.pending_gauge is not None: + for i, c in enumerate(self.pending_counts): + self.pending_gauge.labels(**self.labels, dp_rank=str(i)).set(c) + + @property + def alive_ranks(self) -> List[int]: + # Empty if the result listener died; else ranks not marked dead. + if self._listener_failed: + return [] + return [r for r in range(self.dp_size) if r not in self._dead_ranks] + + @property + def all_ranks_alive(self) -> bool: + # Strict health (only /health uses this); routing still degrades. + return len(self.alive_ranks) == self.dp_size + + def start(self) -> None: + logger.info(f"DP dispatcher started: {self.dp_size} ranks (all remote)") + asyncio.create_task(self._result_listener()) + asyncio.create_task(self._worker_watchdog()) + asyncio.create_task(self._cleanup_stale_mappings()) + + def _drop_pending_and_mapping(self, rank: int, req_id: str) -> None: + # dispatch / broadcast failure: no follow-up /send expected. + self.pending_futures[rank].pop(req_id, None) + self.req_id_to_rank.pop(req_id, None) + self._update_pending_gauge() + + @staticmethod + def _send_req_key(req_id: str, request: dict) -> str: + """One in-flight /send future per decoder TP rank, keyed by the rank's + ZMQ-ack endpoint; a retry from the same rank reuses the key.""" + endpoint = NetworkAddress( + request["prefill_host"], request["embedding_port"] + ).to_host_port_str() + return f"{req_id}_send_{endpoint}" + + @staticmethod + def _register_req_key(req_id: str, request: dict) -> str: + return f"{req_id}_register_{request['receive_url']}" + + def _metadata_req_key(self, req_id: str) -> str: + key = f"{req_id}_metadata_{self._metadata_counter}" + self._metadata_counter += 1 + return key + + @staticmethod + def _pending_req_info(key: str) -> Tuple[str, str]: + marker_index, dp_type = max( + ( + (key.rfind("_send_"), "send"), + (key.rfind("_metadata_"), "wait_metadata"), + (key.rfind("_register_"), "register_destinations"), + ), + key=lambda item: item[0], + ) + if marker_index >= 0: + return key[:marker_index], dp_type + return key, "encode" + + def _fail_pending_for_rank(self, rank: int, reason: str, error_type: str) -> None: + # Resolve a rank's outstanding futures with 503 so awaiters don't hang. + pending = self.pending_futures[rank] + for key, future in list(pending.items()): + if not future.done(): + req_id, dp_type = self._pending_req_info(key) + future.set_result( + { + "req_id": req_id, + "_dp_type": dp_type, + "content": None, + "_error": reason, + "_error_type": error_type, + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + ) + pending.pop(key, None) + self._update_pending_gauge() + + def _fail_all_pending(self, reason: str, error_type: str) -> None: + for rank in range(self.dp_size): + self._fail_pending_for_rank(rank, reason, error_type) + self.req_id_to_rank.clear() + self._pending_send_at.clear() + + @staticmethod + def _timeout_envelope(req_id: str, dp_type: str, reason: str) -> dict: + return { + "req_id": req_id, + "_dp_type": dp_type, + "content": None, + "_error": reason, + "_error_type": "TimeoutError", + "_error_code": int(HTTPStatus.GATEWAY_TIMEOUT), + } + + async def dispatch(self, request: dict) -> dict: + counts = self.pending_counts + # Skip ranks whose worker process has died. + alive_ranks = self.alive_ranks + if not alive_ranks: + raise server_module.MMError( + "All encoder DP workers are dead.", + code=HTTPStatus.SERVICE_UNAVAILABLE, + ) + min_p = min(counts[r] for r in alive_ranks) + candidates = [r for r in alive_ranks if counts[r] == min_p] + rank = candidates[self._rr_counter % len(candidates)] + self._rr_counter += 1 + req_id = request["req_id"] + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][req_id] = future + self._update_pending_gauge() + logger.info( + f"MM-Encoder DP dispatch: req_id={req_id}, " + f"modality={request.get('modality', 'image')}, " + f"dp_rank={rank}, pending={self.pending_counts}" + ) + + try: + # Do not let concurrent metadata/destination control requests route + # to this worker until the corresponding encode is enqueued first. + # They share one PUSH socket, so releasing the condition after send + # preserves the required order. + async with self._mapping_condition: + self.req_id_to_rank[req_id] = rank + try: + await async_sock_send( + self.dispatch_sockets[rank], wrap_as_pickle(request) + ) + except BaseException: + self._drop_pending_and_mapping(rank, req_id) + self._mapping_condition.notify_all() + raise + self._mapping_condition.notify_all() + # An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip + # the watchdog, so bound the wait explicitly. + return await asyncio.wait_for( + future, timeout=server_module.ENCODER_REQ_TIMEOUT + ) + except asyncio.TimeoutError: + self._drop_pending_and_mapping(rank, req_id) + return self._timeout_envelope( + req_id, + "encode", + f"Encoder DP rank={rank} timed out after {server_module.ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self._drop_pending_and_mapping(rank, req_id) + raise + + async def dispatch_register_destinations(self, request: dict) -> dict: + """Route a scheduler receive URL to the DP worker owning ``req_id``.""" + req_id = request["req_id"] + deadline = time.monotonic() + min(5.0, server_module.ENCODER_REQ_TIMEOUT) + async with self._mapping_condition: + while req_id not in self.req_id_to_rank: + remaining = deadline - time.monotonic() + if remaining <= 0: + return { + "req_id": req_id, + "_error": f"Unknown req_id: {req_id}", + "_error_code": int(HTTPStatus.NOT_FOUND), + } + try: + await asyncio.wait_for( + self._mapping_condition.wait(), timeout=remaining + ) + except asyncio.TimeoutError: + return { + "req_id": req_id, + "_error": f"Unknown req_id: {req_id}", + "_error_code": int(HTTPStatus.NOT_FOUND), + } + rank = self.req_id_to_rank[req_id] + + if rank in self._dead_ranks: + return { + "req_id": req_id, + "_error": f"DP worker rank={rank} died before URL registration", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + + key = self._register_req_key(req_id, request) + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][key] = future + self._update_pending_gauge() + worker_request = { + **request, + "_dp_type": "register_destinations", + "_dp_register_key": key, + } + try: + await async_sock_send( + self.dispatch_sockets[rank], wrap_as_pickle(worker_request) + ) + return await asyncio.wait_for( + future, timeout=server_module.ENCODER_REQ_TIMEOUT + ) + except asyncio.TimeoutError: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + return self._timeout_envelope( + req_id, + "register_destinations", + f"Encoder DP rank={rank} URL registration timed out after " + f"{server_module.ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + raise + + async def dispatch_wait_metadata(self, request: dict) -> dict: + """Wait for metadata in the DP worker that owns ``req_id``. + + The worker-local registry publishes preprocessing metadata before the + encoder forward. Keeping the wait in that process preserves Mooncake's + early landing-buffer allocation in DP mode. + """ + req_id = request["req_id"] + deadline = time.monotonic() + min(5.0, server_module.ENCODER_REQ_TIMEOUT) + async with self._mapping_condition: + while req_id not in self.req_id_to_rank: + remaining = deadline - time.monotonic() + if remaining <= 0: + return { + "req_id": req_id, + "_error": f"Unknown req_id: {req_id}", + "_error_code": int(HTTPStatus.NOT_FOUND), + } + try: + await asyncio.wait_for( + self._mapping_condition.wait(), timeout=remaining + ) + except asyncio.TimeoutError: + return { + "req_id": req_id, + "_error": f"Unknown req_id: {req_id}", + "_error_code": int(HTTPStatus.NOT_FOUND), + } + rank = self.req_id_to_rank[req_id] + + if rank in self._dead_ranks: + return { + "req_id": req_id, + "_error": f"DP worker rank={rank} died before metadata became ready", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + + key = self._metadata_req_key(req_id) + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][key] = future + self._update_pending_gauge() + worker_request = { + **request, + "_dp_type": "wait_metadata", + "_dp_metadata_key": key, + } + try: + await async_sock_send( + self.dispatch_sockets[rank], wrap_as_pickle(worker_request) + ) + return await asyncio.wait_for( + future, timeout=server_module.ENCODER_REQ_TIMEOUT + ) + except asyncio.TimeoutError: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + return self._timeout_envelope( + req_id, + "wait_metadata", + f"Encoder DP rank={rank} metadata wait timed out after " + f"{server_module.ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + raise + + async def dispatch_send(self, request: dict) -> dict: + req_id = request["req_id"] + # /send arrived → stop tracking it for stale-mapping GC. + self._pending_send_at.pop(req_id, None) + if self._listener_failed: + return { + "req_id": req_id, + "_error": "encoder DP result listener stopped; cannot route /send", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + rank = self.req_id_to_rank.get(req_id) + if rank is None: + logger.warning( + f"MM-Encoder dispatch_send: unknown req_id={req_id}, " + f"cannot route to worker" + ) + return {"req_id": req_id, "_error": f"Unknown req_id: {req_id}"} + if rank in self._dead_ranks: + # Worker died between encode and /send; embedding is gone. + self.req_id_to_rank.pop(req_id, None) + return { + "req_id": req_id, + "_error": f"DP worker rank={rank} died before /send for req_id={req_id}", + "_error_code": int(HTTPStatus.SERVICE_UNAVAILABLE), + } + key = self._send_req_key(req_id, request) + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][key] = future + self._update_pending_gauge() + request["_dp_type"] = "send" + request["_dp_send_key"] = key + logger.info( + f"MM-Encoder DP dispatch_send: req_id={req_id}, " + f"dp_rank={rank}, send_key={key}, pending={self.pending_counts}" + ) + try: + await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request)) + return await asyncio.wait_for( + future, timeout=server_module.ENCODER_REQ_TIMEOUT + ) + except asyncio.TimeoutError: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + # Siblings still route via req_id_to_rank; the stale sweep evicts it. + self._pending_send_at[req_id] = time.monotonic() + return self._timeout_envelope( + req_id, + "send", + f"Encoder DP rank={rank} /send timed out after {server_module.ENCODER_REQ_TIMEOUT}s", + ) + except BaseException: + self.pending_futures[rank].pop(key, None) + self._update_pending_gauge() + self._pending_send_at[req_id] = time.monotonic() + raise + + async def broadcast( + self, request: dict, timeout: Optional[float] = None + ) -> List[dict]: + # Skip dead ranks: a PUSH to a gone worker would just buffer and then + # surface as a spurious per-rank timeout. All dead → 503 (same as + # dispatch), which the profile endpoints turn into an HTTP error. + eff_timeout = ( + timeout if timeout is not None else server_module.ENCODER_REQ_TIMEOUT + ) + alive_ranks = self.alive_ranks + if not alive_ranks: + raise server_module.MMError( + "All encoder DP workers are dead.", + code=HTTPStatus.SERVICE_UNAVAILABLE, + ) + batch_id = self._broadcast_counter + self._broadcast_counter += 1 + rank_keys: List[Tuple[int, str]] = [] + futures: List[asyncio.Future] = [] + dp_type = request.get("_dp_type", "unknown") + try: + for rank in alive_ranks: + req_id = f"_broadcast_{batch_id}_{rank}" + future = asyncio.get_running_loop().create_future() + self.pending_futures[rank][req_id] = future + self.req_id_to_rank[req_id] = rank + rank_keys.append((rank, req_id)) + request_copy = {**request, "req_id": req_id} + await async_sock_send( + self.dispatch_sockets[rank], wrap_as_pickle(request_copy) + ) + futures.append(future) + # Concurrent wait → total bounded by eff_timeout, not + # dp_size × eff_timeout. + outcomes = await asyncio.gather( + *(asyncio.wait_for(fut, timeout=eff_timeout) for fut in futures), + return_exceptions=True, + ) + results: List[dict] = [] + for (rank, req_id), outcome in zip(rank_keys, outcomes): + if isinstance(outcome, asyncio.TimeoutError): + self._drop_pending_and_mapping(rank, req_id) + results.append( + self._timeout_envelope( + req_id, + dp_type, + f"Encoder DP rank={rank} broadcast timed out " + f"after {eff_timeout}s", + ) + ) + elif isinstance(outcome, BaseException): + self._drop_pending_and_mapping(rank, req_id) + raise outcome + else: + results.append(outcome) + return results + except BaseException: + for rank, req_id in rank_keys: + self._drop_pending_and_mapping(rank, req_id) + raise + + async def _worker_watchdog(self) -> None: + # proc.sentinel becomes readable on process exit; fail this rank's + # pending futures so awaiters don't hang on a dead worker. + loop = asyncio.get_running_loop() + watch: Dict[int, asyncio.Future] = {} + for rank, proc in enumerate(self.worker_processes): + fut: asyncio.Future = loop.create_future() + + # add_reader is level-triggered, so remove_reader inside the + # callback to avoid spinning every loop iteration. + def _on_exit(r=rank, f=fut, p=proc, lp=loop): + try: + lp.remove_reader(p.sentinel) + except (ValueError, OSError): + pass + if not f.done(): + f.set_result(r) + + try: + loop.add_reader(proc.sentinel, _on_exit) + except (ValueError, OSError): + continue + watch[rank] = fut + + while watch: + done, _ = await asyncio.wait( + watch.values(), return_when=asyncio.FIRST_COMPLETED + ) + for fut in done: + rank = fut.result() + proc = self.worker_processes[rank] + logger.error( + f"DP worker rank={rank} (pid={proc.pid}) exited " + f"with code={proc.exitcode}; failing pending requests" + ) + self._dead_ranks.add(rank) + reason = f"DP worker rank={rank} died (exitcode={proc.exitcode})" + self._fail_pending_for_rank(rank, reason, "WorkerDied") + self.req_id_to_rank = { + r: rk for r, rk in self.req_id_to_rank.items() if rk != rank + } + watch.pop(rank, None) + + async def _result_listener(self) -> None: + # Bounded back-off + give-up so a torn-down context exits in ~3s + # rather than spinning forever on recv errors. + consecutive_errors = 0 + while True: + try: + msg = await async_sock_recv(self.result_socket) + consecutive_errors = 0 + except asyncio.CancelledError: + raise + except Exception: + consecutive_errors += 1 + logger.error("_result_listener recv error", exc_info=True) + if consecutive_errors >= 30: + logger.error( + "_result_listener giving up after 30 consecutive errors" + ) + self._listener_failed = True + self._fail_all_pending( + "encoder DP result listener stopped after repeated " + "recv errors", + "ResultListenerStopped", + ) + return + await asyncio.sleep(min(0.1 * consecutive_errors, 1.0)) + continue + req_id = msg.get("req_id", "") + dp_type = msg.get("_dp_type", "encode") + if dp_type == "send": + key = msg.get("_dp_send_key") + if key is None: + # Workers always echo the key; never fall back to req_id, + # which would wrongly resolve the encode future. + logger.warning( + f"_result_listener: send envelope without _dp_send_key " + f"for req_id={req_id}, dropping" + ) + continue + elif dp_type == "register_destinations": + key = msg.get("_dp_register_key") + if key is None: + logger.warning( + f"_result_listener: URL registration envelope without " + f"_dp_register_key for req_id={req_id}, dropping" + ) + continue + elif dp_type == "wait_metadata": + key = msg.get("_dp_metadata_key") + if key is None: + logger.warning( + f"_result_listener: metadata envelope without " + f"_dp_metadata_key for req_id={req_id}, dropping" + ) + continue + else: + key = req_id + rank = self.req_id_to_rank.get(req_id) + if rank is None or key not in self.pending_futures[rank]: + logger.warning( + f"_result_listener: no pending future for " + f"req_id={req_id}, dp_type={dp_type}, key={key}, dropping" + ) + continue + future = self.pending_futures[rank].pop(key) + self._update_pending_gauge() + # Each decoder TP rank sends against the same req_id, so dropping the + # mapping on the first /send leaves the siblings unroutable. Refresh + # the timestamp instead and let the stale-mapping sweep evict it. + register_prefix = f"{req_id}_register_" + has_pending_registration = any( + pending_key.startswith(register_prefix) + for pending_key in self.pending_futures[rank] + ) + metadata_prefix = f"{req_id}_metadata_" + has_pending_metadata = any( + pending_key.startswith(metadata_prefix) + for pending_key in self.pending_futures[rank] + ) + keep_mapping = ( + dp_type in ("send", "register_destinations", "wait_metadata") + or (dp_type == "encode" and msg.get("content") is not None) + or has_pending_registration + or has_pending_metadata + ) + if dp_type == "send" or ( + dp_type == "encode" and msg.get("content") is not None + ): + self._pending_send_at[req_id] = time.monotonic() + if not keep_mapping: + self.req_id_to_rank.pop(req_id, None) + try: + future.set_result(msg) + + except asyncio.InvalidStateError: + logger.warning( + f"_result_listener: future already done for " + f"req_id={req_id}, dp_type={dp_type}, key={key}" + ) + + if dp_type == "register_destinations": + encode_still_pending = req_id in self.pending_futures[rank] + other_registration_pending = any( + pending_key.startswith(register_prefix) + for pending_key in self.pending_futures[rank] + ) + if not encode_still_pending and not other_registration_pending: + self.req_id_to_rank.pop(req_id, None) + elif dp_type == "wait_metadata": + encode_still_pending = req_id in self.pending_futures[rank] + other_metadata_pending = any( + pending_key.startswith(metadata_prefix) + for pending_key in self.pending_futures[rank] + ) + if ( + not encode_still_pending + and not other_metadata_pending + and req_id not in self._pending_send_at + ): + self.req_id_to_rank.pop(req_id, None) + + async def _cleanup_stale_mappings(self) -> None: + # Evict req_id->rank mappings whose /send never came. The worker frees + # its own embedding via the send_timeout cleanup scheduled at encode, + # so both sides key off the same timeout. + ttl = envs.SGLANG_ENCODER_SEND_TIMEOUT.get() + interval = max(ttl / 4, 30) + while True: + await asyncio.sleep(interval) + now = time.monotonic() + stale = [rid for rid, ts in self._pending_send_at.items() if now - ts > ttl] + for rid in stale: + self._pending_send_at.pop(rid, None) + self.req_id_to_rank.pop(rid, None) + if stale: + logger.warning( + f"Evicted {len(stale)} stale encoder DP /send mapping(s) " + f"with no /send within {ttl}s" + ) + + +async def _push_embedding_to_prefill( + enc: MMEncoder, + request: dict, + *, + background_url_send: bool = False, +) -> None: + """Deliver a staged ZMQ result and release it after the send completes.""" + req_id = request["req_id"] + backend = enc.transfer_backend + + if backend == "mooncake": + return + + if backend == "zmq_to_scheduler" and request.get("embedding_port") is None: + send_coro = enc.send_with_url(req_id=req_id) + if background_url_send: + task = asyncio.create_task(send_coro) + enc.background_tasks.add(task) + task.add_done_callback(enc.background_tasks.discard) + else: + await send_coro + return + + if backend == "zmq_to_tokenizer": + try: + await enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=request["embedding_port"], + ) + finally: + await enc.release_request(req_id) + return + + if backend == "zmq_to_scheduler": + ports = request["embedding_port"] + assert isinstance(ports, list) + try: + await asyncio.gather( + *( + enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=p, + ) + for p in ports + ) + ) + finally: + await enc.release_request(req_id) + + +def _record_pipeline_result(modality: Modality, status: str) -> None: + if server_module.encoder_metrics_collector is not None: + server_module.encoder_metrics_collector.inc_requests_total( + modality=modality.name.lower(), status=status + ) + + +async def execute_encode_pipeline( + enc: MMEncoder, + sched: Optional[EncoderScheduler], + request: dict, + *, + send_sockets: Optional[List[zmq.Socket]] = None, +) -> Optional[dict]: + """Run the shared HTTP/DP and Mooncake/ZMQ request lifecycle. + + Every backend publishes preprocess metadata. Mooncake has early consumers + and keeps the result until follow-up /send calls complete. ZMQ has no early + consumer: it waits for encode, sends the embedding, releases it, then returns. + """ + req_id = request["req_id"] + time_stats_json = request.pop("time_stats_json", None) + time_stats = EncoderReqTimeStats() + if time_stats_json: + time_stats.decode_json(time_stats_json) + request["enter_time"] = time.time() + modality = Modality.from_str(request["modality"]) + modality_str = modality.name.lower() + time_stats.modality = modality_str + time_stats.set_metrics_collector(server_module.encoder_metrics_collector) + backend = enc.transfer_backend + + if server_module.encoder_metrics_collector is not None: + server_module.encoder_metrics_collector.inc_requests_received( + modality=modality_str + ) + + time_stats.set_mm_encode_start_time() + try: + if sched is not None and modality in _BATCHABLE_MODALITIES: + result = await sched.submit(request) + elif send_sockets is not None: + # Non-batched requests still own their collective dispatch order + # directly; batched requests take this lock in _dispatch_group. + # Locking direct dispatch together with the rank0 await keeps its + # NCCL launch order matching the ZMQ dispatch order rank>0 sees. + async with enc.encode_dispatch_lock: + for socket in send_sockets: + sock_send(socket, wrap_as_pickle(request)) + result = await enc.encode( + mm_items=request["mm_items"], + modality=modality, + req_id=request["req_id"], + num_parts=request["num_parts"], + part_idx=request["part_idx"], + hashes=request.get("hashes"), + ) + else: + result = await enc.encode( + mm_items=request["mm_items"], + modality=modality, + req_id=request["req_id"], + num_parts=request["num_parts"], + part_idx=request["part_idx"], + hashes=request.get("hashes"), + ) + except asyncio.TimeoutError: + error_msg = "encoder batch timed out" + time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) + await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) + await enc.release_request(req_id, preserve_metadata=backend == "mooncake") + _record_pipeline_result(modality, "error") + raise + except Exception as e: + error_msg = str(e) + time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) + await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) + await enc.release_request(req_id, preserve_metadata=backend == "mooncake") + _record_pipeline_result(modality, "error") + raise + + nbytes, embedding_len, embedding_dim, error_msg, error_code = result + if error_msg: + time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) + await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) + if backend == "mooncake": + await enc.release_request(req_id, preserve_metadata=True) + else: + try: + await _push_embedding_to_prefill( + enc, + request, + background_url_send=True, + ) + except Exception as send_err: + logger.error( + f"Error-send failed for req_id={req_id}: {send_err}", + exc_info=True, + ) + _record_pipeline_result(modality, "error") + raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) + + time_stats.set_mm_encode_end_time() + try: + # Publish the actual result for every backend. ZMQ does not consume this + # early and removes it when its synchronous send releases the request. + await server_module.meta_registry.publish( + req_id, nbytes, embedding_len, embedding_dim + ) + + if backend == "mooncake": + request.pop("mm_items", None) + request.update( + embedding_size=nbytes, + embedding_len=embedding_len, + embedding_dim=embedding_dim, + ) + content = request + else: + await _push_embedding_to_prefill(enc, request) + content = None + except Exception as e: + time_stats.trace_ctx.abort(abort_info={"reason": str(e)}) + await enc.release_request(req_id) + _record_pipeline_result(modality, "error") + raise + + _record_pipeline_result(modality, "success") + return content + + +async def _dp_worker_health_encode(enc: MMEncoder) -> None: + """Functional health probe run on a DP worker. + + Process-liveness (proc.sentinel) can't see a worker that's alive but + wedged — hung GPU, NCCL deadlock, stalled ZMQ, or a blocked event loop. + When idle, run a tiny dummy encode to exercise the VIT forward and surface + those stalls. No prefill destination: the embedding is discarded, mirroring + the non-DP /health path. Raises on encode failure so the worker envelope + carries ``_error`` back to the dispatcher. + """ + if enc.supports_modality(Modality.IMAGE): + mm_items = [f"data:image/png;base64,{MINIMUM_PNG_PICTURE_BASE64}"] + modality = Modality.IMAGE + elif enc.supports_modality(Modality.AUDIO): + mm_items = [f"data:audio/wav;base64,{MINIMUM_WAV_SILENCE_BASE64}"] + modality = Modality.AUDIO + else: + # No processor → can't functionally probe; liveness alone is healthy. + return None + + # uuid keeps rids unique across workers; a bare time.time() can collide. + req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" + try: + async with enc.encode_dispatch_lock: + # Traffic may have started while the probe waited for the lock. + if enc.has_pending_embeddings(): + return None + _, _, _, error_msg, error_code = await enc.encode( + mm_items=mm_items, + modality=modality, + req_id=req_id, + num_parts=1, + part_idx=0, + ) + finally: + # Never leave the dummy embedding sitting in the send map. + await enc.release_request(req_id) + + if error_msg: + raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) + + +async def _dp_worker_handle_profile( + enc: MMEncoder, dp_rank: int, dp_type: str, request: dict +) -> dict: + prefix = f"dp_rank={dp_rank}: " + if dp_type == "start_profile": + req = request.get("profile_req") or ProfileReq() + req.req_type = ProfileReqType.START_PROFILE + if enc.profiler is None: + enc.profiler = EncoderProfiler(dp_rank) + ok, msg = enc.profiler.start(req) + detail = ( + f"started profiling, output_dir={enc.profiler.output_dir}" if ok else msg + ) + else: # stop_profile + if enc.profiler is None: + return {"ok": False, "msg": prefix + "profiling not initialized"} + ok, msg = enc.profiler.stop() + detail = "stopped profiling" if ok else msg + return {"ok": ok, "msg": prefix + detail} + + +async def _dp_worker_handle_request( + enc: MMEncoder, + sched: EncoderScheduler, + send_sock, + send_lock: asyncio.Lock, + dp_rank: int, + request: dict, + dp_type: str, +) -> None: + t0 = time.time() + try: + if dp_type in ("start_profile", "stop_profile"): + content = await _dp_worker_handle_profile(enc, dp_rank, dp_type, request) + elif dp_type == "health_encode": + content = await _dp_worker_health_encode(enc) + elif dp_type == "register_destinations": + await enc.register_embedding_destinations( + request["req_id"], + request["receive_count"], + [request["receive_url"]], + ) + content = None + elif dp_type == "wait_metadata": + try: + content = await server_module.meta_registry.wait(request["req_id"]) + except asyncio.TimeoutError as e: + raise MMError( + "encode metadata not ready", code=HTTPStatus.GATEWAY_TIMEOUT + ) from e + elif dp_type == "send": + req_id = request["req_id"] + sent = await enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=request["embedding_port"], + session_id=request["session_id"], + buffer_address=request["buffer_address"], + ) + if not sent: + # Error envelope, not 200 + phantom count: the decoder must + # fail fast instead of waiting for a ZMQ ack that never comes. + raise MMError( + f"no staged embedding for /send req_id={req_id} " + f"(already released)" + ) + # Releasing on the first /send breaks decoder TP > 1. No count means + # a pre-refcount decoder: stay eager rather than pin until the sweep. + receive_count = request.get("receive_count") + if receive_count: + await server_module.meta_registry.note_send_done(req_id, receive_count) + else: + await enc.release_request(req_id) + content = None + else: + content = await execute_encode_pipeline(enc, sched, request) + + logger.info( + f"MM-Encoder [dp_rank={dp_rank}] {dp_type} done: " + f"req_id={request.get('req_id', '?')}, " + f"modality={request.get('modality', 'image')}, " + f"cost={(time.time() - t0) * 1000:.1f}ms" + ) + envelope = { + "req_id": request.get("req_id", ""), + "_dp_type": dp_type, + "content": content, + } + if dp_type == "send" and request.get("_dp_send_key") is not None: + envelope["_dp_send_key"] = request["_dp_send_key"] + if ( + dp_type == "register_destinations" + and request.get("_dp_register_key") is not None + ): + envelope["_dp_register_key"] = request["_dp_register_key"] + if dp_type == "wait_metadata" and request.get("_dp_metadata_key") is not None: + envelope["_dp_metadata_key"] = request["_dp_metadata_key"] + except Exception as e: + logger.error( + f"DP worker {dp_rank} error on {dp_type} " + f"req_id={request.get('req_id', '?')}: {e}", + exc_info=True, + ) + err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR) + envelope = { + "req_id": request.get("req_id", ""), + "_dp_type": dp_type, + "content": None, + "_error": str(e), + "_error_type": type(e).__name__, + "_error_code": err_code, + } + if dp_type == "send" and request.get("_dp_send_key") is not None: + envelope["_dp_send_key"] = request["_dp_send_key"] + if ( + dp_type == "register_destinations" + and request.get("_dp_register_key") is not None + ): + envelope["_dp_register_key"] = request["_dp_register_key"] + if dp_type == "wait_metadata" and request.get("_dp_metadata_key") is not None: + envelope["_dp_metadata_key"] = request["_dp_metadata_key"] + + # pyzmq async send isn't safe for concurrent senders. + try: + async with send_lock: + await async_sock_send(send_sock, wrap_as_pickle(envelope)) + except Exception: + logger.error( + f"DP worker {dp_rank} failed to send envelope for " + f"req_id={request.get('req_id', '?')}", + exc_info=True, + ) + + +async def run_dp_worker( + server_args: ServerArgs, + dp_rank: int, + gpu_id: int, + dispatch_path: str, + result_path: str, +): + logger.info( + f"DP worker {dp_rank} starting on gpu_id={gpu_id} " + f"(CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES', 'unset')})" + ) + + # gpu_id is the device chosen by maybe_reindex_device_id in the parent: + # 0 when CVD is pinned to one GPU, else the absolute id. + enc = MMEncoder( + server_args, + dist_init_method=f"tcp://127.0.0.1:{get_free_port()}", + rank=0, + gpu_id=gpu_id, + ) + + if get_observability().enable_metrics: + set_prometheus_multiproc_dir() + labels = { + "model_name": get_serving().served_model_name, + "dp_rank": str(dp_rank), + } + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) + server_module.encoder_metrics_collector = EncoderMetricsCollector(labels) + enc.dp_rank = dp_rank + + max_batch_size, coalesce_same_turn = _resolve_encoder_batch_policy( + enc.model_type, + ENCODER_MAX_BATCH_SIZE, + ENCODER_MAX_BATCH_SIZE_EXPLICIT, + ) + sched = EncoderScheduler( + encoder=enc, + send_sockets=[], + max_batch_size=max_batch_size, + coalesce_same_turn=coalesce_same_turn, + ) + + ctx = zmq.asyncio.Context(2) + recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False) + send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False) + send_lock = asyncio.Lock() + inflight: Set[asyncio.Task] = set() + # Acquire-before-recv → back-pressure propagates to the dispatcher + # PUSH buffer. Must be at least max_batch_size or batching degrades. + max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get() + if max_inflight < max_batch_size: + logger.warning( + f"SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT={max_inflight} is below " + f"the effective encoder max_batch_size={max_batch_size}; the encoder " + f"will never assemble a full batch." + ) + inflight_sem = asyncio.Semaphore(max_inflight) + sched.start() + logger.info(f"DP worker {dp_rank} ready") + + try: + while True: + await inflight_sem.acquire() + spawned = False + try: + try: + request = await async_sock_recv(recv_sock) + except asyncio.CancelledError: + raise + except Exception: + logger.error(f"DP worker {dp_rank} recv error", exc_info=True) + continue + if not isinstance(request, dict): + logger.error( + f"DP worker {dp_rank} received non-dict request " + f"({type(request).__name__}); dropping" + ) + continue + dp_type = request.pop("_dp_type", "encode") + + async def _run(req=request, t=dp_type): + try: + await _dp_worker_handle_request( + enc, sched, send_sock, send_lock, dp_rank, req, t + ) + finally: + inflight_sem.release() + + task = asyncio.create_task(_run()) + spawned = True + inflight.add(task) + task.add_done_callback(inflight.discard) + finally: + if not spawned: + inflight_sem.release() + finally: + for task in inflight: + task.cancel() + ctx.destroy(linger=0) + + +def launch_dp_worker( + server_args: ServerArgs, + dp_rank: int, + gpu_id: int, + dispatch_path: str, + result_path: str, +): + try: + configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") + asyncio.run( + run_dp_worker(server_args, dp_rank, gpu_id, dispatch_path, result_path) + ) + except KeyboardInterrupt: + logger.info(f"DP worker {dp_rank} exiting") + except Exception: + traceback.print_exc() + + +def launch_local_runtime(server_args: ServerArgs) -> EncoderRuntime: + """Launch the current non-DP Scheduler and TP Encoder group. + + This function owns backend construction only. HTTP/gRPC middleware, + service registration, and network serving remain Transport concerns. + """ + if get_parallel().dp_size > 1: + raise ValueError( + "launch_local_runtime requires --dp-size 1; got " + f"dp_size={get_parallel().dp_size}." + ) + + # Set up prometheus metrics. + if get_observability().enable_metrics: + set_prometheus_multiproc_dir() + labels = { + "model_name": get_serving().served_model_name, + "dp_rank": "0", + } + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) + server_module.encoder_metrics_collector = EncoderMetricsCollector(labels) + + process_context = mp.get_context("spawn") + zmq_context = zmq.Context(10) + ipc_path_prefix = random_uuid() + port_args = PortArgs.init_new(server_args) + if get_parallel().dist_init_addr: + dist_init_method = NetworkAddress.parse(get_parallel().dist_init_addr).to_tcp() + else: + dist_init_method = NetworkAddress( + get_serving().host or "127.0.0.1", port_args.nccl_port + ).to_tcp() + + if get_observability().enable_trace: + process_tracing_init( + get_observability().otlp_traces_endpoint, + "sglang", + trace_modules=get_observability().trace_modules, + ) + trace_set_thread_info("Encoder") + + send_sockets: List[zmq.Socket] = [] + tp_processes: List[mp.Process] = [] + for rank in range(1, configured_tp_size()): + schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" + send_sockets.append( + get_zmq_socket(zmq_context, zmq.PUSH, schedule_path, bind=False) + ) + process = process_context.Process( + target=launch_encoder, + args=(server_args, schedule_path, dist_init_method, rank), + daemon=True, + ) + process.start() + tp_processes.append(process) + + encoder = MMEncoder(server_args, dist_init_method=dist_init_method) + max_batch_size, coalesce_same_turn = _resolve_encoder_batch_policy( + encoder.model_type, + ENCODER_MAX_BATCH_SIZE, + ENCODER_MAX_BATCH_SIZE_EXPLICIT, + ) + scheduler = EncoderScheduler( + encoder, + send_sockets, + max_batch_size=max_batch_size, + coalesce_same_turn=coalesce_same_turn, + ) + return EncoderRuntime( + encoder=encoder, + scheduler=scheduler, + send_sockets=send_sockets, + zmq_context=zmq_context, + tp_processes=tp_processes, + ) + + +def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: + """Launch the protocol-neutral DP backend and return its dispatcher. + + HTTP uses this entry point today. gRPC can reuse it later without + importing HTTP application state or Uvicorn. + """ + if get_parallel().dp_size <= 1 or server_args.tp_size != 1: + raise ValueError( + "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " + f"dp_size={get_parallel().dp_size}, tp_size={server_args.tp_size}." + ) + dp_size = get_parallel().dp_size + logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") + + # DP mode: workers (subprocesses) write metrics to the shared multiproc dir; + # the main process exposes the aggregated /metrics endpoint. + if get_observability().enable_metrics: + set_prometheus_multiproc_dir() + + ctx = mp.get_context("spawn") + ipc_prefix = random_uuid() + async_zmq_ctx = zmq.asyncio.Context(dp_size + 1) + + result_path = f"ipc:///tmp/{ipc_prefix}_dp_result" + result_socket = get_zmq_socket(async_zmq_ctx, zmq.PULL, result_path, True) + dispatch_sockets: List[zmq.asyncio.Socket] = [ + get_zmq_socket( + async_zmq_ctx, zmq.PUSH, f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{r}", True + ) + for r in range(dp_size) + ] + + worker_processes: List[mp.Process] = [] + + def _kill_workers(): + for process in worker_processes: + if process.is_alive(): + process.kill() + for process in worker_processes: + process.join(timeout=5) + + # Register atexit BEFORE spawn loop so partial spawns get reaped on + # exception (atexit holds the list ref and reads it at exit time). + atexit.register(_kill_workers) + + for dp_rank in range(dp_size): + gpu_id = server_args.base_gpu_id + dp_rank + # Pin the device parent-side around spawn (same convention as the + # scheduler launcher and DP controller) so the child inherits + # CUDA_VISIBLE_DEVICES from its first instruction, before any import + # can enumerate CUDA. No-op unless SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS + # is set, in which case gpu_id is reindexed to 0 and CVD is pinned. + with maybe_reindex_device_id(gpu_id) as gpu_id: + process = ctx.Process( + target=launch_dp_worker, + args=( + server_args, + dp_rank, + gpu_id, + f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}", + result_path, + ), + daemon=False, + ) + process.start() + worker_processes.append(process) + + labels = {"model_name": get_serving().served_model_name} + if server_args.extra_metric_labels: + labels.update(server_args.extra_metric_labels) + return DPDispatcher( + dp_size, + dispatch_sockets, + result_socket, + worker_processes, + enable_metrics=get_observability().enable_metrics, + labels=labels, + ) diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py new file mode 100644 index 000000000..da8c428e3 --- /dev/null +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -0,0 +1,2065 @@ +import asyncio +import concurrent.futures +import ctypes +import logging +import os +import pickle +import time +import traceback +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from http import HTTPStatus +from typing import Any, Awaitable, Callable, Dict, Iterable, List, Optional, Set, Tuple + +import aiohttp +import msgspec +import numpy as np +import torch +import zmq +import zmq.asyncio + +from sglang.srt.configs.device_config import DeviceConfig +from sglang.srt.configs.load_config import LoadConfig +from sglang.srt.configs.model_config import ModelConfig +from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX +from sglang.srt.disaggregation.encoder.preprocessor import ( + EncoderPreprocessor, + EncoderPreprocessResult, + _convert, + _mm_grid_attrs, +) +from sglang.srt.disaggregation.encoder.receiver import ( + EmbeddingData, + video_meta_attrs_for, +) +from sglang.srt.distributed.parallel_state import ( + get_default_distributed_backend, + get_mooncake_transfer_engine, + get_tp_group, + init_distributed_environment, + initialize_model_parallel, +) +from sglang.srt.environ import envs +from sglang.srt.layers.dp_attention import initialize_dp_attention +from sglang.srt.managers.io_struct import ( + ProfileReq, + ProfileReqType, + async_sock_recv, +) +from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache +from sglang.srt.model_executor.model_runner_components.load_model_utils import ( + maybe_precompile_model_kernels_after_loading, +) +from sglang.srt.model_loader import get_model as load_model +from sglang.srt.multimodal.encoder_preprocessing import ( + get_encoder_preprocessed_items, + resolve_encoder_media_processor_config, +) +from sglang.srt.observability.metrics_collector import EncoderMetricsCollector +from sglang.srt.runtime_context import ( + get_device, + get_disagg, + get_exec, + get_mm, + get_model, + publish, +) +from sglang.srt.server_args import ServerArgs +from sglang.srt.utils import configure_media_url_security +from sglang.srt.utils.network import ( + NetworkAddress, + config_socket, + get_local_ip_auto, + get_zmq_socket, +) + +logger = logging.getLogger(__name__) + + +def is_health_check_request(rid: Optional[str]) -> bool: + return isinstance(rid, str) and rid.startswith(HEALTH_CHECK_RID_PREFIX) + + +rid_lock = asyncio.Lock() +rid_to_receive_endpoint: Dict[str, Set[str]] = dict() +rid_to_receive_count: Dict[str, int] = dict() +cond_dict_lock = asyncio.Lock() +rid_to_cond: Dict[str, asyncio.Condition] = {} + + +async def _get_receive_condition(req_id: str) -> asyncio.Condition: + async with cond_dict_lock: + if req_id not in rid_to_cond: + rid_to_cond[req_id] = asyncio.Condition() + return rid_to_cond[req_id] + + +ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get() +ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set() +# Watchdog: max time to wait for a batched /encode result. Bounds HTTP latency +# if the batch worker stalls (NCCL hang, dead worker proc, etc.). +ENCODER_REQ_TIMEOUT = envs.SGLANG_ENCODER_REQ_TIMEOUT.get() + + +class EncoderMetaRegistry: + """Per-part metadata shared by every encoder request lifecycle. + + Mooncake decoder ranks consume it early to allocate landing buffers. ZMQ + publishes the same state for a uniform pipeline but does not consume it + before encode/send completes. + """ + + def __init__(self, *, wait_timeout: float, sweep_timeout: float): + # How long a decoder blocks in /scheduler_receive_meta_data. + self.wait_timeout = wait_timeout + # Backstop for state whose /send calls never all land. + self.sweep_timeout = sweep_timeout + self._rid_to_meta: Dict[str, dict] = {} + self._rid_to_send_done: Dict[str, int] = {} + self._pending_at: Dict[str, float] = {} + self._sweeper_task: Optional[asyncio.Task] = None + # Set only where the embedding also lives; None in the DP main process. + self.on_release: Optional[Callable[[str], Awaitable[None]]] = None + + def _touch(self, req_id: str) -> None: + self._pending_at[req_id] = time.monotonic() + self._ensure_sweeper() + + def _ensure_sweeper(self) -> None: + if self._sweeper_task is not None and not self._sweeper_task.done(): + return + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return + self._sweeper_task = loop.create_task(self._sweep_loop()) + + async def _sweep_loop(self) -> None: + # Same idiom as DPDispatcher._cleanup_stale_mappings: one eternal + # scanner; interval re-read each pass so the MMEncoder override applies. + while True: + await asyncio.sleep(max(self.sweep_timeout / 4, 0.01)) + now = time.monotonic() + stale = [ + rid + for rid, ts in self._pending_at.items() + if now - ts > self.sweep_timeout + ] + for rid in stale: + await self._release(rid) + + async def publish( + self, + req_id: str, + nbytes: int, + embedding_len: int, + embedding_dim: int, + error: Optional[str] = None, + ) -> None: + """Publish per-part metadata (or an error), wake waiters, arm the sweep.""" + meta = ( + {"error": error} + if error is not None + else { + "embedding_size": nbytes, + "embedding_len": embedding_len, + "embedding_dim": embedding_dim, + } + ) + async with rid_lock: + self._rid_to_meta[req_id] = meta + self._touch(req_id) + cond = await _get_receive_condition(req_id) + async with cond: + cond.notify_all() + + async def wait(self, req_id: str) -> Optional[dict]: + """Block until req_id's metadata is published; TimeoutError past wait_timeout. + No _touch here: a pull-first timestamp would let the sweeper pop the very + Condition this waiter holds, stranding it when publish notifies a new one.""" + cond = await _get_receive_condition(req_id) + async with cond: + await asyncio.wait_for( + cond.wait_for(lambda: self._rid_to_meta.get(req_id) is not None), + timeout=self.wait_timeout, + ) + return self._rid_to_meta.get(req_id) + + async def note_send_done(self, req_id: str, receive_count: int) -> None: + """Count one completed ``/send``; release everything at receive_count.""" + async with rid_lock: + count = self._rid_to_send_done.get(req_id, 0) + 1 + self._rid_to_send_done[req_id] = count + if count >= receive_count: + await self._release(req_id) + + async def _release(self, req_id: str) -> None: + if self.on_release is not None: + await self.on_release(req_id) + await self.discard(req_id) + + async def discard(self, req_id: str) -> None: + """Drop the meta rendezvous state for req_id. Idempotent.""" + async with rid_lock: + self._rid_to_meta.pop(req_id, None) + self._rid_to_send_done.pop(req_id, None) + self._pending_at.pop(req_id, None) + async with cond_dict_lock: + rid_to_cond.pop(req_id, None) + + +meta_registry = EncoderMetaRegistry( + wait_timeout=ENCODER_REQ_TIMEOUT, + sweep_timeout=envs.SGLANG_ENCODER_SEND_TIMEOUT.get(), +) + + +class MMError(Exception): + def __init__(self, message, code=HTTPStatus.INTERNAL_SERVER_ERROR): + self.message = message + self.code = code + super().__init__(self.message) + + +class BadRequestError(MMError): + def __init__(self, message): + super().__init__(message, code=HTTPStatus.BAD_REQUEST) + + +class InternalError(MMError): + def __init__(self, message): + super().__init__(message, code=HTTPStatus.INTERNAL_SERVER_ERROR) + + +class EncodeContext(msgspec.Struct): + """One flattened encode batch; a single request is the N=1 case.""" + + req_id: str # first request's id, for cache prefetch keys and logs + modality: Modality + preprocess_result: EncoderPreprocessResult + get_feature_fn: Any + mm_feature: Any + num_items: int + items_per_req: List[int] # grid entries per request, in flatten order + aux_data: dict + str_mm_hashes: Optional[List[str]] + use_global_cache: bool + is_health_check: bool + + +@dataclass +class ReqState: + """The result and in-flight work for one encoder request.""" + + req_id: str + embedding_data: Optional[EmbeddingData] = None + active_encodes: int = 0 + active_sends: int = 0 + release_requested: bool = False + preserve_metadata_on_release: bool = False + embedding_ready: asyncio.Event = field(default_factory=asyncio.Event, repr=False) + lifecycle_condition: asyncio.Condition = field( + default_factory=asyncio.Condition, repr=False + ) + + +@dataclass(frozen=True) +class SendDestination: + """One normalized destination for exactly one transfer.""" + + endpoint: str + session_id: Optional[str] = None + buffer_address: Optional[int] = None + + @classmethod + def from_host_port( + cls, + prefill_host: str, + embedding_port: int, + *, + session_id: Optional[str] = None, + buffer_address: Optional[int] = None, + ) -> "SendDestination": + return cls( + endpoint=NetworkAddress(prefill_host, embedding_port).to_host_port_str(), + session_id=session_id, + buffer_address=buffer_address, + ) + + @classmethod + def from_url(cls, url: str) -> "SendDestination": + return cls(endpoint=NetworkAddress.parse(url).to_host_port_str()) + + +class TensorWrapper: + """Wrapper to keep tensor alive while exposing buffer for zero-copy.""" + + def __init__(self, tensor): + # Ensure tensor is on CPU and contiguous + if tensor.is_cuda: + tensor = tensor.cpu() + if not tensor.is_contiguous(): + tensor = tensor.contiguous() + + # Keep tensor reference + self.tensor = tensor + self.shape = list(tensor.shape) + self.dtype = tensor.dtype + + def __buffer__(self): + data_ptr = self.tensor.data_ptr() + total_bytes = self.tensor.numel() * self.tensor.element_size() + c_obj = (ctypes.c_char * total_bytes).from_address(data_ptr) + c_obj._keep_alive_ref = self + return memoryview(c_obj) + + +class EncoderDelivery(ABC): + """Transfer backend boundary. Send never releases the request.""" + + def __init__(self, encoder: "MMEncoder"): + self.encoder = encoder + + @abstractmethod + async def send( + self, + state: ReqState, + destination: SendDestination, + ) -> None: ... + + @abstractmethod + async def release(self, state: ReqState) -> None: ... + + +class MooncakeDelivery(EncoderDelivery): + async def send( + self, + state: ReqState, + destination: SendDestination, + ) -> None: + mm_data = await self.encoder._wait_for_embedding(state) + await self.encoder._send( + mm_data.embedding, + mm_data, + session_id=destination.session_id, + buffer_address=destination.buffer_address, + url=destination.endpoint, + ) + + async def release(self, state: ReqState) -> None: + mm_data = state.embedding_data + if mm_data is not None and mm_data._mr_ptr is not None: + try: + self.encoder.engine.deregister(mm_data._mr_ptr) + except Exception as dereg_err: + logger.warning( + f"Shared-MR deregister failed for {state.req_id}: {dereg_err}" + ) + finally: + mm_data._mr_ptr = None + + +class ZmqDelivery(EncoderDelivery): + def __init__(self, encoder: "MMEncoder", *, cleanup_receive_state: bool) -> None: + super().__init__(encoder) + self.cleanup_receive_state = cleanup_receive_state + + async def send( + self, + state: ReqState, + destination: SendDestination, + ) -> None: + mm_data = await self.encoder._wait_for_embedding(state) + await self.encoder._send(mm_data.embedding, mm_data, url=destination.endpoint) + + async def release(self, state: ReqState) -> None: + if not self.cleanup_receive_state: + return + async with rid_lock: + rid_to_receive_endpoint.pop(state.req_id, None) + rid_to_receive_count.pop(state.req_id, None) + async with cond_dict_lock: + rid_to_cond.pop(state.req_id, None) + + +_mm_feature_attrs = { + Modality.IMAGE: ["pixel_values"], + Modality.VIDEO: ["pixel_values_videos"], + Modality.AUDIO: ["input_features"], +} + + +def _get_mm_feature(mm_inputs, modality): + for attr in _mm_feature_attrs[modality]: + if attr in mm_inputs: + return mm_inputs[attr] + raise ValueError( + f"Feature attrs ({_mm_feature_attrs[modality]}) not found in {mm_inputs}" + ) + + +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. + aux = { + attr: _normalize_aux_value(mm_inputs.get(attr)) + for attr in video_meta_attrs_for(model_type) + } + if model_type == "kimi_k3": + aux["original_image_sizes"] = _normalize_aux_value( + mm_inputs.get("original_image_sizes") + ) + return aux + + +class MMEncoder: + def __init__( + self, + server_args: ServerArgs, + schedule_path=None, + dist_init_method=None, + rank: int = 0, + gpu_id: Optional[int] = None, + ): + """``gpu_id`` pins this encoder to a device other than + ``base_gpu_id + rank`` — the DP launcher's per-worker placement. It is + this instance's value, not a config change, so it travels as an + argument.""" + # The DP and TP encoder workers are spawned, so this constructor is + # the first publish in those processes. + publish(server_args, role="encoder") + logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") + self.server_args = server_args + configure_media_url_security( + get_mm().allowed_media_domains, + server_args.media_url_max_file_size_mb, + ) + self.transfer_backend = get_disagg().encoder_transfer_backend + self.use_mooncake = self.transfer_backend == "mooncake" + self.rank = rank + # DP rank for metric labels; overridden by runtime.run_dp_worker. + # 0 in the single-instance (non-DP) path. + self.dp_rank = 0 + self.profiler = EncoderProfiler(rank) + + self.model_config = ModelConfig.from_server_args( + server_args, + ) + self.load_config = LoadConfig( + load_format=get_model().load_format, + download_dir=server_args.download_dir, + model_loader_extra_config=server_args.model_loader_extra_config, + remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, + remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, + remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, + ) + self.model_type = getattr( + self.model_config.hf_config, "model_type", "unknown" + ).lower() + + self.device = get_device().device + self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id + + self.device_config = DeviceConfig( + device=self.device, + gpu_id=self.gpu_id, + ) + + torch.get_device_module(self.device).set_device(self.gpu_id) + + init_distributed_environment( + backend=get_default_distributed_backend(self.device), + world_size=server_args.tp_size, + rank=rank, + distributed_init_method=dist_init_method, + local_rank=rank, + ) + initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size) + initialize_dp_attention(server_args, self.model_config) + + self.model = load_model( + model_config=self.model_config, + load_config=self.load_config, + device_config=self.device_config, + ) + encoder_media_processor_config = resolve_encoder_media_processor_config( + self.model + ) + maybe_precompile_model_kernels_after_loading(self.model, self.device) + + # CPU preprocessing pipeline (Rust-replaceable) + self.preprocessor = EncoderPreprocessor( + server_args=server_args, + model_config=self.model_config, + model_preprocessor=getattr(self.model, "preprocess_mm_for_encoder", None), + encoder_media_processor_config=encoder_media_processor_config, + ) + + self.context = zmq.asyncio.Context(2) + self.sync_context = zmq.Context() # Reuse sync context for thread pool + self.scheduler_send_sockets = {} + self.scheduler_send_locks = {} + self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=10) + + embedding_cache_size = int(os.environ.get("SGLANG_VLM_CACHE_SIZE_MB", "4096")) + self.mm_cache = MultiModalStaticCache(embedding_cache_size * 1024 * 1024) + self.mm_cache_lock = asyncio.Lock() + + self.send_timeout = envs.SGLANG_ENCODER_SEND_TIMEOUT.get() + + if schedule_path is not None: + self.schedule_socket = get_zmq_socket( + self.context, zmq.PULL, schedule_path, True + ) + self.background_tasks: Set[asyncio.Task] = set() + + # Embedding dtype = model param dtype. Always available (both transfer + # backends and the global-cache pool rely on it). + self._embedding_dtype = next(self.model.parameters()).dtype + self._element_size = torch.tensor( + [], dtype=self._embedding_dtype + ).element_size() + self._embedding_dims = self._infer_embedding_dims() + + if get_mm().enable_mm_global_cache: + from sglang.srt.mem_cache.embedding_cache_controller import ( + EmbeddingCacheController, + ) + from sglang.srt.mem_cache.embedding_store import EmbeddingStoreFactory + + embedding_store = EmbeddingStoreFactory.create_backend( + get_mm().mm_global_cache_backend, + ) + self.mm_global_cache = EmbeddingCacheController( + rank, + server_args.tp_size, + embedding_store=embedding_store, + hidden_dims=self._embedding_dims, + tp_group=get_tp_group().cpu_group, + all_rank_get=False, + dtype=self._embedding_dtype, + ) + else: + self.mm_global_cache = None + + if self.rank == 0: + logger.info( + f"Using transfer backend: {get_disagg().encoder_transfer_backend}" + ) + + if get_disagg().encoder_transfer_backend == "mooncake": + self.local_ip = get_local_ip_auto() + + self.engine = get_mooncake_transfer_engine() + if self.engine is None: + from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( + init_mooncake_transfer_engine, + ) + + self.engine = init_mooncake_transfer_engine( + hostname=self.local_ip, + gpu_id=self.gpu_id, + ib_device=( + get_disagg().disaggregation_ib_device + or get_exec().moe.mooncake_ib_device + ), + ) + + self.req_states: Dict[str, ReqState] = {} + # Need to ensure the NCCL launch order on rank0 matches the dispatch order rank>0 + self.encode_dispatch_lock = asyncio.Lock() + + if get_disagg().encoder_transfer_backend == "mooncake": + self.delivery = MooncakeDelivery(self) + # Embeddings live here, so registry cleanup uses the common release. + meta_registry.on_release = self.release_request + meta_registry.sweep_timeout = self.send_timeout + else: + self.delivery = ZmqDelivery( + self, + cleanup_receive_state=( + get_disagg().encoder_transfer_backend == "zmq_to_scheduler" + ), + ) + + logger.info(f"rank {rank} init finish ") + + def supports_modality(self, modality: Modality) -> bool: + return self.preprocessor.supports_modality(modality) + + def has_pending_embeddings(self) -> bool: + return bool(getattr(self, "req_states", None)) + + def _require_active_encode_state(self, req_id: str) -> ReqState: + """Return the state holding an encode ref; never resurrect a request.""" + state = self.req_states.get(req_id) + if state is None: + raise InternalError( + f"No request state exists while encoding request: {req_id}" + ) + if state.active_encodes <= 0: + raise InternalError(f"Request state has no active encode work: {req_id}") + return state + + def _acquire_encode_ref(self, req_id: str) -> Optional[ReqState]: + """Acquire a rank 0 encode ref before preprocessing can suspend.""" + if self.rank != 0: + return None + state = self.req_states.get(req_id) + if state is None: + state = ReqState(req_id) + self.req_states[req_id] = state + state.active_encodes += 1 + return state + + async def _release_encode_ref(self, state: Optional[ReqState]) -> None: + if state is None: + return + async with state.lifecycle_condition: + state.active_encodes -= 1 + assert state.active_encodes >= 0 + should_release = state.release_requested and state.active_encodes == 0 + state.lifecycle_condition.notify_all() + if should_release: + await self.release_request(state.req_id) + + def _stage_embedding(self, mm_data: EmbeddingData) -> None: + state = self._require_active_encode_state(mm_data.req_id) + metadata = state.embedding_data + if ( + metadata is not None + and metadata.embedding is None + and mm_data.embedding is not None + and (metadata.shape != mm_data.shape or metadata.dtype != mm_data.dtype) + ): + raise InternalError( + f"Embedding metadata mismatch for {mm_data.req_id}: " + f"expected={metadata.shape}/{metadata.dtype}, " + f"actual={mm_data.shape}/{mm_data.dtype}" + ) + state.embedding_data = mm_data + state.embedding_ready.set() + + async def _wait_for_embedding(self, state: ReqState) -> EmbeddingData: + await state.embedding_ready.wait() + if state.embedding_data is None: + raise InternalError(f"No embedding available for request: {state.req_id}") + return state.embedding_data + + async def send_to_destination( + self, state: ReqState, destination: SendDestination + ) -> None: + async with state.lifecycle_condition: + if ( + self.req_states.get(state.req_id) is not state + or state.release_requested + ): + raise InternalError(f"Encoder request was released: {state.req_id}") + state.active_sends += 1 + try: + await self.delivery.send(state, destination) + finally: + async with state.lifecycle_condition: + state.active_sends -= 1 + state.lifecycle_condition.notify_all() + + async def release_request( + self, req_id: str, *, preserve_metadata: bool = False + ) -> None: + """Release backend resources, then the embedding, through one path.""" + state = self.req_states.get(req_id) + if state is None: + if not preserve_metadata: + await meta_registry.discard(req_id) + return + async with state.lifecycle_condition: + state.release_requested = True + state.preserve_metadata_on_release |= preserve_metadata + if state.active_encodes > 0: + return + await state.lifecycle_condition.wait_for(lambda: state.active_sends == 0) + if self.req_states.get(req_id) is not state: + return + self.req_states.pop(req_id, None) + await self.delivery.release(state) + state.embedding_data = None + if not state.preserve_metadata_on_release: + await meta_registry.discard(req_id) + + async def register_embedding_destinations( + self, + req_id: str, + expected_destination_count: int, + destination_urls: Iterable[str], + ) -> None: + async with rid_lock: + if req_id not in rid_to_receive_endpoint: + rid_to_receive_endpoint[req_id] = set() + rid_to_receive_count[req_id] = expected_destination_count + registered_count = rid_to_receive_count[req_id] + if registered_count != expected_destination_count: + raise BadRequestError( + f"Inconsistent receive_count for req_id={req_id}: " + f"registered {registered_count}, got {expected_destination_count}" + ) + rid_to_receive_endpoint[req_id].update(destination_urls) + + cond = await _get_receive_condition(req_id) + async with cond: + cond.notify_all() + + def _infer_embedding_dims(self) -> dict: + """Infer per-modality embedding dimensions from hf_config at init time.""" + default = self.model_config.hidden_size + hf_cfg = self.model_config.hf_config + thinker_cfg = getattr(hf_cfg, "thinker_config", None) + dims = { + Modality.IMAGE: default, + Modality.VIDEO: default, + Modality.AUDIO: default, + } + + vision_cfg = getattr(thinker_cfg, "vision_config", None) or getattr( + hf_cfg, "vision_config", None + ) + if vision_cfg is not None: + out_hs = getattr(vision_cfg, "out_hidden_size", None) + if out_hs is not None: + ds = getattr(vision_cfg, "deepstack_visual_indexes", None) + vis_dim = ( + out_hs * (1 + len(ds)) + if isinstance(ds, (list, tuple)) and ds + else out_hs + ) + dims[Modality.IMAGE] = vis_dim + dims[Modality.VIDEO] = vis_dim + + audio_cfg = getattr(thinker_cfg, "audio_config", None) or getattr( + hf_cfg, "audio_config", None + ) + if audio_cfg is not None: + for attr in ("output_dim", "d_model"): + val = getattr(audio_cfg, attr, None) + if val and int(val) > 0: + dims[Modality.AUDIO] = int(val) + break + + logger.info(f"Global cache embedding dims: {dims}") + return dims + + def slice_embedding( + self, + mm_embedding: torch.Tensor, + token_counts: Iterable[int], + ) -> List[torch.Tensor]: + """Slice embeddings using preprocessing-owned token counts.""" + slices, offset = [], 0 + for count in token_counts: + slices.append(mm_embedding[offset : offset + count]) + offset += count + if mm_embedding.shape[0] != offset: + raise InternalError( + f"Encoder produced {mm_embedding.shape[0]} tokens, but " + f"preprocessor metadata expected {offset}" + ) + return slices + + def _calculate_hashes_from_features( + self, mm_feature, grid_thw: List, modality: Modality, mm_inputs=None + ) -> List[int]: + """CPU Task: Compute hashes based on processed feature patches.""" + preprocessed_items = ( + get_encoder_preprocessed_items(mm_inputs) if mm_inputs is not None else None + ) + if preprocessed_items is not None: + if len(preprocessed_items) != len(grid_thw): + raise ValueError( + "Encoder preprocess item/grid mismatch: " + f"{len(preprocessed_items)} items != {len(grid_thw)} grids" + ) + hashes = [] + for item in preprocessed_items: + item.set_pad_value() + hashes.append(item.hash) + return hashes + + hashes = [] + if modality == Modality.AUDIO and isinstance(mm_feature, list): + for feature in mm_feature: + tmp_item = MultimodalDataItem(modality=modality, feature=feature) + tmp_item.set_pad_value() + hashes.append(tmp_item.hash) + return hashes + + offset = 0 + logger.info(f"{mm_feature.shape=} with {modality=}") + for grid in grid_thw: + num_patches = self.preprocessor.get_num_patches(grid, modality) + feature_slice = mm_feature[offset : offset + num_patches] + tmp_item = MultimodalDataItem(modality=modality, feature=feature_slice) + tmp_item.set_pad_value() + hashes.append(tmp_item.hash) + offset += num_patches + return hashes + + def _encode_missing( + self, + mm_feature, + preprocess_result: EncoderPreprocessResult, + indices: List[int], + modality: Modality = Modality.IMAGE, + get_feature_fn=None, + ) -> List[torch.Tensor]: + """ + GPU Task: Run ViT inference ONLY on the subset of mm items missing from the cache. + """ + token_counts = preprocess_result.token_counts + mm_items = self._build_model_mm_items( + mm_feature, preprocess_result, indices, modality + ) + + forward_start = time.perf_counter() + with torch.inference_mode(): + new_embeddings = get_feature_fn(mm_items) + if new_embeddings.ndim != 2: + new_embeddings = new_embeddings.reshape(-1, new_embeddings.shape[-1]) + if encoder_metrics_collector is not None: + encoder_metrics_collector.observe_model_forward( + time.perf_counter() - forward_start, modality=modality.name.lower() + ) + + return self.slice_embedding(new_embeddings, (token_counts[i] for i in indices)) + + def _build_model_mm_items( + self, + mm_feature, + preprocess_result: EncoderPreprocessResult, + indices: List[int], + modality: Modality, + ) -> List[MultimodalDataItem]: + """Build the model-facing items selected for one encoder forward. + + Model preprocessors can preserve an item-wise representation with + ``EncoderPreprocessOutput``. This avoids concatenating and re-slicing + features before encoder-DP knows which rank owns each item. Legacy + processor outputs retain their aggregate tensor behavior. + """ + mm_inputs = preprocess_result.mm_inputs + grid_thw = preprocess_result.grid_thw + preprocessed_items = get_encoder_preprocessed_items(mm_inputs) + if preprocessed_items is not None: + if len(preprocessed_items) != len(grid_thw): + raise ValueError( + "Encoder preprocess item/grid mismatch: " + f"{len(preprocessed_items)} items != {len(grid_thw)} grids" + ) + selected = [preprocessed_items[index] for index in indices] + if any(item.modality != modality for item in selected): + raise ValueError("Encoder preprocess output contains wrong modality") + return selected + + split_kimi_k3_images = ( + self.model_type == "kimi_k3" and modality == Modality.IMAGE + ) + + if modality == Modality.AUDIO: + if isinstance(mm_feature, list): + sub_feature = [mm_feature[i] for i in indices] + else: + sub_feature = mm_feature[list(indices)] + else: + feature_slices = [] + offsets = [0] + curr = 0 + for grid in grid_thw: + curr += self.preprocessor.get_num_patches(grid, modality) + offsets.append(curr) + for idx in indices: + feature_slices.append(mm_feature[offsets[idx] : offsets[idx + 1]]) + if not split_kimi_k3_images: + sub_feature = torch.cat(feature_slices, dim=0) + + if split_kimi_k3_images: + mm_items = [ + MultimodalDataItem.from_dict( + {"modality": modality, "feature": _convert(feature)} + ) + for feature in feature_slices + ] + else: + mm_items = [ + MultimodalDataItem.from_dict( + { + "modality": modality, + "feature": ( + sub_feature + if isinstance(sub_feature, list) + else _convert(sub_feature) + ), + } + ) + ] + + for key, value in mm_inputs.items(): + if key in _mm_feature_attrs.get(modality, []): + continue + value = _convert(value) + if key in _mm_grid_attrs.get(modality, []): + if split_kimi_k3_images: + for mm_item, idx in zip(mm_items, indices): + mm_item.set(key, value[idx : idx + 1]) + else: + mm_items[0].set(key, value[indices]) + else: + for mm_item in mm_items: + mm_item.set(key, value) + return mm_items + + async def _prepare_encode_context( + self, + requests: List[dict], + modality: Modality, + *, + use_global_cache: bool, + is_health_check: bool = False, + ) -> EncodeContext: + """Flatten a batch of requests into one EncodeContext (single = N of 1).""" + modality_str = modality.name.lower() + preprocess_start = time.perf_counter() + try: + preprocess_result, items_per_req = ( + await self.preprocessor.process_batch_mm_items(requests, modality) + ) + except NotImplementedError as e: + raise InternalError(f"Not implemented error: {str(e)}") + except Exception as e: + raise BadRequestError(f"Failed to process mm items: {str(e)}") + + if len(items_per_req) != len(requests) or any(n <= 0 for n in items_per_req): + raise InternalError( + f"Invalid batch layout {items_per_req} for {len(requests)} requests" + ) + + if encoder_metrics_collector is not None and not is_health_check: + encoder_metrics_collector.observe_preprocess( + time.perf_counter() - preprocess_start, + modality=modality_str, + ) + for item_count in items_per_req: + encoder_metrics_collector.observe_mm_items_per_request( + item_count, modality=modality_str + ) + encoder_metrics_collector.observe_mm_items_per_batch( + sum(items_per_req), modality=modality_str + ) + target = self.model.thinker if hasattr(self.model, "thinker") else self.model + get_feature_fn = getattr(target, f"get_{modality_str}_feature") + + mm_inputs = preprocess_result.mm_inputs + grid_thw = preprocess_result.grid_thw + token_counts = preprocess_result.token_counts + mm_feature = _convert(_get_mm_feature(mm_inputs, modality)) + num_items = len(grid_thw) + if num_items != sum(items_per_req): + raise InternalError( + f"Batch layout {items_per_req} expects {sum(items_per_req)} " + f"grids, but the processor produced {num_items}" + ) + if len(token_counts) != num_items: + raise InternalError( + f"Preprocessor returned {len(token_counts)} token counts for " + f"{num_items} {modality_str} grid entries" + ) + + str_mm_hashes = None + if use_global_cache: + # Hashes must be grid-space per request (a leaf-space list would + # size-mismatch rank>0's mask and deadlock TP); validate on every + # rank so a bad request fails symmetrically before any collective. + per_req_hashes = [req.get("hashes") for req in requests] + mm_hashes = None + if all(h is not None for h in per_req_hashes): + for req, hashes, n in zip(requests, per_req_hashes, items_per_req): + if len(hashes) != n: + raise BadRequestError( + f"User-supplied hashes length {len(hashes)} != grid " + f"count {n} for req {req['req_id']}; hashes must be " + f"grid-space (1 per encoder grid entry)." + ) + mm_hashes = [h for hashes in per_req_hashes for h in hashes] + if self.rank == 0: + if mm_hashes is None: + mm_hashes = self._calculate_hashes_from_features( + mm_feature, grid_thw, modality, mm_inputs + ) + # Embedding stores use string cache keys. + str_mm_hashes = [str(h) for h in mm_hashes] + + return EncodeContext( + req_id=requests[0]["req_id"], + modality=modality, + preprocess_result=preprocess_result, + get_feature_fn=get_feature_fn, + mm_feature=mm_feature, + num_items=num_items, + items_per_req=items_per_req, + aux_data=_build_mm_aux_data(mm_inputs, self.model_type), + str_mm_hashes=str_mm_hashes, + use_global_cache=use_global_cache, + is_health_check=is_health_check, + ) + + def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): + if self.server_args.tp_size > 1: + torch.distributed.broadcast( + mask_tensor, + src=0, + group=self.mm_global_cache.prefetch_tp_group, + ) + + async def _lookup_global_cache( + self, + ctx: EncodeContext, + ) -> Tuple[List[int], List[int]]: + if self.rank == 0: + exist_mask = await self.mm_global_cache.batch_is_exist(ctx.str_mm_hashes) + mask_tensor = torch.tensor( + [1 if e else 0 for e in exist_mask], dtype=torch.int32 + ) + else: + mask_tensor = torch.zeros(ctx.num_items, dtype=torch.int32) + + self._broadcast_global_cache_mask(mask_tensor) + + 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] + return missing_indices, hit_indices + + def _prefetch_global_cache_hits( + self, + ctx: EncodeContext, + hit_indices: List[int], + ) -> List[str]: + if self.rank != 0 or not hit_indices: + return [] + + hit_hashes = [ctx.str_mm_hashes[i] for i in hit_indices] + hit_tokens = [ctx.preprocess_result.token_counts[i] for i in hit_indices] + self.mm_global_cache.prefetch(ctx.req_id, hit_hashes, hit_tokens, ctx.modality) + return hit_hashes + + async def _wait_global_cache_prefetch( + self, + ctx: EncodeContext, + hit_indices: List[int], + hit_hashes: List[str], + ) -> List[int]: + fallback_mask = torch.zeros(ctx.num_items, dtype=torch.int32) + if self.rank == 0 and hit_indices: + try: + + async def _wait_prefetch(): + while not self.mm_global_cache.check_prefetch_progress(ctx.req_id): + await asyncio.sleep(0.005) + + await asyncio.wait_for(_wait_prefetch(), timeout=60.0) + + for i, idx in enumerate(hit_indices): + if not self.mm_global_cache.has_local_embedding(hit_hashes[i]): + fallback_mask[idx] = 1 + num_partial_fail = int(fallback_mask.sum().item()) + if num_partial_fail > 0: + logger.warning( + f"Req {ctx.req_id}: {num_partial_fail}/{len(hit_indices)} " + f"cache-hit items failed to load, falling back to ViT" + ) + except (asyncio.TimeoutError, Exception) as e: + logger.error( + f"Prefetch failed for req {ctx.req_id}: {e}. " + f"Falling back to ViT for {len(hit_indices)} hit items." + ) + for idx in hit_indices: + fallback_mask[idx] = 1 + + self._broadcast_global_cache_mask(fallback_mask) + fallback_indices = [ + i for i in range(ctx.num_items) if fallback_mask[i].item() == 1 + ] + return fallback_indices + + def _launch_global_cache_insert( + self, + ctx: EncodeContext, + hashes: List[str], + d2h_handles: List[Any], + ): + if not hashes: + return + + async def _background_insert(): + await asyncio.to_thread( + self.mm_global_cache.wait_store_to_pool, + d2h_handles, + ) + await asyncio.to_thread( + self.mm_global_cache.insert_batch, + hashes, + ctx.modality, + ) + + task = asyncio.create_task(_background_insert()) + self.background_tasks.add(task) + task.add_done_callback(self.background_tasks.discard) + + @staticmethod + def _as_2d_tensor(tensor: torch.Tensor) -> torch.Tensor: + if tensor.ndim != 2: + tensor = tensor.reshape(-1, tensor.shape[-1]) + return tensor + + def _assemble_global_cache_cpu( + self, + ctx: EncodeContext, + hit_indices: List[int], + missing_indices: List[int], + fallback_indices: List[int], + new_slices: List[torch.Tensor], + fallback_slices: List[torch.Tensor], + ) -> torch.Tensor: + miss_slice_pos = {idx: pos for pos, idx in enumerate(missing_indices)} + fallback_slice_pos = {idx: pos for pos, idx in enumerate(fallback_indices)} + fallback_index_set = set(fallback_indices) + token_counts = ctx.preprocess_result.token_counts + dim = self.mm_global_cache.get_embedding_dim(ctx.modality) + + mm_embedding = torch.empty( + (sum(token_counts), dim), + dtype=self._embedding_dtype, + pin_memory=True, + ) + + hit_view_hashes = [ + ctx.str_mm_hashes[idx] + for idx in hit_indices + if idx not in fallback_index_set + ] + hit_views = {} + try: + if hit_view_hashes: + cached_slice_lists = self.mm_global_cache.get_pool_views( + hit_view_hashes + ) + for h, slices in zip(hit_view_hashes, cached_slice_lists): + if slices is None: + raise InternalError( + f"Cached embedding {h} not available for req {ctx.req_id}" + ) + hit_views[h] = slices + + offset = 0 + for idx, num_tokens in enumerate(token_counts): + if idx in miss_slice_pos: + src = self._as_2d_tensor(new_slices[miss_slice_pos[idx]]) + mm_embedding[offset : offset + num_tokens].copy_( + src, non_blocking=True + ) + elif idx in fallback_slice_pos: + src = self._as_2d_tensor(fallback_slices[fallback_slice_pos[idx]]) + mm_embedding[offset : offset + num_tokens].copy_( + src, non_blocking=True + ) + else: + copied = 0 + for view in hit_views[ctx.str_mm_hashes[idx]]: + n = view.shape[0] + mm_embedding[offset + copied : offset + copied + n].copy_(view) + copied += n + offset += num_tokens + + torch.cuda.current_stream(self.device).synchronize() + return mm_embedding + finally: + if hit_view_hashes: + self.mm_global_cache.release_pool_views(hit_view_hashes) + + def _assemble_global_cache_gpu( + self, + ctx: EncodeContext, + missing_indices: List[int], + fallback_indices: List[int], + new_slices: List[torch.Tensor], + fallback_slices: List[torch.Tensor], + ) -> torch.Tensor: + miss_slice_pos = {idx: pos for pos, idx in enumerate(missing_indices)} + fallback_slice_pos = {idx: pos for pos, idx in enumerate(fallback_indices)} + token_counts = ctx.preprocess_result.token_counts + embedding_dim = self.mm_global_cache.get_embedding_dim(ctx.modality) + mm_embedding = torch.empty( + (sum(token_counts), embedding_dim), + dtype=self._embedding_dtype, + device=self.device, + ) + + offset = 0 + copy_handles = [] + for idx, num_tokens in enumerate(token_counts): + if idx in miss_slice_pos: + mm_embedding[offset : offset + num_tokens].copy_( + new_slices[miss_slice_pos[idx]], + non_blocking=True, + ) + elif idx in fallback_slice_pos: + mm_embedding[offset : offset + num_tokens].copy_( + fallback_slices[fallback_slice_pos[idx]], + non_blocking=True, + ) + else: + handle = self.mm_global_cache.load_to_device_async( + ctx.str_mm_hashes[idx], mm_embedding, offset + ) + if handle is None: + raise InternalError( + f"Cached embedding {ctx.str_mm_hashes[idx]} disappeared " + f"during assembly for req {ctx.req_id}" + ) + copy_handles.append(handle) + offset += num_tokens + + self.mm_global_cache.wait_load_to_device(copy_handles) + torch.cuda.current_stream(mm_embedding.device).synchronize() + return mm_embedding + + async def _compute_global_cache_embedding( + self, + ctx: EncodeContext, + *, + keep_on_gpu: bool, + ) -> Optional[torch.Tensor]: + """Resolve cache hits, compute misses, assemble output, and insert misses.""" + missing_indices, hit_indices = await self._lookup_global_cache(ctx) + hit_hashes = self._prefetch_global_cache_hits(ctx, hit_indices) + + new_slices = [] + if missing_indices: + new_slices = self._encode_missing( + ctx.mm_feature, + ctx.preprocess_result, + missing_indices, + ctx.modality, + ctx.get_feature_fn, + ) + + miss_d2h_handles = [] + # The CPU output path starts D2H staging before waiting for cache-hit loads. + if self.rank == 0 and new_slices and not keep_on_gpu: + miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] + miss_d2h_handles = self.mm_global_cache.store_to_pool_async( + miss_hashes, new_slices, ctx.modality + ) + + fallback_indices = await self._wait_global_cache_prefetch( + ctx, hit_indices, hit_hashes + ) + + fallback_slices = [] + fallback_d2h_handles = [] + if fallback_indices: + logger.info( + f"Req {ctx.req_id}: All ranks running ViT fallback " + f"for {len(fallback_indices)} items." + ) + fallback_slices = self._encode_missing( + ctx.mm_feature, + ctx.preprocess_result, + fallback_indices, + ctx.modality, + ctx.get_feature_fn, + ) + if self.rank == 0 and not keep_on_gpu: + fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] + fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( + fallback_hashes, fallback_slices, ctx.modality + ) + + if self.rank == 0: + if keep_on_gpu: + # Start staging newly computed GPU slices into the CPU cache + # pool asynchronously before assembling the GPU output. + if new_slices: + miss_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] + miss_d2h_handles = self.mm_global_cache.store_to_pool_async( + miss_hashes, new_slices, ctx.modality + ) + if fallback_slices: + fallback_hashes = [ctx.str_mm_hashes[i] for i in fallback_indices] + fallback_d2h_handles = self.mm_global_cache.store_to_pool_async( + fallback_hashes, fallback_slices, ctx.modality + ) + mm_embedding = self._assemble_global_cache_gpu( + ctx, + missing_indices, + fallback_indices, + new_slices, + fallback_slices, + ) + else: + mm_embedding = self._assemble_global_cache_cpu( + ctx, + hit_indices, + missing_indices, + fallback_indices, + new_slices, + fallback_slices, + ) + + new_hashes = [ctx.str_mm_hashes[i] for i in missing_indices] + new_hashes += [ctx.str_mm_hashes[i] for i in fallback_indices] + self._launch_global_cache_insert( + ctx, + new_hashes, + miss_d2h_handles + fallback_d2h_handles, + ) + return mm_embedding + + return None + + async def _compute_direct_embedding( + self, + ctx: EncodeContext, + *, + keep_on_gpu: bool, + ) -> torch.Tensor: + """Compute without global cache, optionally using the prefix MM cache.""" + modality = ctx.modality + modality_str = modality.name.lower() + try: + mm_embedding = None + mm_hash = None + + mm_items = self._build_model_mm_items( + ctx.mm_feature, + ctx.preprocess_result, + list(range(ctx.num_items)), + modality, + ) + + cache_hit = False + # The prefix cache hashes the whole request; a fused multi-request + # batch has no per-request key, so only N=1 contexts use it. + use_mm_cache = ( + get_mm().enable_prefix_mm_cache + and not ctx.is_health_check + and not keep_on_gpu + and len(ctx.items_per_req) == 1 + ) + if use_mm_cache: + for mm_item in mm_items: + mm_item.set_pad_value() + mm_hashes = [mm_item.hash for mm_item in mm_items] + mm_hash = MultiModalStaticCache.combine_hashes(mm_hashes) + async with self.mm_cache_lock: + mm_cache = self.mm_cache.get(mm_hashes) + if mm_cache is not None: + mm_embedding = mm_cache.embedding + cache_hit = True + + if mm_embedding is None: + forward_start = time.perf_counter() + with torch.inference_mode(): + mm_embedding: torch.Tensor = ctx.get_feature_fn(mm_items) + if not keep_on_gpu: + mm_embedding = mm_embedding.cpu() + if len(mm_embedding.shape) != 2: + mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1]) + if encoder_metrics_collector is not None and not ctx.is_health_check: + encoder_metrics_collector.observe_model_forward( + time.perf_counter() - forward_start, modality=modality_str + ) + + # Per-request cache hit metrics: tokens = embedding rows. + if use_mm_cache and encoder_metrics_collector is not None: + total_tokens = int(mm_embedding.shape[0]) + hit_tokens = total_tokens if cache_hit else 0 + encoder_metrics_collector.record_cache_tokens( + hit_tokens, total_tokens, modality=modality_str + ) + encoder_metrics_collector.record_cache_files( + len(mm_items) if cache_hit else 0, + len(mm_items), + modality=modality_str, + ) + + if use_mm_cache: + async with self.mm_cache_lock: + entries_before = len(self.mm_cache) + already_present = self.mm_cache.has(mm_hash) + inserted = self.mm_cache.set( + mm_hash, EmbeddingResult(embedding=mm_embedding) + ) + entries_after = len(self.mm_cache) + if encoder_metrics_collector is not None: + added = 0 if already_present else (1 if inserted else 0) + evictions = max(0, added - (entries_after - entries_before)) + if evictions > 0: + encoder_metrics_collector.inc_cache_evictions( + modality=modality_str, count=evictions + ) + encoder_metrics_collector.set_cache_state( + self.mm_cache.current_size, entries_after + ) + + if ( + not keep_on_gpu + and modality == Modality.VIDEO + and ctx.preprocess_result.mm_inputs.get("video_audio_features") + ): + target = ( + self.model.thinker if hasattr(self.model, "thinker") else self.model + ) + encode_video_audio_fn = getattr(target, "encode_video_audio", None) + if encode_video_audio_fn is not None: + audio_forward_start = time.perf_counter() + audio_embedding = encode_video_audio_fn( + ctx.preprocess_result.mm_inputs + ) + if ( + encoder_metrics_collector is not None + and not ctx.is_health_check + ): + encoder_metrics_collector.observe_model_forward( + time.perf_counter() - audio_forward_start, modality="audio" + ) + if audio_embedding is not None: + ctx.aux_data["video_audio_embedding"] = audio_embedding + else: + logger.warning( + "Videos carry audio tracks but model has no " + "encode_video_audio; dropping audio for EPD encoding." + ) + + return mm_embedding + except BadRequestError as e: + raise BadRequestError(f"Bad request error: {str(e)}") + except Exception as e: + raise InternalError(f"Internal encoding error: {str(e)}") + + async def _compute_embedding( + self, + ctx: EncodeContext, + *, + keep_on_gpu: bool, + ) -> Optional[torch.Tensor]: + """Compute one flattened request with global cache as an optional stage.""" + if ctx.use_global_cache: + mm_embedding = await self._compute_global_cache_embedding( + ctx, keep_on_gpu=keep_on_gpu + ) + else: + mm_embedding = await self._compute_direct_embedding( + ctx, keep_on_gpu=keep_on_gpu + ) + + expected_tokens = sum(ctx.preprocess_result.token_counts) + if mm_embedding is not None and mm_embedding.shape[0] != expected_tokens: + raise InternalError( + f"Encoder produced {mm_embedding.shape[0]} tokens, but " + f"preprocessor metadata expected {expected_tokens}" + ) + return mm_embedding + + async def _publish_preprocess_metadata( + self, ctx: EncodeContext, requests: List[dict] + ) -> None: + """Publish each request's size after preprocessing, before model forward.""" + if self.rank != 0: + return + embedding_dim = self._embedding_dims[ctx.modality] + item_offset = 0 + for request, item_count in zip(requests, ctx.items_per_req): + item_end = item_offset + item_count + token_count = sum(ctx.preprocess_result.token_counts[item_offset:item_end]) + req_id = request["req_id"] + state = self._require_active_encode_state(req_id) + state.embedding_data = EmbeddingData( + req_id, + request["num_parts"], + request["part_idx"], + ctx.preprocess_result.grid_thw[item_offset:item_end], + ctx.modality, + embedding_shape=[token_count, embedding_dim], + dtype=self._embedding_dtype, + ) + await meta_registry.publish( + req_id, + token_count * embedding_dim * self._element_size, + token_count, + embedding_dim, + ) + item_offset = item_end + + async def _send( + self, + embedding: torch.Tensor, + mm_data: EmbeddingData, + session_id=None, + buffer_address=None, + prefill_host=None, + embedding_port=None, + url=None, + ): + if get_disagg().encoder_transfer_backend == "mooncake": + # Encode is synchronous, so mm_data was staged before /encode returned. + req_id = mm_data.req_id + 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})" + ) + + # Fall back to a per-send registration only if the shared one failed. + mr_already_registered = mm_data._mr_ptr == embedding.data_ptr() + if not mr_already_registered: + self.engine.register(embedding.data_ptr(), embedding.nbytes) + _t_xfer_start = time.monotonic() + xfer_ret = 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 encoder_metrics_collector is not None: + encoder_metrics_collector.observe_transfer( + xfer_ms / 1000.0, backend="mooncake" + ) + if not mr_already_registered: + self.engine.deregister(embedding.data_ptr()) + if xfer_ret < 0: + raise InternalError( + f"Mooncake transfer_sync failed for {req_id} " + f"(session={session_id}, nbytes={embedding.nbytes}, " + f"ret={xfer_ret})" + ) + # Emit at INFO for slow transfers or per-send registrations. + 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}" + ) + + # Sibling ranks re-read mm_data here; meta_registry owns the release. + + # Send ack/data + if url is not None: + endpoint = NetworkAddress.parse(url).to_tcp() + else: + endpoint = NetworkAddress(prefill_host, embedding_port).to_tcp() + logger.info(f"{endpoint = }") + + # Serialize data + if get_disagg().encoder_transfer_backend == "mooncake": + # 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() + if new_mm_data.error_msg is not None: + buffer = None + serialized_data = pickle.dumps(new_mm_data) + else: + embedding_tensor = TensorWrapper(mm_data.embedding) + serialized_data = pickle.dumps(new_mm_data) + buffer = embedding_tensor.__buffer__() + + transfer_start = time.perf_counter() + if self.transfer_backend == "zmq_to_scheduler" and url is not None: + lock = self.scheduler_send_locks.get(endpoint) + if lock is None: + lock = asyncio.Lock() + self.scheduler_send_locks[endpoint] = lock + + async with lock: + sock = self.scheduler_send_sockets.get(endpoint) + if sock is None: + sock = self.context.socket(zmq.PUSH) + config_socket(sock, zmq.PUSH) + sock.setsockopt(zmq.IMMEDIATE, 1) + sock.setsockopt(zmq.SNDTIMEO, int(self.send_timeout * 1000)) + sock.connect(endpoint) + self.scheduler_send_sockets[endpoint] = sock + try: + frames = ( + [serialized_data, buffer] + if buffer is not None + else [serialized_data] + ) + tracker = await sock.send_multipart(frames, copy=False, track=True) + except Exception: + if self.scheduler_send_sockets.get(endpoint) is sock: + self.scheduler_send_sockets.pop(endpoint, None) + sock.close(linger=0) + raise + + # MessageTracker.wait() protects the zero-copy source buffer; it + # is not a receiver acknowledgement. Waiting under the per-peer + # lock serialized every large embedding on that TCP connection. + # Queue sends in order under the lock, then wait for buffer + # ownership independently so libzmq can pipeline the connection. + try: + await asyncio.to_thread(tracker.wait, self.send_timeout) + except Exception: + if self.scheduler_send_sockets.get(endpoint) is sock: + self.scheduler_send_sockets.pop(endpoint, None) + sock.close(linger=0) + raise + + if encoder_metrics_collector is not None: + encoder_metrics_collector.observe_transfer( + time.perf_counter() - transfer_start, + backend=self.transfer_backend, + ) + return + + # Per-request sockets remain for zmq_to_tokenizer and legacy direct + # scheduler sends. Scheduler URL sends use persistent sockets above. + def send_with_socket(): + sock = self.sync_context.socket(zmq.PUSH) + config_socket(sock, zmq.PUSH) + sock.setsockopt(zmq.IMMEDIATE, 1) + sock.setsockopt(zmq.SNDTIMEO, int(self.send_timeout * 1000)) + try: + sock.connect(endpoint) + if buffer is not None: + tracker = sock.send_multipart( + [serialized_data, buffer], copy=False, track=True + ) + else: + tracker = sock.send_multipart( + [serialized_data], copy=False, track=True + ) + tracker.wait(timeout=self.send_timeout) + finally: + sock.close(linger=5000) + + await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) + if ( + encoder_metrics_collector is not None + and get_disagg().encoder_transfer_backend != "mooncake" + ): + encoder_metrics_collector.observe_transfer( + time.perf_counter() - transfer_start, + backend=get_disagg().encoder_transfer_backend, + ) + + def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor): + """Register one MR shared by every rank's /send; _send re-registers on failure.""" + try: + self.engine.register(embedding.data_ptr(), embedding.nbytes) + mm_data._mr_ptr = embedding.data_ptr() + except Exception as reg_err: + logger.warning( + f"Shared-MR register failed for {mm_data.req_id}, " + f"falling back to per-/send register: {reg_err}" + ) + + def _stage_embeddings( + self, + ctx: EncodeContext, + requests: List[dict], + mm_embedding: Optional[torch.Tensor], + *, + keep_on_gpu: bool, + ) -> List[Tuple[int, int, int, Optional[str], Optional[int]]]: + """Split the fused embedding per request and stage one EmbeddingData each. + + Per-request token ranges are contiguous in flatten order, so each + staged embedding is a slice of the batch tensor. + """ + if self.rank != 0: + return [(0, 0, 0, None, None)] * len(requests) + if mm_embedding is None: + raise InternalError(f"Rank 0 produced no embedding for {ctx.req_id}") + + results = [] + staged_embeddings = [] + item_offset = 0 + token_offset = 0 + for req, num_items in zip(requests, ctx.items_per_req): + item_end = item_offset + num_items + num_tokens = sum(ctx.preprocess_result.token_counts[item_offset:item_end]) + embedding = mm_embedding[token_offset : token_offset + num_tokens] + if keep_on_gpu and len(requests) > 1: + # A view would pin the whole batch tensor until the last transfer. + embedding = embedding.clone() + req_aux_data = dict(ctx.aux_data) + if ctx.aux_data.get("original_image_sizes") is not None: + req_aux_data["original_image_sizes"] = ctx.aux_data[ + "original_image_sizes" + ][item_offset:item_end] + mm_data = EmbeddingData( + req["req_id"], + req["num_parts"], + req["part_idx"], + ctx.preprocess_result.grid_thw[item_offset:item_end], + ctx.modality, + embedding, + **req_aux_data, + ) + # Global-cache embeddings keep registering per /send instead. + if keep_on_gpu and not ctx.use_global_cache: + self._register_shared_mr(mm_data, embedding) + staged_embeddings.append(mm_data) + results.append( + (embedding.nbytes, embedding.shape[0], embedding.shape[1], None, None) + ) + item_offset = item_end + token_offset += num_tokens + + # transfer_sync bypasses CUDA streams, so GPU writes (forward and the + # per-request clones) must land before /send reads the buffers. + if keep_on_gpu and mm_embedding.is_cuda: + torch.cuda.current_stream(mm_embedding.device).synchronize() + for mm_data in staged_embeddings: + self._stage_embedding(mm_data) + return results + + def _stage_errors( + self, requests: List[dict], modality: Modality, exc: Exception + ) -> List[Tuple[int, int, int, Optional[str], Optional[int]]]: + """Stage one error EmbeddingData per request so /send reports the failure.""" + code = ( + exc.code if isinstance(exc, MMError) else HTTPStatus.INTERNAL_SERVER_ERROR + ) + msg = str(exc) + logger.error(f"Rank {self.rank} encode failed: {msg} {code = }", exc_info=True) + if self.rank == 0: + for req in requests: + self._stage_embedding( + EmbeddingData( + req["req_id"], + req["num_parts"], + req["part_idx"], + None, + modality, + error_msg=msg, + error_code=code, + ) + ) + return [(0, 0, 0, msg, code)] * len(requests) + + async def batch_encode( + self, requests: List[dict], modality: Modality + ) -> List[Tuple[int, int, int, Optional[str], Optional[int]]]: + """Encode requests through one fused pipeline; encode() is the N=1 case. + + Fuse-or-not is EncoderScheduler policy, not an API fork. Health probes + bypass caches and stage on CPU so completion confirms a model forward. + """ + states = [self._acquire_encode_ref(req["req_id"]) for req in requests] + is_health_check = all( + is_health_check_request(req["req_id"]) for req in requests + ) + keep_on_gpu = self.use_mooncake and not is_health_check + use_global_cache = self.mm_global_cache is not None and not is_health_check + try: + ctx = await self._prepare_encode_context( + requests, + modality, + use_global_cache=use_global_cache, + is_health_check=is_health_check, + ) + await self._publish_preprocess_metadata(ctx, requests) + mm_embedding = await self._compute_embedding(ctx, keep_on_gpu=keep_on_gpu) + + if self.profiler is not None: + for _ in requests: + self.profiler.step() + + return self._stage_embeddings( + ctx, requests, mm_embedding, keep_on_gpu=keep_on_gpu + ) + except Exception as e: + return self._stage_errors(requests, modality, e) + finally: + for state in states: + await self._release_encode_ref(state) + + async def encode( + self, mm_items, modality: Modality, req_id, num_parts, part_idx, hashes=None + ): + """Encode one request: the batch-of-1 case of batch_encode.""" + results = await self.batch_encode( + [ + { + "req_id": req_id, + "num_parts": num_parts, + "part_idx": part_idx, + "mm_items": mm_items, + "hashes": hashes, + } + ], + modality, + ) + return results[0] + + async def encode_request(self, req: dict, modality: Modality): + """Adapt a request dictionary to the single-request encode interface.""" + return await self.encode( + mm_items=req["mm_items"], + modality=modality, + req_id=req["req_id"], + num_parts=req["num_parts"], + part_idx=req["part_idx"], + hashes=req.get("hashes"), + ) + + # For zmq_to_tokenizer zmq_to_scheduler and mooncake + async def send( + self, req_id, prefill_host, embedding_port, session_id=None, buffer_address=None + ): + state = self.req_states.get(req_id) + if state is None: + # False = nothing transferred: callers must not count this send + # nor report success, or the decoder waits on an ack never coming. + logger.warning( + f"MMEncoder.send: no embedding for req_id={req_id} " + f"(already released or unknown)" + ) + return False + await self.send_to_destination( + state, + SendDestination.from_host_port( + prefill_host, + embedding_port, + session_id=session_id, + buffer_address=buffer_address, + ), + ) + return True + + # For zmq_to_scheduler + async def send_with_url( + self, + req_id, + ): + state = self.req_states.get(req_id) + if state is None: + return + sent_urls: Set[str] = set() + all_tasks: List[Tuple[asyncio.Task, str]] = [] + start_time = asyncio.get_running_loop().time() + timeout = self.send_timeout + cond = await _get_receive_condition(req_id) + + try: + while True: + async with rid_lock: + current_targets = rid_to_receive_endpoint.get(req_id, set()).copy() + expected_count = rid_to_receive_count.get(req_id) + + new_targets = current_targets - sent_urls + + if new_targets: + logger.info( + f"Found {len(new_targets)} new endpoints for {req_id}. Starting tasks..." + ) + for url in new_targets: + task = asyncio.create_task( + self.send_to_destination( + state, + SendDestination.from_url(url), + ) + ) + all_tasks.append((task, url)) + sent_urls.add(url) # Mark as handled immediately + if expected_count is not None and len(sent_urls) >= expected_count: + logger.info( + f"All {expected_count} endpoints initiated for {req_id}. Breaking loop." + ) + break + remaining = timeout - (asyncio.get_running_loop().time() - start_time) + if remaining <= 0: + logger.error( + f"[{req_id}] Timeout! Sent {len(sent_urls)}/{expected_count}" + ) + break + + async with cond: + try: + await asyncio.wait_for(cond.wait(), timeout=remaining) + except asyncio.TimeoutError: + continue + + if all_tasks: + logger.info( + f"Loop finished. Awaiting completion of {len(all_tasks)} sending tasks..." + ) + tasks_only = [t[0] for t in all_tasks] + results = await asyncio.gather(*tasks_only, return_exceptions=True) + + # Process results and log errors + for i, result in enumerate(results): + url = all_tasks[i][1] # Retrieve URL associated with the task + if isinstance(result, Exception): + logger.error(f"Failed to send to {url}: {result}") + else: + logger.debug(f"Successfully sent to {url}") + + logger.info(f"All tasks completed for req_id: {req_id}") + + finally: + logger.info(f"Cleaning up resources for req_id {req_id}") + await self.release_request(req_id) + + async def get_embedding_port(self, prefill_url): + async with aiohttp.ClientSession( + timeout=aiohttp.ClientTimeout(total=1800) + ) as session: + response = await session.post( + f"{prefill_url}/embedding_bootstrap", + json={"embedding_port": None}, + ) + response_json = await response.json() + return response_json["embedding_port"] + + +class EncoderProfiler: + def __init__(self, rank: int): + self.rank = rank + self.profiler = None + self.steps_left = None + self.output_dir = None + self.prefix = None + self.profile_id = None + + def start(self, obj: ProfileReq): + if self.profiler is not None: + return False, "profiling already running" + + output_dir = obj.output_dir or os.getenv("SGLANG_TORCH_PROFILER_DIR", "/tmp") + os.makedirs(output_dir, exist_ok=True) + self.output_dir = output_dir + self.prefix = obj.profile_prefix or "encoder" + self.profile_id = str(time.time()) + + activities = obj.activities or ["CPU", "GPU"] + torch_activities = [] + if "CPU" in activities: + torch_activities.append(torch.profiler.ProfilerActivity.CPU) + if "GPU" in activities: + torch_activities.append(torch.profiler.ProfilerActivity.CUDA) + + profile_memory = "MEM" in activities + if not torch_activities and not profile_memory: + return False, "no supported activities" + + self.profiler = torch.profiler.profile( + activities=torch_activities, + with_stack=True if obj.with_stack is None else obj.with_stack, + record_shapes=False if obj.record_shapes is None else obj.record_shapes, + profile_memory=profile_memory, + ) + self.profiler.start() + self.steps_left = obj.num_steps + logger.info( + f"Encoder profiling started. output_dir={self.output_dir} profile_id={self.profile_id}" + ) + return True, None + + def step(self): + if self.profiler is None: + return + self.profiler.step() + if self.steps_left is not None: + self.steps_left -= 1 + if self.steps_left <= 0: + self.stop() + + def stop(self): + if self.profiler is None: + return False, "profiling not running" + self.profiler.stop() + filename = f"{self.prefix}-rank{self.rank}-{self.profile_id}.trace.json" + trace_path = os.path.join(self.output_dir, filename) + self.profiler.export_chrome_trace(trace_path) + logger.info("Encoder profiling saved to: %s", trace_path) + self.profiler = None + self.steps_left = None + return True, None + + +async def run_encoder( + server_args: ServerArgs, schedule_path, dist_init_method, rank: int +): + encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank) + while True: + request = await async_sock_recv(encoder.schedule_socket) + await _handle_encoder_worker_request(encoder, request) + + +async def _handle_encoder_worker_request(encoder: MMEncoder, request): + if isinstance(request, ProfileReq): + if request.req_type == ProfileReqType.START_PROFILE: + if encoder.profiler is None: + encoder.profiler = EncoderProfiler(encoder.rank) + encoder.profiler.start(request) + else: + encoder.profiler.stop() + elif isinstance(request, dict) and request.get("type") == "batch_encode": + await encoder.batch_encode( + request["requests"], + Modality.from_str(request["modality"]), + ) + else: + # Health-check rids need no special routing: batch_encode derives + # health semantics from the rid prefix itself. + await encoder.encode_request(request, Modality.from_str(request["modality"])) + + +def launch_encoder(server_args, schedule_path, dist_init_method, rank): + try: + asyncio.run(run_encoder(server_args, schedule_path, dist_init_method, rank)) + except KeyboardInterrupt: + logger.info(f"Exit rank {rank}") + except Exception: + traceback.print_exc() + + +# Per-process encoder metrics collector. Set by +# runtime.launch_local_runtime (non-DP) and +# runtime.run_dp_worker (DP mode). None when metrics disabled. Kept +# here because MMEncoder GPU methods reference it directly. +encoder_metrics_collector: Optional[EncoderMetricsCollector] = None diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index f95eae089..ca060faf9 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -317,7 +317,6 @@ class GenerateReqInput: # 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[Any]] = None # Snapshot of encoder URLs at the time tokenizer-side computed # ``num_items_assigned``. encoder_urls: Optional[List[str]] = None @@ -1024,10 +1023,6 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True): need_wait_for_mm_inputs: Optional[bool] = None num_items_assigned: Optional[Dict[Modality, List[int]]] = None - # Pickled Optional[List[{"url": MultimodalDataInputItem, "modality": Modality}]] - # from MMReceiverBase._extract_url_data. "url" is ImageData.url, - # dict["url"] when present, or the original raw multimodal item. - mm_data_mooncake: Optional[PickleWrapper] = None # Encoder URL snapshot frozen at tokenizer-side dispatch time so that # encoder_idx assignments stay consistent in the scheduler subprocess. # Internal IPC only. @@ -1044,11 +1039,9 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True): cache_salt: Optional[str] = None def wrap_pickle_fields(self): - self.mm_data_mooncake = wrap_as_pickle(self.mm_data_mooncake) self.time_stats = wrap_as_pickle(self.time_stats) def unwrap_pickle_fields(self): - self.mm_data_mooncake = unwrap_from_pickle(self.mm_data_mooncake) self.time_stats = unwrap_from_pickle(self.time_stats) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 5e8e2f57d..2927824fa 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -709,6 +709,7 @@ def general_mm_embed_routine( if ( isinstance(precomputed_embeddings, torch.Tensor) and precomputed_embeddings.is_cuda + and not mm_item.keep_device_embedding ): mm_item.precomputed_embeddings = ( precomputed_embeddings.to( diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index b6f504088..96241c378 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -343,6 +343,8 @@ class MultimodalDataItem(msgspec.Struct, kw_only=True, dict=True, array_like=Tru # the precomputed embeddings, passed as final encoder embeddings # One and only one of the feature and precomputed_embeddings will be empty precomputed_embeddings: Optional[MultimodalDataValue] = None + # Keep precomputed_embeddings on GPU after use (EPD pool/GPU receive path) + keep_device_embedding: bool = False # Processor-owned tensors/arrays/scalars/transports. msgspec rejects a # precise union with multiple custom types, but accepts Ext-decoded values diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 58a547b9a..b1bba1fdb 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -82,7 +82,7 @@ from sglang.srt.disaggregation.decode import ( from sglang.srt.disaggregation.decode_kvcache_offload_manager import ( DecodeKVCacheOffloadManager, ) -from sglang.srt.disaggregation.encode_receiver import create_mm_receiver +from sglang.srt.disaggregation.encoder.receiver import create_mm_receiver from sglang.srt.disaggregation.prefill import ( PrefillBootstrapQueue, SchedulerDisaggregationPrefillMixin, @@ -4539,6 +4539,10 @@ class Scheduler( self._pending_chunked_abort_req = chunked_req # todo hisparse, release resources for abort requests in hisparse coordinator + # Abort requests still waiting for encoder embeddings (EPD language-only) + if self.mm_receiver is not None: + self.mm_receiver.abort_waiting_requests(recv_req) + # Delete requests in the waiting queue to_del = [] for i, req in enumerate(self.waiting_queue): diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 1adf77b56..1341cb31d 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -47,7 +47,7 @@ from fastapi import BackgroundTasks from sglang.srt.configs.model_config import ModelConfig from sglang.srt.constants import HEALTH_CHECK_RID_PREFIX -from sglang.srt.disaggregation.encode_receiver import create_mm_receiver +from sglang.srt.disaggregation.encoder.receiver import create_mm_receiver from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry @@ -671,7 +671,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Encoder Disaggregation self.encoder_bootstrap_server = None if self.server_args.language_only: - from sglang.srt.disaggregation.encode_receiver import ( + from sglang.srt.disaggregation.encoder.receiver import ( EncoderBootstrapServer, ) @@ -1409,7 +1409,6 @@ 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, encoder_urls=obj.encoder_urls, ) elif isinstance(obj, EmbeddingReqInput): diff --git a/python/sglang/srt/multimodal/processors/mimo_v2.py b/python/sglang/srt/multimodal/processors/mimo_v2.py index 022bb41cc..55e761fe8 100644 --- a/python/sglang/srt/multimodal/processors/mimo_v2.py +++ b/python/sglang/srt/multimodal/processors/mimo_v2.py @@ -1021,7 +1021,7 @@ class MiMoProcessor: "num_video_tokens": num_media_tokens_per_grid, "segment_audio_token_len": segment_audio_token_len, "segment_audio": segment_audio, - # Used by encode_server to trim audio_encoder output. + # Used by encoder.server to trim audio_encoder output. "audio_start_token_idx": audio_start_token_idx, } ) diff --git a/python/sglang/srt/utils/request_logger.py b/python/sglang/srt/utils/request_logger.py index 00cbaf529..b909de732 100644 --- a/python/sglang/srt/utils/request_logger.py +++ b/python/sglang/srt/utils/request_logger.py @@ -206,7 +206,6 @@ class RequestLogger: "image_data", "audio_data", "video_data", - "mm_data_mooncake", "lora_path", "sampling_params", } @@ -220,7 +219,6 @@ class RequestLogger: "image_data", "audio_data", "video_data", - "mm_data_mooncake", "lora_path", } out_skip_names = {"text", "output_ids", "embedding"} diff --git a/test/registered/observability/test_encoder_server_metrics.py b/test/registered/observability/test_encoder_server_metrics.py index 1c3854e67..5b51d2ffe 100644 --- a/test/registered/observability/test_encoder_server_metrics.py +++ b/test/registered/observability/test_encoder_server_metrics.py @@ -10,7 +10,7 @@ import zmq from prometheus_client.parser import text_string_to_metric_families from prometheus_client.samples import Sample -from sglang.srt.disaggregation.encode_server import MINIMUM_PNG_PICTURE_BASE64 +from sglang.srt.disaggregation.encoder.http_server import MINIMUM_PNG_PICTURE_BASE64 from sglang.srt.utils import kill_process_tree from sglang.srt.utils.network import get_zmq_socket_on_host from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci diff --git a/test/registered/unit/disaggregation/test_encode_receiver.py b/test/registered/unit/disaggregation/test_encode_receiver.py index f1d654a96..55515fddd 100644 --- a/test/registered/unit/disaggregation/test_encode_receiver.py +++ b/test/registered/unit/disaggregation/test_encode_receiver.py @@ -4,7 +4,7 @@ import unittest from array import array from types import SimpleNamespace -from sglang.srt.disaggregation.encode_receiver import MMReceiverBase +from sglang.srt.disaggregation.encoder.receiver import MMReceiverBase from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.sampling.sampling_params import SamplingParams from sglang.test.ci.ci_register import register_cpu_ci diff --git a/test/registered/unit/disaggregation/test_encode_server.py b/test/registered/unit/disaggregation/test_encode_server.py index ae45fd93d..dfa7b4ae6 100644 --- a/test/registered/unit/disaggregation/test_encode_server.py +++ b/test/registered/unit/disaggregation/test_encode_server.py @@ -1,30 +1,53 @@ +import asyncio import pickle import unittest from types import SimpleNamespace +from unittest.mock import AsyncMock, patch import numpy as np import torch -from sglang.srt.disaggregation.encode_receiver import EmbeddingData -from sglang.srt.disaggregation.encode_server import MMEncoder, _get_mm_grid_dim +from sglang.srt.disaggregation.encoder.preprocessor import EncoderPreprocessor +from sglang.srt.disaggregation.encoder.receiver import EmbeddingData +from sglang.srt.disaggregation.encoder.runtime import execute_encode_pipeline +from sglang.srt.disaggregation.encoder.server import ( + EncoderDelivery, + InternalError, + MMEncoder, + MooncakeDelivery, + ReqState, + SendDestination, + ZmqDelivery, + meta_registry, + rid_to_cond, + rid_to_receive_count, + rid_to_receive_endpoint, +) from sglang.srt.managers.schedule_batch import Modality from sglang.srt.utils.common import safe_pickle_loads from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=2, suite="base-a-test-cpu") -class TestKimiVLEPDGrid(unittest.TestCase): +class TestEncoderPreprocessorKimiGrid(CustomTestCase): @staticmethod - def _make_encoder(model_type="kimi_vl"): - encoder = MMEncoder.__new__(MMEncoder) - encoder.model_type = model_type - encoder.model_config = SimpleNamespace( + def _make_preprocessor(model_type="kimi_vl"): + preprocessor = EncoderPreprocessor.__new__(EncoderPreprocessor) + preprocessor.model_type = model_type + preprocessor.model_config = SimpleNamespace( hf_config=SimpleNamespace( vision_config=SimpleNamespace(merge_kernel_size=(2, 2)) ) ) - return encoder + preprocessor.image_processor = SimpleNamespace(merge_size=2) + preprocessor._model_preprocessor = None + return preprocessor + + @staticmethod + def _make_encoder(): + return MMEncoder.__new__(MMEncoder) def test_kimi_vl_prefers_and_normalizes_hw_grid(self): mm_inputs = { @@ -33,7 +56,7 @@ class TestKimiVLEPDGrid(unittest.TestCase): "grid_thws": torch.tensor([[1, 10, 15]]), } - grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_vl") + grid = self._make_preprocessor()._get_mm_grid_dim(mm_inputs, Modality.IMAGE) self.assertIsInstance(grid, torch.Tensor) torch.testing.assert_close(grid, torch.tensor([[40, 60]])) @@ -44,49 +67,51 @@ class TestKimiVLEPDGrid(unittest.TestCase): "grid_thws": np.array([[1, 10, 15]], dtype=np.int64), } - grid = _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_k25") + grid = self._make_preprocessor("kimi_k25")._get_mm_grid_dim( + mm_inputs, Modality.IMAGE + ) torch.testing.assert_close(grid, torch.tensor([[1, 10, 15]])) def test_kimi_vl_2d_grid_counting_and_slicing(self): + preprocessor = self._make_preprocessor() encoder = self._make_encoder() grids = torch.tensor([[40, 60], [20, 40]]) embedding = torch.arange(800 * 2).reshape(800, 2) self.assertEqual( - encoder.get_num_patches(grids[0], Modality.IMAGE), + preprocessor.get_num_patches(grids[0], Modality.IMAGE), 2400, ) self.assertEqual( - encoder.get_num_tokens(grids[0], Modality.IMAGE), + preprocessor.get_num_tokens(grids[0], Modality.IMAGE), 600, ) - slices = encoder.slice_embedding(embedding, grids, Modality.IMAGE) + slices = encoder.slice_embedding(embedding, [600, 200]) self.assertEqual([item.shape for item in slices], [(600, 2), (200, 2)]) torch.testing.assert_close(slices[0], embedding[:600]) torch.testing.assert_close(slices[1], embedding[600:]) def test_kimi_3d_grid_remains_supported(self): - encoder = self._make_encoder() + preprocessor = self._make_preprocessor() grid = torch.tensor([1, 40, 60]) - self.assertEqual(encoder.get_num_patches(grid, Modality.IMAGE), 2400) - self.assertEqual(encoder.get_num_tokens(grid, Modality.IMAGE), 600) + self.assertEqual(preprocessor.get_num_patches(grid, Modality.IMAGE), 2400) + self.assertEqual(preprocessor.get_num_tokens(grid, Modality.IMAGE), 600) def test_kimi_k25_3d_patch_counting_is_unchanged(self): - encoder = self._make_encoder("kimi_k25") + preprocessor = self._make_preprocessor("kimi_k25") grid = torch.tensor([2, 12, 16]) - self.assertEqual(encoder.get_num_patches(grid, Modality.IMAGE), 384) - self.assertEqual(encoder.get_num_tokens(grid, Modality.IMAGE), 48) + self.assertEqual(preprocessor.get_num_patches(grid, Modality.IMAGE), 384) + self.assertEqual(preprocessor.get_num_tokens(grid, Modality.IMAGE), 48) def test_grid_metadata_is_safe_to_deserialize(self): - grid = _get_mm_grid_dim( + grid = self._make_preprocessor()._get_mm_grid_dim( {"image_grid_hws": np.array([[40, 60]], dtype=np.int64)}, Modality.IMAGE, - "kimi_vl", ) embedding_data = EmbeddingData( req_id="test-request", @@ -104,5 +129,411 @@ class TestKimiVLEPDGrid(unittest.TestCase): torch.testing.assert_close(restored.grid_dim, torch.tensor([[40, 60]])) +class TestEncoderDelivery(CustomTestCase): + def test_contract_has_two_direct_implementations(self): + self.assertEqual(EncoderDelivery.__abstractmethods__, {"send", "release"}) + self.assertEqual( + set(EncoderDelivery.__subclasses__()), + { + MooncakeDelivery, + ZmqDelivery, + }, + ) + + def test_zmq_delivery_cleanup_is_configurable(self): + async def run(): + req_id = "test-zmq-delivery-cleanup" + rid_to_receive_endpoint[req_id] = {"127.0.0.1:1"} + rid_to_receive_count[req_id] = 1 + rid_to_cond[req_id] = asyncio.Condition() + state = ReqState(req_id) + encoder = SimpleNamespace() + + await ZmqDelivery(encoder, cleanup_receive_state=False).release(state) + self.assertIn(req_id, rid_to_receive_endpoint) + self.assertIn(req_id, rid_to_receive_count) + self.assertIn(req_id, rid_to_cond) + + await ZmqDelivery(encoder, cleanup_receive_state=True).release(state) + self.assertNotIn(req_id, rid_to_receive_endpoint) + self.assertNotIn(req_id, rid_to_receive_count) + self.assertNotIn(req_id, rid_to_cond) + + asyncio.run(run()) + + def test_preprocess_metadata_precedes_embedding(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + encoder._embedding_dims = {Modality.IMAGE: 8} + encoder._embedding_dtype = torch.float16 + encoder._element_size = 2 + first_state = encoder._acquire_encode_ref("req-0") + second_state = encoder._acquire_encode_ref("req-1") + ctx = SimpleNamespace( + req_id="req-0", + modality=Modality.IMAGE, + items_per_req=[1, 2], + preprocess_result=SimpleNamespace( + token_counts=[2, 3, 4], + grid_thw=[[1, 2, 3], [1, 4, 5], [1, 6, 7]], + ), + ) + requests = [ + {"req_id": "req-0", "num_parts": 2, "part_idx": 0}, + {"req_id": "req-1", "num_parts": 2, "part_idx": 1}, + ] + + publish = AsyncMock() + with patch.object(meta_registry, "publish", publish): + await encoder._publish_preprocess_metadata(ctx, requests) + + self.assertIs(encoder.req_states["req-0"], first_state) + self.assertIs(encoder.req_states["req-1"], second_state) + self.assertEqual(first_state.embedding_data.shape, [2, 8]) + self.assertEqual(second_state.embedding_data.shape, [7, 8]) + self.assertEqual(first_state.embedding_data.grid_dim, [[1, 2, 3]]) + self.assertEqual( + second_state.embedding_data.grid_dim, + [[1, 4, 5], [1, 6, 7]], + ) + self.assertEqual(first_state.embedding_data.dtype, torch.float16) + self.assertEqual(second_state.embedding_data.dtype, torch.float16) + self.assertFalse(first_state.embedding_ready.is_set()) + self.assertFalse(second_state.embedding_ready.is_set()) + self.assertEqual( + publish.await_args_list, + [ + unittest.mock.call("req-0", 32, 2, 8), + unittest.mock.call("req-1", 112, 7, 8), + ], + ) + await encoder._release_encode_ref(first_state) + await encoder._release_encode_ref(second_state) + + asyncio.run(run()) + + def test_mooncake_embedding_is_ready_only_after_cuda_sync(self): + class FakeCudaEmbedding: + shape = (2, 4) + dtype = torch.float16 + nbytes = 16 + is_cuda = True + device = "cuda:0" + + def __getitem__(self, key): + return self + + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + events = [] + encoder._stage_embedding = lambda mm_data: events.append("ready") + ctx = SimpleNamespace( + req_id="req", + modality=Modality.IMAGE, + items_per_req=[1], + preprocess_result=SimpleNamespace( + token_counts=[2], + grid_thw=[[1, 2, 3]], + ), + aux_data={}, + use_global_cache=True, + ) + stream = SimpleNamespace(synchronize=lambda: events.append("sync")) + + with patch.object(torch.cuda, "current_stream", return_value=stream): + encoder._stage_embeddings( + ctx, + [{"req_id": "req", "num_parts": 1, "part_idx": 0}], + FakeCudaEmbedding(), + keep_on_gpu=True, + ) + + self.assertEqual(events, ["sync", "ready"]) + + def test_stage_embedding_does_not_resurrect_missing_state(self): + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + + with self.assertRaisesRegex( + InternalError, "No request state exists while encoding request: req" + ): + encoder._stage_embedding( + EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding=torch.ones((1, 1)), + ) + ) + + self.assertNotIn("req", encoder.req_states) + + def test_stage_embedding_requires_active_encode(self): + encoder = MMEncoder.__new__(MMEncoder) + state = ReqState("req") + encoder.req_states = {"req": state} + + with self.assertRaisesRegex( + InternalError, "Request state has no active encode work: req" + ): + encoder._stage_embedding( + EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding=torch.ones((1, 1)), + ) + ) + + self.assertIs(encoder.req_states["req"], state) + self.assertIsNone(state.embedding_data) + self.assertFalse(state.embedding_ready.is_set()) + + def test_release_during_encode_is_deferred_without_resurrecting_state(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + encoder.delivery = SimpleNamespace(release=AsyncMock()) + + state = encoder._acquire_encode_ref("req") + state.embedding_data = EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding_shape=[1, 1], + dtype=torch.float32, + ) + + discard = AsyncMock() + with patch.object(meta_registry, "discard", discard): + await encoder.release_request("req") + self.assertTrue(state.release_requested) + self.assertIn("req", encoder.req_states) + encoder.delivery.release.assert_not_awaited() + + embedding = torch.ones((1, 1)) + encoder._stage_embedding( + EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding=embedding, + ) + ) + await encoder._release_encode_ref(state) + + encoder.delivery.release.assert_awaited_once_with(state) + discard.assert_awaited_once_with("req") + self.assertIsNone(state.embedding_data) + self.assertNotIn("req", encoder.req_states) + + asyncio.run(run()) + + def test_error_metadata_survives_buffer_release_for_waiter(self): + async def run(): + req_id = "test-error-metadata-waiter" + await meta_registry.discard(req_id) + + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + encoder.delivery = SimpleNamespace(release=AsyncMock()) + state = ReqState( + req_id, + EmbeddingData( + req_id, + 1, + 0, + None, + Modality.IMAGE, + error_msg="encode failed", + ), + ) + state.embedding_ready.set() + encoder.req_states[req_id] = state + + waiter = asyncio.create_task(meta_registry.wait(req_id)) + await asyncio.sleep(0) + try: + await meta_registry.publish(req_id, 0, 0, 0, error="encode failed") + await encoder.release_request(req_id, preserve_metadata=True) + meta = await asyncio.wait_for(waiter, timeout=1) + self.assertEqual(meta, {"error": "encode failed"}) + self.assertNotIn(req_id, encoder.req_states) + finally: + if not waiter.done(): + waiter.cancel() + await meta_registry.discard(req_id) + + asyncio.run(run()) + + def test_zmq_pipeline_sends_only_after_encode_completes(self): + async def run(): + events = [] + finish_encode = asyncio.Event() + encoder = MMEncoder.__new__(MMEncoder) + encoder.transfer_backend = "zmq_to_tokenizer" + + async def encode(**kwargs): + events.append("metadata_published") + await finish_encode.wait() + events.append("encode_completed") + return 16, 2, 4, None, None + + async def send(**kwargs): + events.append("send") + return True + + async def release_request(req_id, **kwargs): + events.append("release") + + encoder.encode = AsyncMock(side_effect=encode) + encoder.send = AsyncMock(side_effect=send) + encoder.release_request = AsyncMock(side_effect=release_request) + + publish = AsyncMock() + request = { + "req_id": "req", + "mm_items": ["item"], + "modality": "image", + "num_parts": 1, + "part_idx": 0, + "prefill_host": "127.0.0.1", + "embedding_port": 1234, + } + with patch.object(meta_registry, "publish", publish): + task = asyncio.create_task( + execute_encode_pipeline(encoder, None, request) + ) + await asyncio.sleep(0) + self.assertEqual(events, ["metadata_published"]) + encoder.send.assert_not_awaited() + + finish_encode.set() + self.assertIsNone(await task) + + self.assertEqual( + events, + ["metadata_published", "encode_completed", "send", "release"], + ) + publish.assert_awaited_once_with("req", 16, 2, 4) + + asyncio.run(run()) + + def test_send_waits_for_embedding_published_by_encode(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + state = encoder._acquire_encode_ref("req") + state.embedding_data = EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding_shape=[1, 1], + dtype=torch.float32, + ) + + delivered = [] + + async def send(current_state, destination): + delivered.append(await encoder._wait_for_embedding(current_state)) + + encoder.delivery = SimpleNamespace( + send=AsyncMock(side_effect=send), + release=AsyncMock(), + ) + send_task = asyncio.create_task( + encoder.send_to_destination(state, SendDestination("127.0.0.1:1")) + ) + await asyncio.sleep(0) + self.assertFalse(send_task.done()) + + embedding = torch.ones((1, 1)) + encoder._stage_embedding( + EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding=embedding, + ) + ) + await send_task + await encoder._release_encode_ref(state) + + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder.release_request("req") + + self.assertEqual(len(delivered), 1) + self.assertIs(delivered[0].embedding, embedding) + + asyncio.run(run()) + + def test_release_waits_for_send_then_clears_embedding(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + + send_started = asyncio.Event() + finish_send = asyncio.Event() + + async def send(state, destination): + send_started.set() + await finish_send.wait() + + embedding_seen_by_release = [] + + async def release(state): + embedding_seen_by_release.append(state.embedding_data.embedding) + + encoder.delivery = SimpleNamespace( + send=AsyncMock(side_effect=send), + release=AsyncMock(side_effect=release), + ) + embedding = torch.ones((1, 1)) + state = ReqState( + "req", + EmbeddingData("req", 1, 0, None, Modality.IMAGE, embedding=embedding), + ) + state.embedding_ready.set() + encoder.req_states[state.req_id] = state + + send_task = asyncio.create_task( + encoder.send_to_destination(state, SendDestination("127.0.0.1:1")) + ) + await send_started.wait() + release_task = asyncio.create_task(encoder.release_request("req")) + await asyncio.sleep(0) + + encoder.delivery.release.assert_not_awaited() + self.assertIs(state.embedding_data.embedding, embedding) + + finish_send.set() + await send_task + await release_task + + encoder.delivery.release.assert_awaited_once_with(state) + self.assertEqual(len(embedding_seen_by_release), 1) + self.assertIs(embedding_seen_by_release[0], embedding) + self.assertIsNone(state.embedding_data) + self.assertNotIn("req", encoder.req_states) + + asyncio.run(run()) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/disaggregation/test_encoder_health.py b/test/registered/unit/disaggregation/test_encoder_health.py index b9b64df85..c5d3872b7 100644 --- a/test/registered/unit/disaggregation/test_encoder_health.py +++ b/test/registered/unit/disaggregation/test_encoder_health.py @@ -3,7 +3,8 @@ import sys import pytest -from sglang.srt.disaggregation import encode_server +from sglang.srt.disaggregation.encoder import http_server +from sglang.srt.managers.schedule_batch import Modality from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="base-a-test-cpu") @@ -17,18 +18,27 @@ class _FakeEncoder: self.encode_dispatch_lock = asyncio.Lock() self.encode_calls = [] + def has_pending_embeddings(self): + return bool(self.embedding_to_send) + + def supports_modality(self, modality): + return modality == Modality.IMAGE + async def encode(self, **kwargs): self.encode_calls.append(kwargs) return 1, 1, 1, None, None + async def release_request(self, _req_id): + return None + def _install_tp_encoder(monkeypatch, encoder): broadcasts = [] - monkeypatch.setattr(encode_server, "dp_dispatcher", None) - monkeypatch.setattr(encode_server, "encoder", encoder) - monkeypatch.setattr(encode_server, "send_sockets", [object()]) + monkeypatch.setattr(http_server, "dp_dispatcher", None) + monkeypatch.setattr(http_server, "encoder", encoder) + monkeypatch.setattr(http_server, "send_sockets", [object()]) monkeypatch.setattr( - encode_server, + http_server, "sock_send", lambda socket, payload: broadcasts.append((socket, payload)), ) @@ -41,7 +51,7 @@ def test_health_encode_waits_for_collective_dispatch_lock(monkeypatch): broadcasts = _install_tp_encoder(monkeypatch, encoder) await encoder.encode_dispatch_lock.acquire() - task = asyncio.create_task(encode_server.health_generate()) + task = asyncio.create_task(http_server.health_generate()) await asyncio.sleep(0) assert broadcasts == [] assert encoder.encode_calls == [] @@ -61,7 +71,7 @@ def test_health_encode_rechecks_busy_state_after_waiting(monkeypatch): broadcasts = _install_tp_encoder(monkeypatch, encoder) await encoder.encode_dispatch_lock.acquire() - task = asyncio.create_task(encode_server.health_generate()) + task = asyncio.create_task(http_server.health_generate()) await asyncio.sleep(0) encoder.embedding_to_send["real-request"] = object() encoder.encode_dispatch_lock.release() diff --git a/test/registered/unit/disaggregation/test_encoder_scheduler.py b/test/registered/unit/disaggregation/test_encoder_scheduler.py index 1430f31b0..1e31694a4 100644 --- a/test/registered/unit/disaggregation/test_encoder_scheduler.py +++ b/test/registered/unit/disaggregation/test_encoder_scheduler.py @@ -3,7 +3,7 @@ import sys import pytest -from sglang.srt.disaggregation.encode_server import ( +from sglang.srt.disaggregation.encoder.runtime import ( EncoderScheduler, PendingRequest, _resolve_encoder_batch_policy, diff --git a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py index 022b92c4c..5cfc958c5 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -15,14 +15,18 @@ import zmq.asyncio from fastapi import HTTPException from PIL import Image -from sglang.srt.disaggregation.encode_receiver import ( +from sglang.srt.disaggregation.encoder.preprocessor import ( + EncoderPreprocessor, + EncoderPreprocessResult, +) +from sglang.srt.disaggregation.encoder.receiver import ( EmbeddingData, MMReceiverHTTP, MultiModalEmbeddingData, _encoder_media_item, _select_mm_processor_prompt, ) -from sglang.srt.disaggregation.encode_server import MMEncoder, _get_mm_grid_dim +from sglang.srt.disaggregation.encoder.server import MMEncoder from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.managers.tokenizer_manager import ( _reject_missing_dispatched_encoder_embedding, @@ -220,16 +224,19 @@ def test_epd_allows_local_processing_when_request_was_not_dispatched(): def _encoder(model_type="kimi_k3"): encoder = MMEncoder.__new__(MMEncoder) encoder.model_type = model_type - encoder.model_config = SimpleNamespace( + preprocessor = EncoderPreprocessor.__new__(EncoderPreprocessor) + preprocessor.model_type = model_type + preprocessor.model_config = SimpleNamespace( hf_config=SimpleNamespace( vision_config=SimpleNamespace(merge_kernel_size=(2, 2)) ) ) - encoder.encoder_media_processor_config = ( + preprocessor.encoder_media_processor_config = ( KimiK3ForConditionalGeneration.encoder_media_processor_config if model_type == "kimi_k3" else EncoderMediaProcessorConfig() ) + encoder.preprocessor = preprocessor return encoder @@ -237,11 +244,11 @@ def test_kimi_k3_encoder_normalizes_pillow_images_to_media_dicts(): image = Image.new("RGB", (2, 2)) encoder = _encoder() - assert encoder._grid_count_per_leaf( + assert encoder.preprocessor._grid_count_per_leaf( [image, {"type": "image", "image": [image, image]}], Modality.IMAGE ) == [1, 2] - normalized = encoder._normalize_kimi_encoder_images( + normalized = encoder.preprocessor._normalize_kimi_encoder_images( [image, {"type": "image", "image": [image, image]}] ) assert len(normalized) == 3 @@ -258,14 +265,15 @@ def test_kimi_k3_encoder_passes_media_dicts_to_image_processor(): return {"pixel_values": torch.ones(1, 3), "grid_thws": [[1, 1, 1]]} encoder = _encoder() - encoder.image_processor = image_processor - encoder.vision_config = {"image": {"return_tensors": "pt"}} - encoder._flatten_and_load_images = AsyncMock(return_value=[image]) - encoder.preproc_executor = ThreadPoolExecutor(max_workers=1) + preprocessor = encoder.preprocessor + preprocessor.image_processor = image_processor + preprocessor.vision_config = {"image": {"return_tensors": "pt"}} + preprocessor._flatten_and_load_images = AsyncMock(return_value=[image]) + preprocessor.preproc_executor = ThreadPoolExecutor(max_workers=1) try: - output = asyncio.run(encoder._process_image_items([image], None)) + output = asyncio.run(preprocessor._process_image_items([image], None)) finally: - encoder.preproc_executor.shutdown() + preprocessor.preproc_executor.shutdown() assert "pixel_values" in output assert output["original_image_sizes"] == [[3, 2]] @@ -354,27 +362,28 @@ def test_kimi_k3_epd_model_preprocessor_receives_image_processor(): return prepare_kimi_k3_encoder_inputs(mm_data, image_processor) encoder = _encoder() - encoder.image_processor = image_processor - encoder.use_image_processor_gpu = False - encoder.vision_config = {"image": {"return_tensors": "pt"}} - encoder._flatten_and_load_images = AsyncMock(return_value=[image]) - encoder.preproc_executor = ThreadPoolExecutor(max_workers=1) + preprocessor = encoder.preprocessor + preprocessor.image_processor = image_processor + preprocessor.use_image_processor_gpu = False + preprocessor.vision_config = {"image": {"return_tensors": "pt"}} + preprocessor._flatten_and_load_images = AsyncMock(return_value=[image]) + preprocessor.preproc_executor = ThreadPoolExecutor(max_workers=1) try: with patch( - "sglang.srt.disaggregation.encode_server.get_parallel", + "sglang.srt.disaggregation.encoder.preprocessor.get_parallel", return_value=SimpleNamespace(attn_tp_rank=0, attn_tp_size=1), ): output = asyncio.run( - encoder._process_image_items([image], model_preprocessor) + preprocessor._process_image_items([image], model_preprocessor) ) finally: - encoder.preproc_executor.shutdown() + preprocessor.preproc_executor.shutdown() assert len(calls) == 1 assert calls[0][0][0] == {"type": "image", "image": image} assert calls[0][1:] == ( Modality.IMAGE, - encoder.vision_config, + preprocessor.vision_config, image_processor, False, ) @@ -481,13 +490,13 @@ def test_kimi_k3_epd_selects_matching_jpeg_decode_mode( ): expected = torch.zeros((3, 2, 3), dtype=torch.uint8) encoder = _encoder() - encoder.use_image_processor_gpu = use_image_processor_gpu + encoder.preprocessor.use_image_processor_gpu = use_image_processor_gpu with patch( - "sglang.srt.disaggregation.encode_server.load_image", + "sglang.srt.disaggregation.encoder.preprocessor.load_image", return_value=(expected, None), ) as load: - output = encoder._load_single_item(b"jpeg", Modality.IMAGE) + output = encoder.preprocessor._load_single_item(b"jpeg", Modality.IMAGE) assert output is expected load.assert_called_once_with(b"jpeg", expected_decode_mode) @@ -498,13 +507,13 @@ def test_kimi_k3_epd_verifies_content_hash_before_decode(): digest = snapshot_media(payload).content_digest expected = torch.zeros((3, 2, 3), dtype=torch.uint8) encoder = _encoder() - encoder.use_image_processor_gpu = False + encoder.preprocessor.use_image_processor_gpu = False with patch( - "sglang.srt.disaggregation.encode_server.load_image", + "sglang.srt.disaggregation.encoder.preprocessor.load_image", return_value=(expected, None), ) as load: - output = encoder._load_single_item( + output = encoder.preprocessor._load_single_item( {"url": payload, "content_hash": digest}, Modality.IMAGE ) @@ -589,8 +598,9 @@ def test_kimi_k3_encoder_prefers_grid_thws_and_uses_temporal_pool_length(): stale_grid = torch.tensor([[1, 2, 2]]) mm_inputs = {"grid_thws": grid_thws, "image_grid_thw": stale_grid} - assert _get_mm_grid_dim(mm_inputs, Modality.IMAGE, "kimi_k3") is grid_thws - assert _encoder().get_num_tokens(grid_thws[0], Modality.IMAGE) == 24 + preprocessor = _encoder().preprocessor + assert preprocessor._get_mm_grid_dim(mm_inputs, Modality.IMAGE) is grid_thws + assert preprocessor.get_num_tokens(grid_thws[0], Modality.IMAGE) == 24 def test_kimi_k3_encoder_splits_cross_request_batch_into_single_grid_items(): @@ -606,12 +616,14 @@ def test_kimi_k3_encoder_splits_cross_request_batch_into_single_grid_items(): output = encoder._encode_missing( feature, - {"pixel_values": feature, "grid_thws": grid_thws}, + EncoderPreprocessResult( + mm_inputs={"pixel_values": feature, "grid_thws": grid_thws}, + grid_thw=grid_thws, + token_counts=[1, 2, 2], + ), indices=[2, 0, 1], modality=Modality.IMAGE, get_feature_fn=get_feature_fn, - grid_thw=grid_thws, - keep_on_gpu=True, ) items = captured["items"] @@ -643,7 +655,7 @@ def test_encoder_preprocessed_items_follow_dp_owner_selection_order(): {"pixel_values": [item.feature for item in items], "grid_thws": grid_thws}, mm_items=items, ) - embeddings = torch.arange(4, dtype=torch.float32).reshape(4, 1) + embeddings = torch.arange(3, dtype=torch.float32).reshape(3, 1) captured = {} def get_feature_fn(selected_items): @@ -652,17 +664,19 @@ def test_encoder_preprocessed_items_follow_dp_owner_selection_order(): output = encoder._encode_missing( mm_inputs["pixel_values"], - mm_inputs, + EncoderPreprocessResult( + mm_inputs=mm_inputs, + grid_thw=grid_thws, + token_counts=[1, 2, 2], + ), indices=[2, 0], modality=Modality.IMAGE, get_feature_fn=get_feature_fn, - grid_thw=grid_thws, - keep_on_gpu=True, ) assert captured["items"] == [items[2], items[0]] assert [part.shape[0] for part in output] == [2, 1] - torch.testing.assert_close(torch.cat(output), embeddings[:3]) + torch.testing.assert_close(torch.cat(output), embeddings) def test_encoder_preprocessed_items_hash_individually(): @@ -767,6 +781,8 @@ def test_epd_encoder_reuses_scheduler_zmq_peer(): ) with config_override as server_args: encoder.server_args = server_args + encoder.transfer_backend = "zmq_to_scheduler" + encoder.use_mooncake = False encoder.send_timeout = 3 encoder.context = context encoder.scheduler_send_sockets = {} @@ -841,6 +857,8 @@ def test_epd_encoder_pipelines_zero_copy_sends_per_peer(): ) with config_override as server_args: encoder.server_args = server_args + encoder.transfer_backend = "zmq_to_scheduler" + encoder.use_mooncake = False encoder.send_timeout = 1 encoder.context = FakeContext(socket) encoder.scheduler_send_sockets = {} diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index 3d447b66e..495c0507b 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -108,7 +108,7 @@ _CONFIGURED_SIZE_CALL_SITES = { "the consumer count is configured fan-out arithmetic (tp_size // " "dp_size), which is what the record answered before" ), - ("srt/disaggregation/encode_server.py", "configured_tp_size"): ( + ("srt/disaggregation/encoder/runtime.py", "configured_tp_size"): ( "the encode server's launch entry sizes its workers before it has " "spawned any of them" ), diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py index a9383f248..377429cf0 100644 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ b/test/registered/unit/test_publish_precedes_bag_reads.py @@ -77,8 +77,8 @@ _KNOWN_ENTRIES = frozenset( "run_data_parallel_controller_process", ), ("srt/ray/scheduler_actor.py", "__init__"), - ("srt/disaggregation/encode_server.py", "__init__"), - ("srt/disaggregation/encode_server.py", "launch_server"), + ("srt/disaggregation/encoder/server.py", "__init__"), + ("srt/disaggregation/encoder/http_server.py", "launch_server"), ("srt/managers/tokenizer_manager.py", "__init__"), ("srt/entrypoints/engine.py", "_launch_subprocesses"), (