[EPD] Cross-request batching for image/audio encoder (#25964)

This commit is contained in:
Zhonghua Deng
2026-05-26 15:38:59 +08:00
committed by GitHub
parent d34d4d9f5f
commit dabdd91ef3
3 changed files with 505 additions and 66 deletions
+487 -50
View File
@@ -1,5 +1,6 @@
import asyncio import asyncio
import concurrent.futures import concurrent.futures
import contextlib
import ctypes import ctypes
import logging import logging
import multiprocessing as mp import multiprocessing as mp
@@ -7,6 +8,7 @@ import os
import pickle import pickle
import time import time
import traceback import traceback
from collections import defaultdict
from http import HTTPStatus from http import HTTPStatus
from typing import Dict, List, Optional, Set, Tuple, Union 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() cond_dict_lock = asyncio.Lock()
rid_to_cond: Dict[str, asyncio.Condition] = {} rid_to_cond: Dict[str, asyncio.Condition] = {}
use_image_processor_gpu = ( use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
) 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): class MMError(Exception):
@@ -223,6 +228,8 @@ class MMEncoder:
use_image_processor_gpu and not server_args.disable_fast_image_processor use_image_processor_gpu and not server_args.disable_fast_image_processor
) )
self._build_vision_config(server_args.mm_process_config) 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( init_distributed_environment(
backend=get_default_distributed_backend(self.device), backend=get_default_distributed_backend(self.device),
@@ -341,6 +348,44 @@ class MMEncoder:
logger.info(f"Global cache embedding dims: {dims}") logger.info(f"Global cache embedding dims: {dims}")
return 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): def _build_vision_config(self, mm_process_config):
""" """
Validate vision config, used for image/video/audio. Validate vision config, used for image/video/audio.
@@ -440,7 +485,6 @@ class MMEncoder:
data, data,
modality: Modality, modality: Modality,
frame_count_limit=None, frame_count_limit=None,
audio_sample_rate: Optional[int] = None,
discard_alpha_channel=True, discard_alpha_channel=True,
): ):
""" """
@@ -463,7 +507,7 @@ class MMEncoder:
elif modality == Modality.VIDEO: elif modality == Modality.VIDEO:
return load_video(data, frame_count_limit) return load_video(data, frame_count_limit)
elif modality == Modality.AUDIO: elif modality == Modality.AUDIO:
return load_audio(data, audio_sample_rate) return load_audio(data, self.model_audio_sr)
except Exception as e: except Exception as e:
raise RuntimeError(f"Error while loading data {data}: {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 ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (feature_lens // 100) * 13
) )
return output_lengths 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: else:
# fallback to original HF audio sample logic for other models # fallback to original HF audio sample logic for other models
logger.warning( logger.warning(
@@ -631,10 +680,18 @@ class MMEncoder:
return slices return slices
def _calculate_hashes_from_features( 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]: ) -> List[str]:
"""CPU Task: Compute hashes based on processed feature patches.""" """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=}") logger.info(f"{mm_feature.shape=} with {modality=}")
for grid in grid_thw: for grid in grid_thw:
num_patches = self.get_num_patches(grid, modality) num_patches = self.get_num_patches(grid, modality)
@@ -647,7 +704,7 @@ class MMEncoder:
async def _encode_missing( async def _encode_missing(
self, self,
mm_feature: torch.Tensor, mm_feature,
mm_inputs: dict, mm_inputs: dict,
indices: List[int], indices: List[int],
modality: Modality = Modality.IMAGE, modality: Modality = Modality.IMAGE,
@@ -658,23 +715,34 @@ class MMEncoder:
""" """
grid_thw = _get_mm_grid_dim(mm_inputs, modality, self.model_type) 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 # 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 = [] sub_feature_list = []
offsets = [0] offsets = [0]
curr = 0 curr = 0
for g in grid_thw: for g in grid_thw:
curr += self.get_num_patches(g, modality) curr += self.get_num_patches(g, modality)
offsets.append(curr) offsets.append(curr)
for idx in indices: for idx in indices:
sub_feature_list.append(mm_feature[offsets[idx] : offsets[idx + 1]]) sub_feature_list.append(mm_feature[offsets[idx] : offsets[idx + 1]])
sub_feature = torch.cat(sub_feature_list, dim=0) sub_feature = torch.cat(sub_feature_list, dim=0)
mm_item = MultimodalDataItem.from_dict( mm_item = MultimodalDataItem.from_dict(
{ {
"modality": modality, "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)) mm_feature = _convert(_get_mm_feature(mm_inputs, modality))
num_items = len(grid_thw) 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. # Step 1: Rank 0 checks global cache and broadcasts hit/miss mask to all ranks.
if self.rank == 0: if self.rank == 0:
if hashes is None: if hashes is None:
@@ -852,7 +929,8 @@ class MMEncoder:
async def _flatten_and_load_audios(self, mm_items): 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) return await self._flatten_and_load_data_by_modality(mm_items, Modality.AUDIO)
@@ -893,6 +971,28 @@ class MMEncoder:
flat.append(item) flat.append(item)
return flat 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): def _normalize_kimi_encoder_images(self, images):
"""Normalize Kimi image inputs for the image processor call.""" """Normalize Kimi image inputs for the image processor call."""
from PIL import Image as PILImage from PIL import Image as PILImage
@@ -1039,18 +1139,21 @@ class MMEncoder:
return processor_input return processor_input
async def _process_audio_items(self, mm_items, model_preprocessor): 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: 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: if not self.audio_processor:
raise ValueError("No audio processor available") raise ValueError("No audio processor available")
audios = await self._flatten_and_load_audios(mm_items)
audio_config = self.vision_config.get("audio", {}) audio_config = self.vision_config.get("audio", {})
processor_input = self.audio_processor.feature_extractor(audios, **audio_config) processor_input = self.audio_processor.feature_extractor(audios, **audio_config)
processor_input["feature_attention_mask"] = processor_input.pop( processor_input["feature_attention_mask"] = processor_input.pop(
"attention_mask" "attention_mask"
) )
# convert to same format as image/video
input_lengths = torch.tensor( input_lengths = torch.tensor(
processor_input["feature_attention_mask"].sum(-1), dtype=torch.long processor_input["feature_attention_mask"].sum(-1), dtype=torch.long
) )
@@ -1225,6 +1328,122 @@ class MMEncoder:
logger.debug(f"Created error EmbeddingData: {mm_data}") logger.debug(f"Created error EmbeddingData: {mm_data}")
return 0, 0, 0, error_msg, error_code 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 # For zmq_to_tokenizer zmq_to_scheduler and mooncake
async def send( async def send(
self, req_id, prefill_host, embedding_port, session_id=None, buffer_address=None self, req_id, prefill_host, embedding_port, session_id=None, buffer_address=None
@@ -1396,9 +1615,239 @@ class EncoderProfiler:
return True, None 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 encoder: Optional[MMEncoder] = None
send_sockets: List[zmq.Socket] = [] 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( async def run_encoder(
@@ -1414,23 +1863,14 @@ async def run_encoder(
encoder.profiler.start(request) encoder.profiler.start(request)
else: else:
encoder.profiler.stop() encoder.profiler.stop()
else: elif isinstance(request, dict) and request.get("type") == "batch_encode":
if encoder.mm_global_cache is not None: await encoder.batch_encode(
await encoder.encode_with_global_cache( request["requests"],
mm_items=request["mm_items"], Modality.from_str(request["modality"]),
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: else:
await encoder.encode( await encoder.encode_request(
mm_items=request["mm_items"], request, Modality.from_str(request["modality"])
modality=Modality.from_str(request["modality"]),
req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
) )
@@ -1489,30 +1929,27 @@ async def handle_encode_request(request: dict):
encoder.background_tasks.add(task) encoder.background_tasks.add(task)
task.add_done_callback(encoder.background_tasks.discard) task.add_done_callback(encoder.background_tasks.discard)
# broadcast request
request.update({"enter_time": time.time()}) request.update({"enter_time": time.time()})
for socket in send_sockets: modality = Modality.from_str(request["modality"])
socket.send_pyobj(request) if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
if encoder.mm_global_cache is not None: try:
nbytes, embedding_len, embedding_dim, error_msg, error_code = ( nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode_with_global_cache( await encoder_scheduler.submit(request)
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),
) )
except asyncio.TimeoutError:
return ORJSONResponse(
status_code=HTTPStatus.GATEWAY_TIMEOUT,
content={
"status": "error",
"message": "encoder batch timed out",
"req_id": req_id,
},
) )
else: else:
for socket in send_sockets:
socket.send_pyobj(request)
nbytes, embedding_len, embedding_dim, error_msg, error_code = ( nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode( await encoder.encode_request(request, modality)
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"],
)
) )
if error_msg: if error_msg:
+3
View File
@@ -712,6 +712,9 @@ class Envs:
SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0) SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0)
SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0) SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0)
SGLANG_ENCODER_DISPATCH_MIN_ITEMS = EnvInt(2) 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 # Elastic EP Backup Port
SGLANG_BACKUP_PORT_BASE = EnvInt(10000) SGLANG_BACKUP_PORT_BASE = EnvInt(10000)
@@ -28,6 +28,7 @@ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import (
Qwen2_5_VLVisionConfig, Qwen2_5_VLVisionConfig,
) )
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
Modality, Modality,
MultimodalDataItem, MultimodalDataItem,
@@ -1815,9 +1816,7 @@ class MiMoV2Processor(BaseMultimodalProcessor):
self.video_end_token_id = self._require_config_value( self.video_end_token_id = self._require_config_value(
processor_config, "video_end_token_id" processor_config, "video_end_token_id"
) )
self.use_image_processor_gpu = ( self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
)
device = server_args.device if self.use_image_processor_gpu else None device = server_args.device if self.use_image_processor_gpu else None
self.mimo_processor = MiMoProcessor( self.mimo_processor = MiMoProcessor(