[EPD] Cross-request batching for image/audio encoder (#25964)
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import contextlib
|
||||
import ctypes
|
||||
import logging
|
||||
import multiprocessing as mp
|
||||
@@ -7,6 +8,7 @@ import os
|
||||
import pickle
|
||||
import time
|
||||
import traceback
|
||||
from collections import defaultdict
|
||||
from http import HTTPStatus
|
||||
from typing import Dict, List, Optional, Set, Tuple, Union
|
||||
|
||||
@@ -78,9 +80,12 @@ rid_to_err_msg: Dict[str, str] = dict()
|
||||
cond_dict_lock = asyncio.Lock()
|
||||
rid_to_cond: Dict[str, asyncio.Condition] = {}
|
||||
|
||||
use_image_processor_gpu = (
|
||||
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
|
||||
)
|
||||
use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
|
||||
|
||||
ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get()
|
||||
# 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):
|
||||
@@ -223,6 +228,8 @@ class MMEncoder:
|
||||
use_image_processor_gpu and not server_args.disable_fast_image_processor
|
||||
)
|
||||
self._build_vision_config(server_args.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),
|
||||
@@ -341,6 +348,44 @@ class MMEncoder:
|
||||
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.
|
||||
@@ -440,7 +485,6 @@ class MMEncoder:
|
||||
data,
|
||||
modality: Modality,
|
||||
frame_count_limit=None,
|
||||
audio_sample_rate: Optional[int] = None,
|
||||
discard_alpha_channel=True,
|
||||
):
|
||||
"""
|
||||
@@ -463,7 +507,7 @@ class MMEncoder:
|
||||
elif modality == Modality.VIDEO:
|
||||
return load_video(data, frame_count_limit)
|
||||
elif modality == Modality.AUDIO:
|
||||
return load_audio(data, audio_sample_rate)
|
||||
return load_audio(data, self.model_audio_sr)
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Error while loading data {data}: {e}")
|
||||
@@ -500,6 +544,11 @@ class MMEncoder:
|
||||
((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(
|
||||
@@ -631,10 +680,18 @@ class MMEncoder:
|
||||
return slices
|
||||
|
||||
def _calculate_hashes_from_features(
|
||||
self, mm_feature: torch.Tensor, grid_thw: List, modality: Modality
|
||||
self, mm_feature, grid_thw: List, modality: Modality
|
||||
) -> List[str]:
|
||||
"""CPU Task: Compute hashes based on processed feature patches."""
|
||||
hashes, offset = [], 0
|
||||
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)
|
||||
@@ -647,7 +704,7 @@ class MMEncoder:
|
||||
|
||||
async def _encode_missing(
|
||||
self,
|
||||
mm_feature: torch.Tensor,
|
||||
mm_feature,
|
||||
mm_inputs: dict,
|
||||
indices: List[int],
|
||||
modality: Modality = Modality.IMAGE,
|
||||
@@ -658,23 +715,34 @@ class MMEncoder:
|
||||
"""
|
||||
grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type)
|
||||
|
||||
# 1. Slice mm_feature to get only the patches for missing mm items
|
||||
sub_feature_list = []
|
||||
offsets = [0]
|
||||
curr = 0
|
||||
for g in grid_thw:
|
||||
curr += self.get_num_patches(g, modality)
|
||||
offsets.append(curr)
|
||||
|
||||
for idx in indices:
|
||||
sub_feature_list.append(mm_feature[offsets[idx] : offsets[idx + 1]])
|
||||
|
||||
sub_feature = torch.cat(sub_feature_list, dim=0)
|
||||
# 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 g in grid_thw:
|
||||
curr += self.get_num_patches(g, modality)
|
||||
offsets.append(curr)
|
||||
for idx in indices:
|
||||
sub_feature_list.append(mm_feature[offsets[idx] : offsets[idx + 1]])
|
||||
sub_feature = torch.cat(sub_feature_list, dim=0)
|
||||
|
||||
mm_item = MultimodalDataItem.from_dict(
|
||||
{
|
||||
"modality": modality,
|
||||
"feature": _convert(sub_feature),
|
||||
"feature": (
|
||||
sub_feature
|
||||
if isinstance(sub_feature, list)
|
||||
else _convert(sub_feature)
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -710,6 +778,15 @@ class MMEncoder:
|
||||
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)."
|
||||
)
|
||||
|
||||
# Step 1: Rank 0 checks global cache and broadcasts hit/miss mask to all ranks.
|
||||
if self.rank == 0:
|
||||
if hashes is None:
|
||||
@@ -852,7 +929,8 @@ class MMEncoder:
|
||||
|
||||
async def _flatten_and_load_audios(self, mm_items):
|
||||
"""
|
||||
Flatten mm_items structure, load audios concurrently, and restore original structure.
|
||||
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)
|
||||
|
||||
@@ -893,6 +971,28 @@ class MMEncoder:
|
||||
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/K25 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_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
|
||||
@@ -1039,18 +1139,21 @@ class MMEncoder:
|
||||
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(mm_items, Modality.AUDIO, self.vision_config)
|
||||
return model_preprocessor(audios, Modality.AUDIO, self.vision_config)
|
||||
|
||||
if not self.audio_processor:
|
||||
raise ValueError("No audio processor available")
|
||||
|
||||
audios = await self._flatten_and_load_audios(mm_items)
|
||||
audio_config = self.vision_config.get("audio", {})
|
||||
processor_input = self.audio_processor.feature_extractor(audios, **audio_config)
|
||||
processor_input["feature_attention_mask"] = processor_input.pop(
|
||||
"attention_mask"
|
||||
)
|
||||
# convert to same format as image/video
|
||||
input_lengths = torch.tensor(
|
||||
processor_input["feature_attention_mask"].sum(-1), dtype=torch.long
|
||||
)
|
||||
@@ -1225,6 +1328,122 @@ class MMEncoder:
|
||||
logger.debug(f"Created error EmbeddingData: {mm_data}")
|
||||
return 0, 0, 0, error_msg, error_code
|
||||
|
||||
async def encode_request(self, req: dict, modality: Modality):
|
||||
"""Single-request encode dispatcher: picks cache vs no-cache path."""
|
||||
if self.mm_global_cache is not None:
|
||||
return await self.encode_with_global_cache(
|
||||
mm_items=req["mm_items"],
|
||||
modality=modality,
|
||||
req_id=req["req_id"],
|
||||
num_parts=req["num_parts"],
|
||||
part_idx=req["part_idx"],
|
||||
hashes=req.get("hashes"),
|
||||
)
|
||||
return await self.encode(
|
||||
mm_items=req["mm_items"],
|
||||
modality=modality,
|
||||
req_id=req["req_id"],
|
||||
num_parts=req["num_parts"],
|
||||
part_idx=req["part_idx"],
|
||||
)
|
||||
|
||||
async def batch_encode(
|
||||
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/K25 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)
|
||||
|
||||
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 = await 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()
|
||||
# No aux_data here: batch_encode only handles IMAGE/AUDIO
|
||||
# (_BATCHABLE_MODALITIES), and _build_mm_aux_data only extracts
|
||||
# video-meta fields — which never appear in image/audio mm_inputs.
|
||||
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)
|
||||
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,
|
||||
)
|
||||
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
|
||||
@@ -1396,9 +1615,239 @@ class EncoderProfiler:
|
||||
return True, None
|
||||
|
||||
|
||||
app = FastAPI()
|
||||
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}
|
||||
|
||||
|
||||
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,
|
||||
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.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(
|
||||
f"EncoderScheduler started with max_batch_size={self.max_batch_size}"
|
||||
)
|
||||
|
||||
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()]
|
||||
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()
|
||||
for sock in self.send_sockets:
|
||||
sock.send_pyobj(
|
||||
{
|
||||
"type": "batch_encode",
|
||||
"modality": modality.name,
|
||||
"requests": requests,
|
||||
"enter_time": start,
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
|
||||
|
||||
try:
|
||||
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:
|
||||
for p in group:
|
||||
req = p.request
|
||||
try:
|
||||
for sock in self.send_sockets:
|
||||
sock.send_pyobj(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
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _lifespan(app: FastAPI):
|
||||
global encoder_scheduler
|
||||
if encoder is not None:
|
||||
encoder_scheduler = EncoderScheduler(
|
||||
encoder, send_sockets, max_batch_size=ENCODER_MAX_BATCH_SIZE
|
||||
)
|
||||
encoder_scheduler.start()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if encoder_scheduler is not None:
|
||||
await encoder_scheduler.stop()
|
||||
|
||||
|
||||
app = FastAPI(lifespan=_lifespan)
|
||||
|
||||
|
||||
async def run_encoder(
|
||||
@@ -1414,24 +1863,15 @@ async def run_encoder(
|
||||
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:
|
||||
if encoder.mm_global_cache is not None:
|
||||
await encoder.encode_with_global_cache(
|
||||
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", None),
|
||||
)
|
||||
else:
|
||||
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"],
|
||||
)
|
||||
await encoder.encode_request(
|
||||
request, Modality.from_str(request["modality"])
|
||||
)
|
||||
|
||||
|
||||
def launch_encoder(server_args, schedule_path, dist_init_method, rank):
|
||||
@@ -1489,30 +1929,27 @@ async def handle_encode_request(request: dict):
|
||||
encoder.background_tasks.add(task)
|
||||
task.add_done_callback(encoder.background_tasks.discard)
|
||||
|
||||
# broadcast request
|
||||
request.update({"enter_time": time.time()})
|
||||
for socket in send_sockets:
|
||||
socket.send_pyobj(request)
|
||||
if encoder.mm_global_cache is not None:
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
await encoder.encode_with_global_cache(
|
||||
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", None),
|
||||
modality = Modality.from_str(request["modality"])
|
||||
if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
|
||||
try:
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
await encoder_scheduler.submit(request)
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
return ORJSONResponse(
|
||||
status_code=HTTPStatus.GATEWAY_TIMEOUT,
|
||||
content={
|
||||
"status": "error",
|
||||
"message": "encoder batch timed out",
|
||||
"req_id": req_id,
|
||||
},
|
||||
)
|
||||
)
|
||||
else:
|
||||
for socket in send_sockets:
|
||||
socket.send_pyobj(request)
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
await encoder.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"],
|
||||
)
|
||||
await encoder.encode_request(request, modality)
|
||||
)
|
||||
|
||||
if error_msg:
|
||||
|
||||
@@ -712,6 +712,9 @@ class Envs:
|
||||
SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0)
|
||||
SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0)
|
||||
SGLANG_ENCODER_DISPATCH_MIN_ITEMS = EnvInt(2)
|
||||
SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU = EnvBool(False)
|
||||
SGLANG_ENCODER_MAX_BATCH_SIZE = EnvInt(8)
|
||||
SGLANG_ENCODER_REQ_TIMEOUT = EnvFloat(180.0)
|
||||
|
||||
# Elastic EP Backup Port
|
||||
SGLANG_BACKUP_PORT_BASE = EnvInt(10000)
|
||||
|
||||
@@ -28,6 +28,7 @@ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import (
|
||||
Qwen2_5_VLVisionConfig,
|
||||
)
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
Modality,
|
||||
MultimodalDataItem,
|
||||
@@ -1815,9 +1816,7 @@ class MiMoV2Processor(BaseMultimodalProcessor):
|
||||
self.video_end_token_id = self._require_config_value(
|
||||
processor_config, "video_end_token_id"
|
||||
)
|
||||
self.use_image_processor_gpu = (
|
||||
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
|
||||
)
|
||||
self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
|
||||
device = server_args.device if self.use_image_processor_gpu else None
|
||||
|
||||
self.mimo_processor = MiMoProcessor(
|
||||
|
||||
Reference in New Issue
Block a user