[EPD][VLM] support video/audio input (#17824)

Co-authored-by: siyu <liusy58@linux.alibaba.com>
This commit is contained in:
Zheng Wengang
2026-03-16 14:18:21 +08:00
committed by GitHub
parent 738cbde902
commit 135af6dc92
12 changed files with 1924 additions and 343 deletions

View File

@@ -26,6 +26,7 @@ from sglang.srt.disaggregation.encode_server import (
handle_scheduler_receive_url_request,
launch_encoder,
)
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid
from sglang.srt.utils.network import get_zmq_socket
@@ -96,6 +97,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
for socket in self.send_sockets:
await socket.send_pyobj(request_dict)
# gRPC encode is image-only; encoder.encode() requires modality
(
nbytes,
embedding_len,
@@ -104,6 +106,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
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,

View File

@@ -1,4 +1,5 @@
import asyncio
import itertools
import logging
import pickle
import random
@@ -6,9 +7,10 @@ import threading
import time
import uuid
from abc import ABC, abstractmethod
from collections import OrderedDict, defaultdict
from enum import IntEnum
from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional
from typing import TYPE_CHECKING, Dict, List, Optional
import aiohttp
import torch
@@ -23,7 +25,7 @@ from sglang.srt.distributed.parallel_state import (
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReqInput
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -125,54 +127,38 @@ class EmbeddingData:
req_id,
num_parts,
part_idx,
image_grid_dim,
grid_dim,
modality,
embedding=None,
embedding_shape=None,
error_msg=None,
error_code=None,
**kwargs,
):
self.req_id = req_id
self.num_parts = num_parts
self.part_idx = part_idx
self.image_grid_dim = image_grid_dim
self.grid_dim = grid_dim
self.modality = modality
self.embedding = embedding
self.send_time = None
self.dtype = embedding.dtype if embedding is not None else None
self.shape = list(embedding.shape) if embedding is not None else None
# aggregated data
self.ready_list = [i == self.part_idx for i in range(self.num_parts)]
self.embedding_list = [
embedding if i == self.part_idx else None for i in range(self.num_parts)
]
self.image_grid_dim_list = [
self.image_grid_dim if i == self.part_idx else None
for i in range(self.num_parts)
]
if embedding_shape is not None:
self.shape = embedding_shape
else:
self.shape = list(embedding.shape) if embedding is not None else None
self.error_msg = error_msg
self.error_code = error_code
# Store additional metadata (e.g., video_timestamps for qwen3_vl)
for key, value in kwargs.items():
setattr(self, key, value)
def add(self, embedding_data):
assert self.req_id == embedding_data.req_id
assert not self.ready_list[embedding_data.part_idx]
self.ready_list[embedding_data.part_idx] = True
self.image_grid_dim_list[embedding_data.part_idx] = (
embedding_data.image_grid_dim
)
self.embedding_list[embedding_data.part_idx] = embedding_data.embedding
def get_grid(self):
"""Get the grid dimension of the embedding, used for image/video/audio."""
return self.grid_dim
def get_embedding(self, is_concat=False):
if is_concat:
return torch.concat(
[embedding.cuda() for embedding in self.embedding_list]
).to("cpu", non_blocking=True)
else:
return self.embedding_list
def get_img_grid(self):
return torch.concatenate(self.image_grid_dim_list)
@property
def ready(self):
return sum(self.ready_list) == self.num_parts
def get_embedding(self):
return self.embedding
def __repr__(self):
return f"EmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx}) error_msg={self.error_msg}"
@@ -182,16 +168,173 @@ class EmbeddingData:
req_id=self.req_id,
num_parts=self.num_parts,
part_idx=self.part_idx,
image_grid_dim=self.image_grid_dim,
grid_dim=self.grid_dim,
modality=self.modality,
embedding=None,
embedding_shape=self.shape,
error_msg=self.error_msg,
error_code=self.error_code,
)
new_data.send_time = self.send_time
new_data.dtype = self.dtype
new_data.shape = self.shape
for key, value in self.__dict__.items():
if key.startswith("_") or key == "embedding":
continue
setattr(new_data, key, value)
return new_data
# Modality -> (list attr name, whether to flatten grid for that list)
_MODALITY_GRID_ATTRS = {
Modality.IMAGE: ("img_grid_thw", False),
Modality.VIDEO: ("video_grid_thw", False),
Modality.AUDIO: ("audio_feature_lens", True),
}
_VIDEO_META_ATTRS = ("video_timestamps", "second_per_grid_ts")
def _cat_grid(dims, flatten_items=False):
"""Concatenate non-None tensors from a list; optionally flatten each before cat."""
valid = (
[g.flatten() for g in dims if g is not None]
if flatten_items
else [g for g in dims if g is not None]
)
return torch.cat(valid, dim=0) if valid else None
class MultiModalEmbeddingData(EmbeddingData):
def __init__(
self,
part_idx,
num_parts,
req_id,
grid_dim,
modality,
embedding,
embedding_shape,
**kwargs,
):
super().__init__(
req_id,
num_parts,
part_idx,
grid_dim,
modality,
embedding,
embedding_shape,
**kwargs,
)
self.img_grid_thw = [None] * num_parts
self.video_grid_thw = [None] * num_parts
self.audio_feature_lens = [None] * num_parts
self.modality_list = [
modality if part_idx == i else None for i in range(num_parts)
]
self.ready_list = [i == part_idx for i in range(num_parts)]
self.embedding_list = [
embedding if i == part_idx else None for i in range(num_parts)
]
self.embedding_shape_list = [
embedding_shape if i == part_idx else None for i in range(num_parts)
]
self.video_timestamps = [None] * num_parts
self.second_per_grid_ts = [None] * num_parts
self._set_part_grid(part_idx, modality, self.get_grid())
if modality == Modality.VIDEO:
self._set_video_meta_for_part(part_idx, kwargs)
def _set_part_grid(self, part_idx, modality, grid):
"""Set the grid for one part according to modality (IMAGE/VIDEO/AUDIO)."""
spec = _MODALITY_GRID_ATTRS.get(modality)
if spec is None:
raise ValueError(f"Invalid modality: {modality}")
attr_name, flatten = spec
value = grid.flatten() if flatten else grid
getattr(self, attr_name)[part_idx] = value
def _set_video_meta_for_part(self, part_idx, source):
"""Copy video_timestamps and second_per_grid_ts from source (dict or object)."""
for attr_name in _VIDEO_META_ATTRS:
val = (
source.get(attr_name)
if isinstance(source, dict)
else getattr(source, attr_name, None)
)
if val is not None:
getattr(self, attr_name)[part_idx] = val
@classmethod
def from_embedding_data(cls, embedding_data: EmbeddingData):
"""Create MultiModalEmbeddingData from an EmbeddingData instance."""
# Only forward known optional attrs (e.g. video metadata) so they land on the instance
extra = {}
for attr in _VIDEO_META_ATTRS:
val = getattr(embedding_data, attr, None)
if val is not None:
extra[attr] = val
mm_data = cls(
part_idx=embedding_data.part_idx,
num_parts=embedding_data.num_parts,
req_id=embedding_data.req_id,
grid_dim=embedding_data.grid_dim,
modality=embedding_data.modality,
embedding=embedding_data.embedding,
embedding_shape=embedding_data.shape,
**extra,
)
mm_data.send_time = embedding_data.send_time
return mm_data
def __repr__(self):
return f"MultiModalEmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx}, modality={self.modality})"
def get_embedding(self, is_concat=False):
if is_concat:
groups = defaultdict(list)
for i, e in enumerate(self.embedding_list):
if e is not None:
groups[self.modality_list[i]].append(e.cuda())
return {
mod: torch.concat(tensors).to("cpu", non_blocking=True)
for mod, tensors in groups.items()
}
return self.embedding_list
@property
def ready(self):
return sum(self.ready_list) == self.num_parts
def get_mm_extra_meta(self):
"""Build kwargs for mm_processor.get_mm_data() from grid and optional video meta."""
kwargs = {
"img_grid_thw": _cat_grid(self.img_grid_thw),
"video_grid_thw": _cat_grid(self.video_grid_thw),
"audio_feature_lens": _cat_grid(
self.audio_feature_lens, flatten_items=True
),
}
for attr in _VIDEO_META_ATTRS:
lst = getattr(self, attr, None)
if not lst:
continue
valid = [a for a in lst if a is not None]
if valid:
kwargs[attr] = list(itertools.chain(*valid))
return kwargs
def add(self, embedding_data: EmbeddingData):
assert self.req_id == embedding_data.req_id
assert not self.ready_list[embedding_data.part_idx]
pid = embedding_data.part_idx
self.ready_list[pid] = True
self.modality_list[pid] = embedding_data.modality
self.embedding_list[pid] = embedding_data.get_embedding()
self.embedding_shape_list[pid] = embedding_data.shape
self._set_part_grid(pid, embedding_data.modality, embedding_data.get_grid())
if embedding_data.modality == Modality.VIDEO:
self._set_video_meta_for_part(pid, embedding_data)
class WaitingImageRequestStatus(IntEnum):
FAIL = -1
PENDING = 0
@@ -199,6 +342,41 @@ class WaitingImageRequestStatus(IntEnum):
TIMEOUT = -2
def create_part_req_id(original_req_id: str, part_idx: int) -> str:
"""Create a unique part request ID by appending part index suffix."""
return f"{original_req_id}_local_part_{part_idx}"
def extract_original_req_id(part_req_id: str) -> str:
"""Extract the original request ID from a part request ID."""
if "_local_part_" in part_req_id:
return part_req_id.rsplit("_local_part_", 1)[0]
return part_req_id
def calculate_modality_num_parts(modalities, num_items_assigned):
"""
Calculate total number of parts and number of parts per modality.
Args:
modalities: List of modalities in order
num_items_assigned: Dictionary mapping modality to list of assignment counts per encoder
Returns:
Tuple of (total_num_parts, modality_num_parts_dict)
- total_num_parts: Total number of parts across all modalities
- modality_num_parts: Dictionary mapping modality to number of parts for that modality
"""
total_num_parts = 0
modality_num_parts = {}
for modality in modalities:
num_items_assigned_modality = num_items_assigned.get(modality)
num_parts = sum(1 for x in num_items_assigned_modality if x != 0)
modality_num_parts[modality] = num_parts
total_num_parts += num_parts
return total_num_parts, modality_num_parts
# For zmq_to_scheduler
class WaitingImageRequest:
def __init__(
@@ -247,21 +425,38 @@ class WaitingImageRequest:
) as session:
tasks = []
logger.info(f"{self.num_items_assigned = } ")
for idx, assigned_num in enumerate(self.num_items_assigned):
if assigned_num == 0:
continue
encoder_url = self.encoder_urls[idx]
target_url = f"{encoder_url}/scheduler_receive_url"
payload = {
"req_id": req_id,
"receive_count": receive_count,
"receive_url": f"{host_name}:{embedding_port}",
}
logger.info(f"Preparing to send to {target_url}")
# Calculate part_idx_offset similar to encode() method
modalities = list(self.num_items_assigned.keys())
_, modality_num_parts = calculate_modality_num_parts(
modalities, self.num_items_assigned
)
task = _send_single_request(session, target_url, payload)
tasks.append(task)
part_idx_offset = 0
for modality in modalities:
assigned_nums = self.num_items_assigned[modality]
num_parts = modality_num_parts[modality]
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(req_id, part_idx)
encoder_url = self.encoder_urls[idx]
target_url = f"{encoder_url}/scheduler_receive_url"
payload = {
"req_id": part_req_id, # use part_req_id to match encode request
"receive_count": receive_count,
"receive_url": f"{host_name}:{embedding_port}",
"modality": modality.name,
}
logger.info(
f"Preparing to send to {target_url} with part_req_id={part_req_id}"
)
task = _send_single_request(session, target_url, payload)
tasks.append(task)
cum_idx += 1
part_idx_offset += num_parts
if not tasks:
logger.info("No tasks to send.")
@@ -305,20 +500,30 @@ class WaitingImageRequest:
return
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
recv_obj.embedding = (
torch.frombuffer(buffer, dtype=recv_obj.dtype)
.reshape(recv_obj.shape)
.clone()
)
recv_obj.embedding_list[recv_obj.part_idx] = recv_obj.embedding
# 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.recv_embedding_data is None:
self.recv_embedding_data = recv_obj
self.recv_embedding_data = MultiModalEmbeddingData.from_embedding_data(
recv_obj
)
else:
self.recv_embedding_data.add(recv_obj)
recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True)
img_grid_thw = self.recv_embedding_data.get_img_grid()
mm_inputs = self.mm_processor.get_mm_data(
self.recv_req.input_text, recv_embedding, img_grid_thw
self.recv_req.input_text,
recv_embedding,
**self.recv_embedding_data.get_mm_extra_meta(),
)
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = mm_inputs["input_ids"]
@@ -330,8 +535,11 @@ class WaitingImageRequestGrpc(WaitingImageRequest):
def send_encode_request(self):
async def send_embedding_port(req_id, receive_count, host_name, embedding_port):
tasks = []
logger.info(f"{self.num_items_assigned = } ")
for idx, assigned_num in enumerate(self.num_items_assigned):
# gRPC image-only: flatten modality dict to flat list
assigned = list(self.num_items_assigned.values())[0]
logger.info(f"num_items_assigned={assigned}")
for idx, assigned_num in enumerate(assigned):
if assigned_num == 0:
continue
encoder_url = self.encoder_urls[idx]
@@ -394,11 +602,22 @@ class MMReceiverBase(ABC):
self.context = zmq.asyncio.Context(20)
self.encoder_transfer_backend = server_args.encoder_transfer_backend
self.encode_urls = server_args.encoder_urls
self.encode_idx = list(range(len(self.encode_urls)))
self.host = get_local_ip_auto(server_args.host)
if self.encoder_transfer_backend == "mooncake":
self.dtype = dtype
self.embeddings_engine = get_mooncake_transfer_engine()
if self.embeddings_engine is None:
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
init_mooncake_transfer_engine,
)
self.embeddings_engine = init_mooncake_transfer_engine(
hostname=self.host,
ib_device=(
server_args.disaggregation_ib_device
or server_args.mooncake_ib_device
),
)
self.embeddings_buffer = dict()
elif self.encoder_transfer_backend == "zmq_to_scheduler":
self.pp_rank = pp_rank
@@ -455,20 +674,17 @@ class MMReceiverBase(ABC):
pass
async def recv_mm_data(
self, img_data, mm_processor, prompt, need_wait_for_image=True
self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True
):
req_id = None
try:
if len(self.encode_urls) == 0 or not need_wait_for_image:
if len(self.encode_urls) == 0 or not need_wait_for_mm_inputs:
return None
req_id = uuid.uuid4().hex
embedding_port, recv_socket = get_zmq_socket_on_host(self.context, zmq.PULL)
if not isinstance(img_data, list):
img_data = [img_data.url]
else:
img_data = [img.url for img in img_data]
mm_data = self._extract_url_data(request_obj)
asyncio.create_task(
self.encode(req_id, img_data, embedding_port, "encode", "send")
self.encode(req_id, mm_data, embedding_port, "encode", "send")
)
return await asyncio.wait_for(
self._recv_mm_data(req_id, recv_socket, mm_processor, prompt),
@@ -501,7 +717,7 @@ class MMReceiverBase(ABC):
recv_embedding = None
recv_embedding_data: EmbeddingData = None
recv_embedding_data: MultiModalEmbeddingData = None
try:
while recv_embedding_data is None or not recv_embedding_data.ready:
@@ -517,6 +733,11 @@ class MMReceiverBase(ABC):
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(
@@ -534,8 +755,9 @@ class MMReceiverBase(ABC):
.clone()
)
if recv_embedding_data is None:
recv_obj.embedding_list[recv_obj.part_idx] = recv_obj.embedding
recv_embedding_data = recv_obj
recv_embedding_data = MultiModalEmbeddingData.from_embedding_data(
recv_obj
)
else:
recv_embedding_data.add(recv_obj)
@@ -545,14 +767,32 @@ class MMReceiverBase(ABC):
"mooncake: embeddings_buffer missing req_id=%s", req_id
)
return None
recv_embedding = self.embeddings_buffer[req_id]
del self.embeddings_buffer[req_id]
self.embeddings_engine.deregister(recv_embedding.data_ptr())
elif self.encoder_transfer_backend == "zmq_to_tokenizer":
recv_embedding = recv_embedding_data.get_embedding(is_concat=True)
raw_buffer = self.embeddings_buffer.pop(req_id)
self.embeddings_engine.deregister(raw_buffer.data_ptr())
byte_offset = 0
for i in range(recv_embedding_data.num_parts):
shape = recv_embedding_data.embedding_shape_list[i]
if shape is None:
continue
part_bytes = (
shape[0]
* shape[1]
* torch.tensor([], dtype=self.dtype).element_size()
)
recv_embedding_data.embedding_list[i] = (
raw_buffer[byte_offset : byte_offset + part_bytes]
.view(self.dtype)
.reshape(shape)
)
byte_offset += part_bytes
img_grid_thw = recv_embedding_data.get_img_grid()
mm_inputs = mm_processor.get_mm_data(prompt, recv_embedding, img_grid_thw)
recv_embedding = recv_embedding_data.get_embedding(is_concat=True)
mm_inputs = mm_processor.get_mm_data(
prompt,
recv_embedding,
**recv_embedding_data.get_mm_extra_meta(),
)
return mm_inputs
finally:
recv_socket.close()
@@ -561,30 +801,24 @@ class MMReceiverBase(ABC):
self._send_encode_request(obj)
def _send_encode_request(self, obj):
if obj.image_data is None:
image_urls = []
elif not isinstance(obj.image_data, list):
image_urls = [obj.image_data.url]
else:
image_urls = [img.url for img in obj.image_data]
mm_data = self._extract_url_data(obj)
if obj.rid is None:
obj.rid = uuid.uuid4().hex
if image_urls and self.encode_urls:
logger.info(f"Processing {len(image_urls)} images for request {obj.rid}")
obj.need_wait_for_image = True
if mm_data and self.encode_urls:
logger.info(f"Processing {len(mm_data)} mm items for request {obj.rid}")
obj.need_wait_for_mm_inputs = True
encode_idx = list(range(len(self.encode_urls)))
random.shuffle(encode_idx)
obj.num_items_assigned = [
(idx + len(image_urls)) // len(self.encode_urls) for idx in encode_idx
]
num_items_assigned = self._assign_items_by_modality(
mm_data, len(self.encode_urls)
)
obj.num_items_assigned = num_items_assigned
encode_thread = threading.Thread(
target=self._run_encode_in_thread,
args=(
obj.rid,
image_urls,
mm_data,
"encode",
obj.num_items_assigned,
num_items_assigned,
None,
),
daemon=True,
@@ -597,7 +831,7 @@ class MMReceiverBase(ABC):
for recv_req in recv_reqs:
if (
isinstance(recv_req, TokenizedGenerateReqInput)
and recv_req.need_wait_for_image is True
and recv_req.need_wait_for_mm_inputs is True
):
waiting_req = waiting_cls(
rid=recv_req.rid,
@@ -666,13 +900,13 @@ class MMReceiverBase(ABC):
return new_recv_reqs, abort_reqs
def _run_encode_in_thread(
self, req_id, img_data, endpoint_encode, num_items_assigned, embedding_port
self, req_id, mm_data, endpoint_encode, num_items_assigned, embedding_port
):
try:
asyncio.run(
self.encode(
req_id=req_id,
img_data=img_data,
mm_data=mm_data,
embedding_port=embedding_port,
endpoint_encode=endpoint_encode,
endpoint_send=None,
@@ -718,11 +952,8 @@ class MMReceiverBase(ABC):
req.tokenizer = self.scheduler.tokenizer
return req
async def allocate_embedding_buffer(self, req_id, embedding_length, embedding_dim):
embeddings = torch.zeros(
(embedding_length, embedding_dim),
dtype=self.dtype,
)
async def allocate_embedding_buffer(self, req_id, total_bytes):
embeddings = torch.empty(total_bytes, dtype=torch.uint8)
self.embeddings_engine.register(
embeddings.data_ptr(),
embeddings.nbytes,
@@ -730,6 +961,76 @@ class MMReceiverBase(ABC):
self.embeddings_buffer[req_id] = embeddings
return embeddings.data_ptr()
def _assign_items_by_modality(
self, mm_data, encoder_num, random_shuffle=True
) -> Dict:
"""
Assign multimodal items across encoders by modality with cross-modality load balancing.
Args:
mm_data: List of multimodal data items, each with a "modality" key
encoder_num: Number of encoders
random_shuffle: Whether to shuffle the encoder indices
Returns:
Dictionary mapping modality to list of assignment counts per encoder
Format: {modality: [count_for_encoder_0, count_for_encoder_1, ...]}
"""
encode_idx = list(range(encoder_num))
if random_shuffle:
random.shuffle(encode_idx)
# Get unique modalities with order preserved
modalities = list(dict.fromkeys(mm_item.get("modality") for mm_item in mm_data))
# Use OrderedDict to explicitly maintain modality order
num_items_assigned = OrderedDict()
current_offset = 0
for modality in modalities:
mm_data_modality = [
mm_item for mm_item in mm_data if mm_item.get("modality") == modality
]
num_items = len(mm_data_modality)
if num_items == 0:
continue
base = num_items // len(encode_idx)
remainder = num_items % len(encode_idx)
# Rotate assignments based on current_offset to balance load across modalities
assignments = [0] * len(encode_idx)
for i in range(len(encode_idx)):
# keep shuffle order when assigning items to encoders
pos_in_shuffled = (current_offset + i) % len(encode_idx)
actual_encoder_idx = encode_idx[pos_in_shuffled]
assignments[actual_encoder_idx] = base + (1 if i < remainder else 0)
num_items_assigned[modality] = assignments
current_offset = (current_offset + remainder) % len(encode_idx)
return num_items_assigned
def _extract_url_data(self, request_obj) -> List[Dict]:
mm_data = []
for attr, modality in [
("image_data", Modality.IMAGE),
("video_data", Modality.VIDEO),
("audio_data", Modality.AUDIO),
]:
mm_items = getattr(request_obj, attr, None)
if mm_items:
if not isinstance(mm_items, list):
mm_items = [mm_items]
for mm_item in mm_items:
mm_data.append(
{
"url": (
mm_item.url
if isinstance(mm_item, ImageData)
else mm_item
),
"modality": modality,
}
)
return mm_data
class MMReceiverHTTP(MMReceiverBase):
def __init__(
@@ -759,42 +1060,65 @@ class MMReceiverHTTP(MMReceiverBase):
async def encode(
self,
req_id,
img_data,
mm_data,
embedding_port,
endpoint_encode,
endpoint_send,
num_items_assigned=None,
):
if len(img_data) == 0:
if len(mm_data) == 0:
return
# Split mm_items
# get unique modalities with order preserved
modalities = [mm_item.get("modality") for mm_item in mm_data]
modalities = list(dict.fromkeys(modalities))
encode_requests = []
if num_items_assigned is None:
random.shuffle(self.encode_idx)
num_items_assigned = [
(idx + len(img_data)) // len(self.encode_urls)
for idx in self.encode_idx
]
num_parts = sum(1 for x in num_items_assigned if x != 0)
cum_num_items = 0
cum_idx = 0
for idx, assigned_num in enumerate(num_items_assigned):
if assigned_num == 0:
continue
encode_requests.append(
{
"encoder_idx": idx,
"mm_items": img_data[cum_num_items : cum_num_items + assigned_num],
"num_parts": num_parts,
"part_idx": cum_idx,
"req_id": req_id,
"prefill_host": self.host,
"embedding_port": embedding_port,
}
num_items_assigned = self._assign_items_by_modality(
mm_data, len(self.encode_urls)
)
cum_idx += 1
cum_num_items += assigned_num
# Calculate total num_parts across all modalities
total_num_parts, modality_num_parts = calculate_modality_num_parts(
modalities, num_items_assigned
)
part_idx_offset = 0
for modality in modalities:
num_items_assigned_modality = num_items_assigned.get(modality)
mm_data_modality = [
mm_item for mm_item in mm_data if mm_item.get("modality") == modality
]
num_parts = modality_num_parts[modality]
cum_num_items = 0
cum_idx = 0
for idx, assigned_num in enumerate(num_items_assigned_modality):
if assigned_num == 0:
continue
part_idx = part_idx_offset + cum_idx
part_req_id = create_part_req_id(req_id, part_idx)
encode_requests.append(
{
"encoder_idx": idx,
"mm_items": [
mm_item.get("url")
for mm_item 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, # use part_req_id to avoid key collision
"modality": modality.name, # convert enum to string for json serialization
"prefill_host": self.host,
"embedding_port": embedding_port,
}
)
cum_idx += 1
cum_num_items += assigned_num
part_idx_offset += num_parts
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(
@@ -832,21 +1156,21 @@ class MMReceiverHTTP(MMReceiverBase):
# mooncake backend: send bootstrap info
embedding_size_list_sort = [None for _ in range(num_parts)]
embedding_length_tot = 0
response_json_list_sort = [None for _ in range(num_parts)]
embedding_size_list_sort = [None for _ in range(total_num_parts)]
response_json_list_sort = [None for _ in range(total_num_parts)]
for response_json in response_json_list_unsort:
idx = response_json["part_idx"]
embedding_size_list_sort[idx] = response_json["embedding_size"]
embedding_length_tot += response_json["embedding_len"]
response_json_list_sort[idx] = response_json
total_embedding_bytes = sum(
s for s in embedding_size_list_sort if s is not None
)
offset = 0
metadata_tasks = []
buffer_address = await self.allocate_embedding_buffer(
req_id,
embedding_length_tot,
response_json_list_sort[0]["embedding_dim"],
total_embedding_bytes,
)
for idx in range(len(tasks)):
response_json = response_json_list_sort[idx]
@@ -903,21 +1227,38 @@ class MMReceiverGrpc(MMReceiverBase):
async def encode(
self,
req_id,
img_data,
mm_data,
embedding_port,
endpoint_encode,
endpoint_send,
num_items_assigned=None,
):
if not img_data:
if not mm_data:
return
# gRPC currently only supports image; flatten new dict formats to simple lists
if mm_data and isinstance(mm_data[0], dict):
non_image = [
item.get("modality")
for item in mm_data
if item.get("modality") != Modality.IMAGE
]
if non_image:
raise NotImplementedError(
f"gRPC encode only supports IMAGE modality, got: {non_image}"
)
img_data = [item.get("url") for item in mm_data]
else:
img_data = mm_data
if isinstance(num_items_assigned, dict):
num_items_assigned = list(num_items_assigned.values())[0]
encode_requests = []
if num_items_assigned is None:
random.shuffle(self.encode_idx)
encode_idx = list(range(len(self.encode_urls)))
random.shuffle(encode_idx)
num_items_assigned = [
(idx + len(img_data)) // len(self.encode_urls)
for idx in self.encode_idx
(idx + len(img_data)) // len(self.encode_urls) for idx in encode_idx
]
num_parts = sum(1 for x in num_items_assigned if x != 0)
cum_num_items = 0
@@ -972,19 +1313,17 @@ class MMReceiverGrpc(MMReceiverBase):
return
embedding_size_by_part = [None for _ in range(num_parts)]
embedding_length_tot = 0
response_json_sorted = [None for _ in range(num_parts)]
for response_json in response_json_unsorted:
idx = response_json["part_idx"]
embedding_size_by_part[idx] = response_json["embedding_size"]
embedding_length_tot += response_json["embedding_len"]
response_json_sorted[idx] = response_json
total_embedding_bytes = sum(s for s in embedding_size_by_part if s is not None)
offset = 0
buffer_address = await self.allocate_embedding_buffer(
req_id,
embedding_length_tot,
response_json_sorted[0]["embedding_dim"],
total_embedding_bytes,
)
grpc_metadata_tasks = []
for response_json in response_json_sorted:

View File

@@ -18,7 +18,7 @@ import zmq
import zmq.asyncio
from fastapi import FastAPI
from fastapi.responses import ORJSONResponse, Response
from transformers import AutoImageProcessor
from transformers import AutoProcessor
from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.load_config import LoadConfig
@@ -37,6 +37,7 @@ from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileRe
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_loader import get_model
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
from sglang.srt.server_args import (
PortArgs,
ServerArgs,
@@ -117,18 +118,46 @@ def _convert(data):
return data
_image_grid_attrs = ["image_grid_thw", "image_grid_hws"]
_mm_grid_attrs = {
Modality.IMAGE: ["image_grid_thw", "image_grid_hws"],
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_image_grid_dim(images_input):
for attr in _image_grid_attrs:
if attr in images_input:
return images_input[attr]
def _get_mm_grid_dim(mm_inputs, modality):
for attr in _mm_grid_attrs[modality]:
if attr in mm_inputs:
return 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"Image grid dim ({_image_grid_attrs}) not found in {images_input}"
f"Feature attrs ({_mm_feature_attrs[modality]}) not found in {mm_inputs}"
)
def _build_mm_aux_data(mm_inputs):
"""
Build auxiliary data for video modality.
"""
aux_data = {
"video_timestamps": mm_inputs.get("video_timestamps", None),
"second_per_grid_ts": mm_inputs.get("second_per_grid_ts", None),
}
return aux_data
class MMEncoder:
def __init__(
self,
@@ -142,17 +171,11 @@ class MMEncoder:
set_global_server_args_for_scheduler(server_args)
self.rank = rank
self.profiler = EncoderProfiler(rank)
self.image_processor = AutoImageProcessor.from_pretrained(
server_args.model_path,
trust_remote_code=server_args.trust_remote_code,
use_fast=True,
)
self._load_mm_processor(server_args)
self.model_config = ModelConfig.from_server_args(
server_args,
)
self.load_config = LoadConfig(
load_format=server_args.load_format,
download_dir=server_args.download_dir,
@@ -161,6 +184,9 @@ class MMEncoder:
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 = server_args.device
self.gpu_id = server_args.base_gpu_id + rank
@@ -172,7 +198,10 @@ class MMEncoder:
torch.get_device_module(self.device).set_device(self.gpu_id)
self.use_image_processor_gpu = use_image_processor_gpu
self.use_image_processor_gpu = (
use_image_processor_gpu and not server_args.disable_fast_image_processor
)
self._build_vision_config(server_args.mm_process_config)
init_distributed_environment(
backend=get_default_distributed_backend(self.device),
@@ -214,10 +243,11 @@ class MMEncoder:
EmbeddingCacheController,
)
hidden_dims = self._infer_embedding_dims()
self.mm_global_cache = EmbeddingCacheController(
rank,
server_args.tp_size,
hidden_dim=self.model_config.hidden_size,
hidden_dims=hidden_dims,
tp_group=get_tp_group().cpu_group,
all_rank_get=False,
)
@@ -233,11 +263,157 @@ class MMEncoder:
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=(
self.server_args.disaggregation_ib_device
or self.server_args.mooncake_ib_device
),
)
self.embedding_to_send = dict()
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 _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
try:
self.image_processor = AutoImageProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
use_fast=not server_args.disable_fast_image_processor,
)
except Exception as e:
logger.warning(f"Failed to load image processor: {e}")
self.image_processor = None
try:
self.video_processor = AutoVideoProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
use_fast=not server_args.disable_fast_image_processor,
)
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(
server_args.tokenizer_path or server_args.model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
use_fast=not server_args.disable_fast_image_processor,
)
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,
@@ -282,16 +458,69 @@ class MMEncoder:
task_info.append((modality, data))
return futures, task_info
async def _flatten_and_load_images(self, mm_items):
def _get_feat_extract_output_lengths(self, feature_lens):
"""
Flatten mm_items structure, load images concurrently, and restore original structure.
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_omni_moe
elif self.model_type == "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
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_images would return
Same structure as load_mm_items would return, support for image/audio
"""
# Handle single image (not a list)
# Handle single mm_item (not a list)
if not isinstance(mm_items, (list, tuple)):
futures, _ = self.submit_data_loading_tasks([mm_items], [Modality.IMAGE])
futures, _ = self.submit_data_loading_tasks([mm_items], [modality])
return await asyncio.wrap_future(futures[0])
# Handle nested list (list of lists)
@@ -299,14 +528,14 @@ class MMEncoder:
# Flatten nested structure
flat_data = []
flat_indices = [] # Track which group each item belongs to
for group_idx, image_group in enumerate(mm_items):
for item in image_group:
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.IMAGE] * len(flat_data)
flat_data, [modality] * len(flat_data)
)
# Wait for all tasks to complete asynchronously
@@ -323,123 +552,137 @@ class MMEncoder:
# Handle simple list
else:
futures, _ = self.submit_data_loading_tasks(
mm_items, [Modality.IMAGE] * len(mm_items)
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]]) -> int:
"""Calculate number of raw patches (before 2x2 merge). Used for pixel_values slicing."""
return int(grid[0] * grid[1] * grid[2])
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())
else:
return int(grid[0] * grid[1] * grid[2])
def get_num_tokens(self, grid: Union[torch.Tensor, List[int]]) -> int:
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."""
merge_size = getattr(self.image_processor, "merge_size", 2)
return self.get_num_patches(grid) // (merge_size**2)
if modality == Modality.AUDIO:
input_length = self.get_num_patches(grid, modality)
return self._get_feat_extract_output_lengths(input_length)
else:
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
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)
count = self.get_num_tokens(grid, modality)
slices.append(mm_embedding[offset : offset + count])
offset += count
return slices
def _calculate_hashes_from_features(
self, pixel_values: torch.Tensor, grid_thw: List
self, mm_feature: torch.Tensor, grid_thw: List, modality: Modality
) -> List[str]:
"""CPU Task: Compute hashes based on processed feature patches (pixel_values)."""
"""CPU Task: Compute hashes based on processed feature patches."""
hashes, offset = [], 0
logger.info(f"{mm_feature.shape=} with {modality=}")
for grid in grid_thw:
num_patches = self.get_num_patches(grid)
feature_slice = pixel_values[offset : offset + num_patches]
tmp_item = MultimodalDataItem(
modality=Modality.IMAGE, feature=feature_slice
)
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
async def _encode_missing(
self, pixel_values: torch.Tensor, images_input: dict, indices: List[int]
self,
mm_feature: torch.Tensor,
mm_inputs: dict,
indices: List[int],
modality: Modality = Modality.IMAGE,
get_feature_fn=None,
) -> List[torch.Tensor]:
"""
GPU Task: Run ViT inference ONLY on the subset of images missing from the cache.
GPU Task: Run ViT inference ONLY on the subset of mm items missing from the cache.
"""
grid_thw = images_input["image_grid_thw"]
grid_thw = _get_mm_grid_dim(mm_inputs, modality)
# 1. Slice pixel_values to get only the patches for missing images
sub_pixel_list = []
# 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)
curr += self.get_num_patches(g, modality)
offsets.append(curr)
for idx in indices:
sub_pixel_list.append(pixel_values[offsets[idx] : offsets[idx + 1]])
sub_feature_list.append(mm_feature[offsets[idx] : offsets[idx + 1]])
sub_feature = torch.cat(sub_pixel_list, dim=0)
sub_feature = torch.cat(sub_feature_list, dim=0)
mm_item = MultimodalDataItem.from_dict(
{
"modality": Modality.IMAGE,
"modality": modality,
"feature": _convert(sub_feature),
}
)
for k, v in images_input.items():
if k == "pixel_values":
for k, v in mm_inputs.items():
if k in _mm_feature_attrs.get(modality, []):
continue
val = _convert(v)
if k in _image_grid_attrs:
if k in _mm_grid_attrs.get(modality, []):
mm_item.set(k, val[indices])
else:
mm_item.set(k, val)
with torch.inference_mode():
new_embeddings = self.model.get_image_feature([mm_item]).cpu()
new_embeddings = get_feature_fn([mm_item]).cpu()
if new_embeddings.ndim != 2:
new_embeddings = new_embeddings.reshape(-1, new_embeddings.shape[-1])
sub_grids = [grid_thw[i] for i in indices]
return self.slice_embedding(new_embeddings, sub_grids)
return self.slice_embedding(new_embeddings, sub_grids, modality)
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:
images = await self._flatten_and_load_images(mm_items)
kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
images_input = self.image_processor(images=images, **kwargs)
pixel_values = images_input["pixel_values"]
grid_thw = images_input["image_grid_thw"]
num_images = len(grid_thw)
mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, modality)
grid_thw = _get_mm_grid_dim(mm_inputs, modality)
mm_feature = _convert(_get_mm_feature(mm_inputs, modality))
num_items = len(grid_thw)
# Step 1: Rank 0 checks global cache and broadcasts hit/miss mask to all ranks.
if self.rank == 0:
if hashes is None:
image_hashes = self._calculate_hashes_from_features(
pixel_values, grid_thw
mm_hashes = self._calculate_hashes_from_features(
mm_feature, grid_thw, modality
)
else:
image_hashes = hashes
exist_mask = await self.mm_global_cache.batch_is_exist(image_hashes)
mm_hashes = hashes
exist_mask = await self.mm_global_cache.batch_is_exist(mm_hashes)
mask_tensor = torch.tensor(
[1 if e else 0 for e in exist_mask], dtype=torch.int32
)
else:
image_hashes = None
mask_tensor = torch.zeros(num_images, dtype=torch.int32)
mm_hashes = None
mask_tensor = torch.zeros(num_items, dtype=torch.int32)
if self.server_args.tp_size > 1:
torch.distributed.broadcast(
@@ -456,7 +699,7 @@ class MMEncoder:
new_slices = []
if missing_indices:
new_slices = await self._encode_missing(
pixel_values, images_input, missing_indices
mm_feature, mm_inputs, missing_indices, modality, get_feature_fn
)
# Step 3: Rank 0 prefetches cache-hit embeddings from global cache.
@@ -464,9 +707,11 @@ class MMEncoder:
if self.rank == 0:
if hit_indices:
hit_hashes = [image_hashes[i] for i in hit_indices]
hit_tokens = [self.get_num_tokens(grid_thw[i]) for i in hit_indices]
self.mm_global_cache.prefetch(req_id, hit_hashes, hit_tokens)
hit_hashes = [mm_hashes[i] for i in hit_indices]
hit_tokens = [
self.get_num_tokens(grid_thw[i], modality) for i in hit_indices
]
self.mm_global_cache.prefetch(req_id, hit_hashes, hit_tokens, modality)
try:
@@ -478,7 +723,7 @@ class MMEncoder:
except (asyncio.TimeoutError, Exception) as e:
logger.error(
f"Prefetch failed for req {req_id}: {e}. "
f"Falling back to ViT for {len(hit_indices)} hit images."
f"Falling back to ViT for {len(hit_indices)} hit items."
)
prefetch_status[0] = 0
@@ -490,21 +735,21 @@ class MMEncoder:
group=self.mm_global_cache.prefetch_tp_group,
)
# Step 5: If prefetch failed, all ranks fallback to ViT for the hit images.
# Step 5: If prefetch failed, all ranks fallback to ViT for the hit mm items.
if prefetch_status.item() == 0 and hit_indices:
logger.info(
f"Req {req_id}: Prefetch failed, all ranks running ViT fallback "
f"for {len(hit_indices)} images."
f"for {len(hit_indices)} mm items."
)
fallback_slices = await self._encode_missing(
pixel_values, images_input, hit_indices
mm_feature, mm_inputs, hit_indices, modality, get_feature_fn
)
else:
fallback_slices = None
# Step 6: Rank 0 assembles final embedding and prepares for sending.
if self.rank == 0:
final_slices = [None] * num_images
final_slices = [None] * num_items
for i, idx in enumerate(missing_indices):
final_slices[idx] = new_slices[i]
@@ -512,7 +757,7 @@ class MMEncoder:
# Fill in cache-hit embeddings (from prefetch or fallback)
if prefetch_status.item() == 1 and hit_indices:
cached_slices = self.mm_global_cache.get_embeddings(
[image_hashes[i] for i in hit_indices]
[mm_hashes[i] for i in hit_indices]
)
for i, idx in enumerate(hit_indices):
final_slices[idx] = cached_slices[i]
@@ -524,10 +769,10 @@ class MMEncoder:
# Background insert: store newly computed embeddings into global cache.
# Includes both original misses and fallback-recomputed hits.
all_new_hashes = [image_hashes[i] for i in missing_indices]
all_new_hashes = [mm_hashes[i] for i in missing_indices]
all_new_slices = list(new_slices)
if fallback_slices is not None:
all_new_hashes += [image_hashes[i] for i in hit_indices]
all_new_hashes += [mm_hashes[i] for i in hit_indices]
all_new_slices += list(fallback_slices)
if all_new_hashes:
@@ -543,8 +788,15 @@ class MMEncoder:
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
aux_data = _build_mm_aux_data(mm_inputs)
self.embedding_to_send[req_id] = EmbeddingData(
req_id, num_parts, part_idx, grid_thw, mm_embedding
req_id,
num_parts,
part_idx,
grid_thw,
modality,
mm_embedding,
**aux_data,
)
return (
mm_embedding.nbytes,
@@ -556,31 +808,152 @@ class MMEncoder:
else:
return (0, 0, 0, None, None)
async def _encode(self, mm_items) -> torch.Tensor:
try:
images = await self._flatten_and_load_images(mm_items)
except Exception as e:
raise BadRequestError(f"Failed to load images from input: {str(e)}")
async def _flatten_and_load_audios(self, mm_items):
"""
Flatten mm_items structure, load audios concurrently, and restore original structure.
"""
return await self._flatten_and_load_data_by_modality(mm_items, Modality.AUDIO)
try:
kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
images_input = self.image_processor(images=images, **kwargs)
feature = images_input["pixel_values"]
mm_item = MultimodalDataItem.from_dict(
{
"modality": Modality.IMAGE,
"feature": _convert(feature),
}
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)
)
for k, v in images_input.items():
if k == "pixel_values":
continue
mm_item.set(k, _convert(v))
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
async def _process_mm_items(self, mm_items, modality):
if modality == Modality.IMAGE and self.image_processor:
images = await self._flatten_and_load_images(mm_items)
image_config = self.vision_config.get("image", {})
processor_input = self.image_processor(images=images, **image_config)
feature = processor_input["pixel_values"]
if hasattr(self.model, "thinker"): # for omni models
get_feature_method = self.model.thinker.get_image_feature
else:
get_feature_method = self.model.get_image_feature
elif modality == Modality.VIDEO and self.video_processor:
videos, video_processor_kwargs = await self._flatten_and_load_videos(
mm_items
)
processor_input = self.video_processor(
videos=videos, **video_processor_kwargs
)
# Get additional video metadata
if (
self.model_type in ["qwen3_vl", "qwen3_vl_moe"]
and video_processor_kwargs.get("video_metadata", None) is not None
):
# For qwen3-vl models, we need to store the video timestamps
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
):
# For omni/qwen2_5_vl models, calculate second_per_grid_ts for rotary embedding
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
# get sampled fps, default: 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
feature = processor_input["pixel_values_videos"]
if hasattr(self.model, "thinker"): # for omni models
get_feature_method = self.model.thinker.get_video_feature
else:
get_feature_method = self.model.get_video_feature
elif modality == Modality.AUDIO and self.audio_processor:
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
)
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
feature = processor_input["input_features"]
if hasattr(self.model, "thinker"): # for omni models
get_feature_method = self.model.thinker.get_audio_feature
else:
get_feature_method = self.model.get_audio_feature
else:
raise ValueError(
f"Currently only support image, video and audio modalities, {modality} modality has no processor available."
)
return processor_input, get_feature_method
async def _encode(self, mm_items, modality: Modality) -> torch.Tensor:
try:
mm_inputs, get_feature_fn = await self._process_mm_items(mm_items, 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)}")
try:
# support mm_cache
mm_embedding = None
mm_hash = None
mm_item = MultimodalDataItem.from_dict(
{
"modality": modality,
"feature": _convert(_get_mm_feature(mm_inputs, modality)),
}
)
for k, v in mm_inputs.items():
if k in _mm_feature_attrs[modality]:
continue
mm_item.set(k, _convert(v))
if self.server_args.enable_prefix_mm_cache:
mm_item.set_pad_value()
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
@@ -591,7 +964,7 @@ class MMEncoder:
if mm_embedding is None:
with torch.inference_mode():
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
mm_embedding: torch.Tensor = get_feature_fn([mm_item])
mm_embedding = mm_embedding.cpu()
if len(mm_embedding.shape) != 2:
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
@@ -602,7 +975,8 @@ class MMEncoder:
if self.profiler is not None:
self.profiler.step()
return _get_image_grid_dim(images_input), mm_embedding
aux_data = _build_mm_aux_data(mm_inputs)
return _get_mm_grid_dim(mm_inputs, modality), mm_embedding, aux_data
except BadRequestError as e:
raise BadRequestError(f"Bad request error: {str(e)}")
except Exception as e:
@@ -626,7 +1000,6 @@ class MMEncoder:
self.engine.deregister(embedding.data_ptr())
mm_data.embedding = None
mm_data.embedding_list[mm_data.part_idx] = None
# Send ack/data
endpoint = (
@@ -665,16 +1038,19 @@ class MMEncoder:
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
async def encode_with_hash(self, mm_items, req_id, num_parts, part_idx, hashes):
images = await self._flatten_and_load_images(mm_items)
async def encode(self, mm_items, req_id, num_parts, part_idx):
async def encode(self, mm_items, modality: Modality, req_id, num_parts, part_idx):
try:
image_grid_dim, mm_embedding = await self._encode(mm_items)
grid_dim, mm_embedding, aux_data = await self._encode(mm_items, modality)
if self.rank == 0:
mm_data = EmbeddingData(
req_id, num_parts, part_idx, image_grid_dim, mm_embedding
req_id,
num_parts,
part_idx,
grid_dim,
modality,
mm_embedding,
**aux_data,
)
self.embedding_to_send[req_id] = mm_data
return (
@@ -694,6 +1070,7 @@ class MMEncoder:
num_parts,
part_idx,
None,
modality,
error_msg=error_msg,
error_code=error_code,
)
@@ -894,6 +1271,7 @@ async def run_encoder(
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"],
@@ -902,6 +1280,7 @@ async def run_encoder(
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"],
@@ -966,6 +1345,7 @@ async def handle_encode_request(request: dict):
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"],
@@ -976,6 +1356,7 @@ async def handle_encode_request(request: dict):
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"],

View File

@@ -28,7 +28,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
import torch
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import BaseFinishReason
from sglang.srt.managers.schedule_batch import BaseFinishReason, Modality
from sglang.srt.multimodal.mm_utils import has_valid_data
from sglang.srt.observability.req_time_stats import (
APIServerReqTimeStats,
@@ -226,8 +226,8 @@ class GenerateReqInput(BaseReq):
received_time: Optional[float] = None
# For EPD-disaggregated inference
need_wait_for_image: Optional[bool] = None
num_items_assigned: Optional[List] = None
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
# Multimodal tiling controls (extensions)
max_dynamic_patch: Optional[int] = None
@@ -731,8 +731,8 @@ class TokenizedGenerateReqInput(BaseReq):
# Whether to return entropy
return_entropy: bool = False
need_wait_for_image: bool = False
num_items_assigned: Optional[List] = None
need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
# For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None

View File

@@ -427,6 +427,7 @@ def get_embedding_chunk(
def _get_precomputed_embedding(
items: List[MultimodalDataItem],
items_size: List[int],
prefix_length: List[int],
extend_length: List[int],
items_offset_list: List[List[Tuple[int, int]]],
@@ -437,38 +438,32 @@ def _get_precomputed_embedding(
If none have precomputed_embeddings, return None.
"""
precomputed_embeddings = []
for idx, item in enumerate(items):
if item.precomputed_embeddings is None:
precomputed_embeddings.append(None)
max_iterations = min(len(items_size) - 1, len(prefix_length))
for i in range(max_iterations):
if items_size[i] == items_size[i + 1]:
continue
seq_start_idx = prefix_length[idx]
seq_end_idx = seq_start_idx + extend_length[idx] - 1
prefix_embedding_length = []
extend_embedding_length = []
for mm_start_idx, mm_end_idx in items_offset_list[idx]:
if mm_start_idx > seq_end_idx:
break
if seq_start_idx > mm_start_idx:
prefix_embedding_length.append(
min(seq_start_idx - mm_start_idx, mm_end_idx - mm_start_idx + 1)
)
if mm_end_idx >= seq_start_idx:
extend_embedding_length.append(
min(
mm_end_idx - seq_start_idx + 1,
seq_end_idx - mm_start_idx + 1,
mm_end_idx - mm_start_idx + 1,
seq_end_idx - seq_start_idx + 1,
)
)
prefix_embedding_length = int(np.sum(prefix_embedding_length))
extend_embedding_length = int(np.sum(extend_embedding_length))
precomputed_embeddings.append(
item.precomputed_embeddings[
prefix_embedding_length : prefix_embedding_length
+ extend_embedding_length
]
)
items_per_req = items[items_size[i] : items_size[i + 1]]
extend_len = extend_length[i] if i < len(extend_length) else 0
items_offset = items_offset_list[i]
if any(item.precomputed_embeddings is None for item in items_per_req):
chunk = None
else:
req_embeddings = torch.concat(
[item.precomputed_embeddings for item in items_per_req]
)
chunk, _, _ = get_embedding_chunk(
embedding=req_embeddings,
extend_prefix_len=prefix_length[i],
extend_seq_len=extend_len,
items_offset=items_offset,
)
if chunk is None and len(items_per_req) > 1:
return None
precomputed_embeddings.append(chunk)
if any(feature is not None for feature in precomputed_embeddings):
if not all(feature is not None for feature in precomputed_embeddings):
@@ -886,7 +881,7 @@ def get_embedding_and_mask(
"""
# 1. Get embedding
embedding = _get_precomputed_embedding(
embedding_items, prefix_length, extend_length, items_offset_list
embedding_items, items_size, prefix_length, extend_length, items_offset_list
)
if embedding is None:
embedding, input_ids = _get_chunked_prefill_embedding(

View File

@@ -724,10 +724,10 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
):
if self.server_args.language_only:
mm_inputs = await self.mm_receiver.recv_mm_data(
img_data=obj.image_data,
request_obj=obj,
mm_processor=self.mm_processor,
prompt=(input_text or input_ids),
need_wait_for_image=obj.need_wait_for_image,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
)
if mm_inputs is None:
mm_inputs: Dict = await self.mm_data_processor.process(
@@ -740,7 +740,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
elif (
self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
and not obj.need_wait_for_image
and not obj.need_wait_for_mm_inputs
):
# In language_only mode with zmq_to_scheduler, if we didn't dispatch
# to encoder (e.g., only one image), process locally like non-language_only mode
@@ -982,7 +982,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
priority=obj.priority,
extra_key=obj.extra_key,
routing_key=obj.routing_key,
need_wait_for_image=obj.need_wait_for_image,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned,
)
elif isinstance(obj, EmbeddingReqInput):
@@ -2362,14 +2362,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if self.server_args.enable_adaptive_dispatch_to_encoder:
should_dispatch = self._should_dispatch_to_encoder(obj)
# Set need_wait_for_image flag based on whether we dispatch to encoder
# Set need_wait_for_mm_inputs flag based on whether we dispatch to encoder
# This flag will be used in _tokenize_one_request to determine processing path
if should_dispatch:
obj.need_wait_for_image = True
obj.need_wait_for_mm_inputs = True
if self.server_args.encoder_transfer_backend == "zmq_to_scheduler":
self.mm_receiver.send_encode_request(obj)
else:
obj.need_wait_for_image = False
obj.need_wait_for_mm_inputs = False
def convert_to_span_attrs(
self,

View File

@@ -95,14 +95,14 @@ class EmbeddingCacheController:
tp_rank,
tp_size,
max_pool_size_gb=4.0,
hidden_dim=1024,
hidden_dims: dict = None,
tp_group=None,
all_rank_get=False,
):
self.tp_world_size = tp_size
self.tp_group = tp_group
self.all_rank_get = all_rank_get
self.hidden_dim = hidden_dim
self.hidden_dims = hidden_dims or {}
self.element_size = torch.float32.itemsize
# 1. Mooncake Backend & Pinned Buffer
@@ -115,7 +115,8 @@ class EmbeddingCacheController:
# 2. Variable Size Memory Management
self.allocator = ContiguousMemoryAllocator(self.total_pool_size_bytes)
self.hash_to_metadata = {} # {image_hash: (offset, num_tokens, size_bytes)}
# {hash: (offset, num_tokens, embedding_dim, size_bytes)}
self.hash_to_metadata = {}
# 3. Task Tracking
self.ongoing_prefetch = {} # {req_id: EmbeddingPrefetchOperation}
@@ -142,9 +143,19 @@ class EmbeddingCacheController:
self.prefetch_tp_group = None
def prefetch(
self, req_id: str, image_hashes: List[str], expected_tokens: List[int]
self,
req_id: str,
image_hashes: List[str],
expected_tokens: List[int],
modality=None,
):
"""Issues ONE batch GET for all missing images in the request."""
dim = self.hidden_dims.get(modality) if modality is not None else None
if not dim:
logger.warning(
f"Req {req_id}: Unknown dim for modality={modality}, skipping prefetch (will fallback to ViT)."
)
return
keys, ptrs, sizes = [], [], []
with self.lock:
@@ -155,12 +166,12 @@ class EmbeddingCacheController:
)
continue
size_bytes = num_tokens * self.hidden_dim * self.element_size
size_bytes = num_tokens * dim * self.element_size
offset = self.allocator.allocate(size_bytes)
if offset is None:
continue
self.hash_to_metadata[h] = (offset, num_tokens, size_bytes)
self.hash_to_metadata[h] = (offset, num_tokens, dim, size_bytes)
keys.append(h)
ptrs.append(self.cpu_pool.data_ptr() + offset)
sizes.append(size_bytes)
@@ -187,20 +198,20 @@ class EmbeddingCacheController:
if h in self.hash_to_metadata:
continue
num_tokens = tensor.shape[0]
size_bytes = num_tokens * self.hidden_dim * self.element_size
num_tokens, dim = tensor.shape[0], tensor.shape[1]
size_bytes = num_tokens * dim * self.element_size
offset = self.allocator.allocate(size_bytes)
if offset is None:
continue
# Copy to pinned pool for RDMA
self.hash_to_metadata[h] = (offset, num_tokens, size_bytes)
target_view = (
self.cpu_pool[offset : offset + size_bytes]
.view(torch.float32)
.view(num_tokens, self.hidden_dim)
.view(num_tokens, dim)
)
target_view.copy_(tensor.cpu())
self.hash_to_metadata[h] = (offset, num_tokens, dim, size_bytes)
keys.append(h)
ptrs.append(self.cpu_pool.data_ptr() + offset)
@@ -277,11 +288,11 @@ class EmbeddingCacheController:
with self.lock:
tensors = []
for h in image_hashes:
offset, num_tokens, size_bytes = self.hash_to_metadata[h]
offset, num_tokens, dim, size_bytes = self.hash_to_metadata[h]
tensors.append(
self.cpu_pool[offset : offset + size_bytes]
.view(torch.float32)
.view(num_tokens, self.hidden_dim)
.view(num_tokens, dim)
)
return tensors

View File

@@ -256,15 +256,20 @@ class BaseMultimodalProcessor(ABC):
def spatial_merge_size(self):
return self.hf_config.vision_config.spatial_merge_size
def build_input_ids(self, prompt, img_grid_thw):
def build_input_ids(
self, prompt, img_grid_thw=None, video_grid_thw=None, audio_seq_lens=None
):
"""
Use prompt and img_grid_thw to build input_ids
Use prompt, img_grid_thw, video_grid_thw, and audio_seq_lens to build input_ids.
Supports image, video, and audio tokens.
"""
if not isinstance(prompt, list):
prompt = self._tokenizer.encode(prompt)
img_token_id = self.IM_TOKEN_ID
spatial_merge_size = self.spatial_merge_size
img_token_id = getattr(self, "IM_TOKEN_ID", None)
video_token_id = getattr(self, "VIDEO_TOKEN_ID", None)
audio_token_id = getattr(self, "audio_token_id", None)
spatial_merge_size = getattr(self, "spatial_merge_size", 1)
input_ids = []
offsets = []
@@ -273,34 +278,81 @@ class BaseMultimodalProcessor(ABC):
# Use img_token_id instead of im_start_id, because a dummy im_start_id
# may be generated by the tokenizer.
img_start_indices = list(
filter(lambda i: prompt[i + 1] == img_token_id, range(len(prompt) - 1))
)
vision_start_indices = []
for i in range(len(prompt) - 1):
if img_token_id is not None and prompt[i + 1] == img_token_id:
vision_start_indices.append((i, Modality.IMAGE))
elif video_token_id is not None and prompt[i + 1] == video_token_id:
vision_start_indices.append((i, Modality.VIDEO))
elif audio_token_id is not None and prompt[i + 1] == audio_token_id:
vision_start_indices.append((i, Modality.AUDIO))
# get modality list with order preserved
modality_list = [modality for _, modality in vision_start_indices]
for cur_img_idx, img_start_idx in enumerate(img_start_indices):
assert cur_idx <= img_start_idx
# include img_start_id
input_ids.extend(prompt[cur_idx : img_start_idx + 1])
img_offset_start = len(input_ids)
img_token_num = img_grid_thw[cur_img_idx].prod() // (spatial_merge_size**2)
input_ids.extend([img_token_id] * img_token_num)
# jump to img_end_id
cur_idx = img_start_idx + 2
offsets.append((img_offset_start, len(input_ids) - 1))
img_idx = 0
video_idx = 0
audio_idx = 0
for mm_start_idx, modality in vision_start_indices:
if modality == Modality.IMAGE:
mm_token_num = img_grid_thw[img_idx].prod() // (spatial_merge_size**2)
mm_token_id = img_token_id
img_idx += 1
elif modality == Modality.VIDEO:
mm_token_num = video_grid_thw[video_idx].prod() // (
spatial_merge_size**2
)
mm_token_id = video_token_id
video_idx += 1
elif modality == Modality.AUDIO:
mm_token_num = int(audio_seq_lens[audio_idx].item())
mm_token_id = audio_token_id
audio_idx += 1
else:
raise ValueError(f"Invalid modality: {modality}")
assert cur_idx <= mm_start_idx
input_ids.extend(prompt[cur_idx : mm_start_idx + 1])
mm_offset_start = len(input_ids)
input_ids.extend([mm_token_id] * mm_token_num)
cur_idx = (
mm_start_idx + 2
) # jump to img_end_id, video_end_id, or audio_end_id
offsets.append((mm_offset_start, len(input_ids) - 1))
else:
input_ids.extend(prompt[cur_idx:])
return input_ids, offsets
return input_ids, offsets, modality_list
def get_mm_data(self, prompt, embeddings, img_grid_thw):
input_ids, offsets = self.build_input_ids(prompt, img_grid_thw)
mm_items = [
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=offsets,
precomputed_embeddings=embeddings,
def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None)
video_grid_thw = kwargs.get("video_grid_thw", None)
audio_feature_lens = kwargs.get("audio_feature_lens", None)
input_ids, offsets, modality_list = self.build_input_ids(
prompt,
img_grid_thw=img_grid_thw,
video_grid_thw=video_grid_thw,
audio_seq_lens=audio_feature_lens,
)
assert all(isinstance(modality, Modality) for modality in modality_list)
mm_items = []
consumed_per_modality = {}
for modality, offset in zip(modality_list, offsets):
num_tokens = offset[1] - offset[0] + 1
embedding_start = consumed_per_modality.get(modality, 0)
embedding_slice = embeddings[modality][
embedding_start : embedding_start + num_tokens
]
consumed_per_modality[modality] = embedding_start + num_tokens
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
precomputed_embeddings=embedding_slice,
)
)
]
return {
"input_ids": input_ids,
@@ -308,6 +360,7 @@ class BaseMultimodalProcessor(ABC):
"im_start_id": self.IM_START_TOKEN_ID,
"im_end_id": self.IM_END_TOKEN_ID,
"im_token_id": self.IM_TOKEN_ID,
"video_token_id": getattr(self, "VIDEO_TOKEN_ID", None),
}
def process_mm_data(

View File

@@ -1,6 +1,6 @@
import re
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.models.qwen2_audio import Qwen2AudioForConditionalGeneration
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
@@ -31,6 +31,52 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor):
self.ATTR_NAME_TO_MODALITY.update({"feature_attention_mask": Modality.AUDIO})
def get_mm_data(self, prompt, embeddings, **kwargs):
audio_feature_lens = kwargs.get("audio_feature_lens", None)
# Convert audio_feature_lens to token counts for build_input_ids
output_lengths = None
input_lengths = None
if audio_feature_lens is not None:
if audio_feature_lens.dim() > 1:
audio_feature_lens = audio_feature_lens.flatten()
input_lengths = (audio_feature_lens - 1) // 2 + 1
output_lengths = (input_lengths - 2) // 2 + 1
input_ids, offsets, modality_list = self.build_input_ids(
prompt,
audio_seq_lens=output_lengths,
)
mm_items = []
consumed_per_modality = {}
for modality, offset in zip(modality_list, offsets):
num_tokens = offset[1] - offset[0] + 1
embedding_start = consumed_per_modality.get(modality, 0)
embedding_slice = embeddings[modality][
embedding_start : embedding_start + num_tokens
]
consumed_per_modality[modality] = embedding_start + num_tokens
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
precomputed_embeddings=embedding_slice,
)
)
if mm_items:
mm_items[0].audio_feature_lens = output_lengths
return {
"mm_items": mm_items,
"input_ids": input_ids,
"audio_start_id": self.audio_start_id,
"audio_token_id": self.audio_token_id,
"audio_end_id": self.audio_end_id,
}
async def process_mm_data_async(
self,
audio_data,

View File

@@ -254,6 +254,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
self.IM_START_TOKEN_ID = hf_config.vision_start_token_id
self.IM_END_TOKEN_ID = hf_config.vision_end_token_id
self.IM_TOKEN_ID = hf_config.image_token_id
self.VIDEO_TOKEN_ID = hf_config.video_token_id
self.vision_start_token_id = hf_config.vision_start_token_id
self.vision_end_token_id = getattr(hf_config, "vision_end_token_id", None)
@@ -271,12 +272,138 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
image_token_regex=re.compile(
r"<\|vision_start\|>(?:<\|image_pad\|>)+<\|vision_end\|>"
),
video_token_id=hf_config.video_token_id,
video_token_id=self.VIDEO_TOKEN_ID,
audio_token_id=self.audio_token_id,
).build(_processor)
def get_mm_data(self, prompt, embeddings, img_grid_thw):
input_ids, offsets = self.build_input_ids(prompt, img_grid_thw)
def build_input_ids_with_timestamps(
self, prompt, embeddings, img_grid_thw, video_grid_thw, video_timestamps
):
"""
Build input_ids with timestamps for qwen3_vl models.
"""
if not isinstance(prompt, list):
prompt = self._processor.tokenizer.encode(prompt)
img_token_id = getattr(self, "IM_TOKEN_ID", None)
video_token_id = getattr(self, "VIDEO_TOKEN_ID", None)
audio_token_id = getattr(self, "audio_token_id", None)
spatial_merge_size = getattr(self, "spatial_merge_size", 1)
vision_start_token_id = getattr(self, "vision_start_token_id", None)
vision_end_token_id = getattr(self, "vision_end_token_id", None)
input_ids = []
offsets = []
modality_list = []
cur_idx = 0
vision_start_indices = []
for i in range(len(prompt) - 1):
if img_token_id is not None and prompt[i + 1] == img_token_id:
vision_start_indices.append((i, Modality.IMAGE))
elif video_token_id is not None and prompt[i + 1] == video_token_id:
vision_start_indices.append((i, Modality.VIDEO))
img_idx = 0
video_idx = 0
model_type = getattr(self, "model_type", None)
for mm_start_idx, modality in vision_start_indices:
modality_list.append(modality)
video_tokens = None
if modality == Modality.IMAGE:
mm_token_num = img_grid_thw[img_idx].prod() // (spatial_merge_size**2)
mm_token_id = img_token_id
img_idx += 1
elif modality == Modality.VIDEO:
curr_timestamps = video_timestamps[video_idx]
num_frames = video_grid_thw[video_idx][0]
frame_seqlen = video_grid_thw[video_idx][1:].prod().item() // (
spatial_merge_size**2
)
video_tokens = []
_current_offset = len(input_ids) + mm_start_idx + 1 - cur_idx
# take single frame as one mm_item
for frame_idx in range(num_frames):
if frame_idx > 0:
modality_list.append(Modality.VIDEO)
curr_time = curr_timestamps[frame_idx]
timestamp_text = f"<{curr_time:.1f} seconds>"
timestamp_tokens = self._processor.tokenizer.encode(
timestamp_text, add_special_tokens=False
)
video_tokens.extend(timestamp_tokens)
_current_offset += len(timestamp_tokens)
if vision_start_token_id is not None:
video_tokens.append(vision_start_token_id)
_current_offset += 1
video_tokens.extend([video_token_id] * frame_seqlen)
if vision_end_token_id is not None:
video_tokens.append(vision_end_token_id)
offsets.append(
(_current_offset, _current_offset + frame_seqlen - 1)
)
_current_offset += (
frame_seqlen + 1
if vision_end_token_id is not None
else frame_seqlen
) # for vision_end_token_id
mm_token_num = len(video_tokens)
mm_token_id = None
video_idx += 1
else:
logger.warning(
f"{modality} modality is not supported for qwen3_vl models with timestamps."
)
continue
assert cur_idx <= mm_start_idx
input_ids.extend(prompt[cur_idx : mm_start_idx + 1])
if modality == Modality.VIDEO:
input_ids.extend(video_tokens)
else:
mm_offset_start = len(input_ids)
input_ids.extend([mm_token_id] * mm_token_num)
offsets.append((mm_offset_start, len(input_ids) - 1))
cur_idx = mm_start_idx + 2 # jump to vision_end_id
else:
input_ids.extend(prompt[cur_idx:])
return input_ids, offsets, modality_list
def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None)
video_grid_thw = kwargs.get("video_grid_thw", None)
audio_feature_lens = kwargs.get("audio_feature_lens", None)
video_timestamps = kwargs.get("video_timestamps", None)
second_per_grid_ts = kwargs.get("second_per_grid_ts", None)
audio_seq_lens = None
if audio_feature_lens is not None:
if self.model_type == "qwen3_omni_moe":
# apply _get_feat_extract_lengths to get seq_lens
input_lengths_leave = audio_feature_lens % 100
feat_lengths = (input_lengths_leave - 1) // 2 + 1
audio_seq_lens = (
((feat_lengths - 1) // 2 + 1 - 1) // 2
+ 1
+ (audio_feature_lens // 100) * 13
)
elif self.model_type == "qwen2_5_omni":
audio_seq_lens = (audio_feature_lens - 1) // 2 + 1
audio_seq_lens = (audio_seq_lens - 2) // 2 + 1
if (
self.model_type in ["qwen3_vl", "qwen3_vl_moe"]
and video_timestamps is not None
):
input_ids, offsets, modality_list = self.build_input_ids_with_timestamps(
prompt, embeddings, img_grid_thw, video_grid_thw, video_timestamps
)
else:
input_ids, offsets, modality_list = self.build_input_ids(
prompt, img_grid_thw, video_grid_thw, audio_seq_lens=audio_seq_lens
)
assert all(isinstance(modality, Modality) for modality in modality_list)
mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index(
spatial_merge_size=self.hf_config.vision_config.spatial_merge_size,
image_token_id=self.mm_tokens.image_token_id,
@@ -285,19 +412,41 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
model_type=self.model_type,
input_ids=torch.tensor(input_ids, dtype=torch.long).unsqueeze(0),
image_grid_thw=img_grid_thw,
video_grid_thw=video_grid_thw,
second_per_grid_ts=second_per_grid_ts,
use_audio_in_video=False,
audio_seqlens=(
audio_feature_lens if self.model_type == "qwen3_omni_moe" else None
),
audio_token_id=getattr(self.hf_config, "audio_token_id", None),
audio_start_token_id=self.audio_start_token_id,
position_id_per_seconds=getattr(
self.hf_config, "position_id_per_seconds", None
),
tokens_per_second=getattr(
self.hf_config.vision_config, "tokens_per_second", None
),
)
mrope_positions = mrope_positions.squeeze(1)
mm_items = [
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=offsets,
precomputed_embeddings=embeddings,
mm_items = []
consumed_per_modality = {}
for modality, offset in zip(modality_list, offsets):
num_tokens = offset[1] - offset[0] + 1
embedding_start = consumed_per_modality.get(modality, 0)
embedding_slice = embeddings[modality][
embedding_start : embedding_start + num_tokens
]
consumed_per_modality[modality] = embedding_start + num_tokens
logger.info(f"Get embedding slice for {modality}, num_tokens={num_tokens}")
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
precomputed_embeddings=embedding_slice,
)
)
]
return {
"input_ids": input_ids,

View File

@@ -2984,6 +2984,22 @@ class ServerArgs:
self.disaggregation_ib_device
)
# Validate model type: only support Qwen models for now
hf_config = self.get_model_config().hf_config
model_arch = hf_config.architectures[0]
if (self.encoder_only or self.language_only) and model_arch not in [
"Qwen2VLForConditionalGeneration",
"Qwen3VLForConditionalGeneration",
"Qwen2_5_VLForConditionalGeneration",
"Qwen3VLMoeForConditionalGeneration",
"Qwen3OmniMoeForConditionalGeneration",
"Qwen2AudioForConditionalGeneration",
"Qwen2_5OmniForConditionalGeneration",
]:
raise ValueError(
f"Model type {model_arch} is not supported for encoder disaggregation, only Qwen models are supported for now."
)
def _validate_ib_devices(self, device_str: str) -> Optional[str]:
"""
Validate IB devices before passing to mooncake.