Add external multimodal processors to the Rust frontend (#39329)

This commit is contained in:
Lianmin Zheng
2026-09-15 15:38:47 -07:00
committed by GitHub
parent 2929a39927
commit b803cfa0c4
20 changed files with 1027 additions and 197 deletions
+4 -1
View File
@@ -2276,7 +2276,7 @@ class Scheduler(
self.rust_server = None
return
rust_server = RustServer.launch(self)
rust_server = self.get_rust_server_class().launch(self)
self.rust_server = rust_server
# The rust server *is* the ingress source: SchedulerRequestReceiver
# drains its request ring (rust_server_mode) instead of a zmq socket.
@@ -2284,6 +2284,9 @@ class Scheduler(
# Park the idle loop on the request ring within the rank-0 rust-server
self.idle_sleeper = RustServerIdleSleeper(rust_server)
def get_rust_server_class(self) -> type[RustServer]:
return RustServer
def rust_server_tokenizer_path(self) -> str:
return get_serving().tokenizer_path
+14 -6
View File
@@ -1,11 +1,12 @@
"""Configuration handoff and CPU placement for the embedded Rust server."""
"""Configuration and CPU placement for the embedded Rust server."""
from __future__ import annotations
import json
import logging
import os
from typing import TYPE_CHECKING, List, Optional, Tuple
from types import ModuleType
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.managers.utils import compute_num_reserved_tokens
@@ -25,8 +26,10 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _build_server_args(scheduler: Scheduler) -> ServerArgs:
"""The typed launch handoff for the scheduler's embedded Rust server:
def _build_server_args(
scheduler: Scheduler, *, extension: Optional[ModuleType] = None
) -> ServerArgs:
"""The typed launch configuration for the scheduler's embedded Rust server:
the ``server_args`` fields it reads, the already-resolved
``model_config``, and launch-time facts — as the Rust extension's own
``ServerArgs`` class. Its constructor takes every field as a required
@@ -35,7 +38,7 @@ def _build_server_args(scheduler: Scheduler) -> ServerArgs:
running on a silently-defaulted knob."""
from sglang.srt.rust_extensions import load_rust_extension
ext = load_rust_extension("sglang.srt.rust_extensions._server")
ext = extension or load_rust_extension("sglang.srt.rust_extensions._server")
sa = resolving_view(scheduler.server_args)
mc = scheduler.model_config
@@ -100,6 +103,7 @@ def _build_server_args(scheduler: Scheduler) -> ServerArgs:
def _partition_cores(
mm_workers: int = 0,
server_core_budget: Optional[Callable[[int, int], int]] = None,
) -> Tuple[Optional[List[int]], Optional[List[int]]]:
"""Split this rank's allowed cores into ``(launch_cores, server_cores)``.
@@ -138,7 +142,11 @@ def _partition_cores(
# once bounded. The budget covers the CPU-hot threads (MM workers, plus
# the I/O-shaped tokenizer/ingress/egress/api ones that are rarely all hot
# at once) and leaves the rest of the node to the scheduler ranks.
pool_budget = max(8, mm_workers + 4)
pool_budget = (
server_core_budget(len(allowed), mm_workers)
if server_core_budget is not None
else max(8, mm_workers + 4)
)
server_cores = allowed[reserve : reserve + pool_budget]
logger.info(
"rust server cores=%s, scheduler launch cores=%s",
+69 -42
View File
@@ -3,7 +3,7 @@
The Rust server replaces the Python api-server + `TokenizerManager` +
`DetokenizerManager` stack, running them as Rust threads inside the scheduler
process. This wrapper keeps all `SGLANG_RUST_SERVER` plumbing — startup,
CPU-core partitioning, the typed `server_args` handoff, and control-response
CPU-core partitioning, the typed `server_args`, and control-response
routing — out of `scheduler.py`. The scheduler holds an `Optional[RustServer]`
and delegates to it.
"""
@@ -14,6 +14,7 @@ import logging
import os
from array import array
from itertools import chain
from types import ModuleType
from typing import TYPE_CHECKING, Any, List, Optional
import msgspec
@@ -39,8 +40,9 @@ from sglang.srt.utils.network import NetworkAddress
if TYPE_CHECKING:
from sglang.srt.managers.io_struct import BatchTokenIDOutput
from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.rust_extensions._server import MmSpec, Server
from sglang.srt.rust_extensions._server import MmEncodedResult, MmSpec, Server
logger = logging.getLogger(__name__)
@@ -62,8 +64,50 @@ class RustServer:
self.server = server
self.http_port = http_port
self.mm_spec = mm_spec
self._multimodal_enabled = mm_spec is not None
self._max_per_poll = max_per_poll
@classmethod
def _load_extension(cls) -> ModuleType:
from sglang.srt.rust_extensions import load_rust_extension
return load_rust_extension("sglang.srt.rust_extensions._server")
def _start_multimodal(self, scheduler: Scheduler) -> None:
"""Start the model's Rust workers and retain their scheduler-side state."""
mm_host = RustMmProcessor(
server_args=scheduler.server_args,
model_config=scheduler.model_config,
processor=scheduler.processor,
)
mm_spec = mm_host.resolve_spec()
if mm_spec is None:
supported = sorted(
set(chain.from_iterable(f.model_types for f in RUST_MM_FAMILIES))
)
raise RuntimeError(
"SGLANG_RUST_SERVER=1: no Rust MM pipeline for "
f"model_type={scheduler.model_config.hf_config.model_type!r} "
f"(supported: {', '.join(supported)}; "
"images only). Unset SGLANG_RUST_SERVER to serve this model."
)
self.server.start_mm_workers(self._build_mm_spec(mm_spec), mm_host.mm_workers)
self.mm_spec = mm_spec
@staticmethod
def _server_core_budget(allowed_core_count: int, mm_workers: int) -> int:
"""Maximum cores available to the Rust frontend and MM workers."""
return max(8, mm_workers + 4)
@classmethod
def _partition_cores(
cls, mm_workers: int = 0
) -> tuple[Optional[List[int]], Optional[List[int]]]:
return _partition_cores(
mm_workers=mm_workers,
server_core_budget=cls._server_core_budget,
)
@classmethod
def launch(cls, scheduler: Scheduler) -> RustServer:
"""Start the embedded Rust server threads and bind the listen port.
@@ -71,14 +115,9 @@ class RustServer:
The caller gates this (``SGLANG_RUST_SERVER`` + rank 0); this always
creates.
"""
from sglang.srt.rust_extensions import load_rust_extension
Server = load_rust_extension("sglang.srt.rust_extensions._server").Server
# Force turn off HF tokenizers rayon's unpinned global thread pool.
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
server_args = scheduler.server_args
# Preserve the DP startup log; ports use node-local offsets.
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
if get_exec().moe.is_ep_scale_joiner:
@@ -94,7 +133,7 @@ class RustServer:
listen_port = get_serving().port + local_dp_rank
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str()
launch_cores, server_cores = _partition_cores(
launch_cores, server_cores = cls._partition_cores(
mm_workers=(
(get_mm().mm_processor_worker_num or RustMmProcessor.AUTO_MM_WORKERS)
if scheduler.model_config.is_multimodal
@@ -102,16 +141,17 @@ class RustServer:
)
)
server = Server(
_build_server_args(scheduler),
# None -> run unpinned; the list carries the pinning decision.
extension = cls._load_extension()
server = extension.Server(
_build_server_args(scheduler, extension=extension),
# None runs unpinned; otherwise the list carries the pinning decision.
cores=server_cores,
port_offset=local_dp_rank,
)
instance = cls(server, http_port=listen_port)
# Multimodal models must have a Rust pipeline — there is no Python
# fallback.
mm_spec = None
if scheduler.model_config.is_multimodal:
# New threads inherit the spawning thread's affinity, and this launch
# thread still holds the full mask. Narrow it first so every MM thread
@@ -125,23 +165,8 @@ class RustServer:
logger.warning(
"rust server: cannot confine mm threads to server cores: %s", e
)
mm_host = RustMmProcessor(
server_args=server_args,
model_config=scheduler.model_config,
processor=scheduler.processor,
)
mm_spec = mm_host.resolve_spec()
if mm_spec is None:
supported = sorted(
set(chain.from_iterable(f.model_types for f in RUST_MM_FAMILIES))
)
raise RuntimeError(
"SGLANG_RUST_SERVER=1: no Rust MM pipeline for "
f"model_type={scheduler.model_config.hf_config.model_type!r} "
f"(supported: {', '.join(supported)}; "
"images only). Unset SGLANG_RUST_SERVER to serve this model."
)
server.start_mm_workers(cls._build_mm_spec(mm_spec), mm_host.mm_workers)
instance._start_multimodal(scheduler)
instance._multimodal_enabled = True
# Narrow the scheduler thread only after the server threads are launched.
if launch_cores is not None:
@@ -162,7 +187,11 @@ class RustServer:
dp_note,
)
return cls(server, http_port=listen_port, mm_spec=mm_spec)
return instance
def _wrap_mm_result(self, entry: MmEncodedResult) -> MultimodalProcessorOutput:
assert self.mm_spec is not None
return RustMmProcessor.wrap_encoded(self.mm_spec, entry)
def wait_request(self, timeout_ms: int) -> None:
"""Block until a request is pushed into the in-process ring or the timeout
@@ -215,20 +244,20 @@ class RustServer:
ids.frombytes(ids_view[pos : pos + nbytes])
obj.input_ids = ids
pos += nbytes
if self.mm_spec is not None and isinstance(obj, TokenizedGenerateReqInput):
if self._multimodal_enabled and isinstance(obj, TokenizedGenerateReqInput):
# The buffers were parked in the Rust result store before the
# ring push; wrapping them into tensors is the only Python step
# of the Rust path. `None` for a text-only request on a
# multimodal model.
encoded = self.server.take_mm_result(obj.rid)
if encoded is not None:
obj.mm_inputs = RustMmProcessor.wrap_encoded(self.mm_spec, encoded)
mm_result = self.server.take_mm_result(obj.rid)
if mm_result is not None:
obj.mm_inputs = self._wrap_mm_result(mm_result)
out.append(obj)
return out
def push_control_output(self, recv_req, output) -> None:
"""Push a control-request response through the egress ring to the waiting
request (routed by rid), encoded as **msgpack** (the ring's native
request (routed by rid), encoded as **msgpack** (the ring's message
format).
A msgspec struct is converted to a *named map* (``structs.asdict``, since
@@ -407,16 +436,14 @@ class RustServer:
len(rids),
)
@staticmethod
def _build_mm_spec(spec: RustMmSpec) -> MmSpec:
"""The typed MM handoff for ``Server.start_mm_workers``: the
@classmethod
def _build_mm_spec(cls, spec: RustMmSpec) -> MmSpec:
"""The typed MM configuration for ``Server.start_mm_workers``: the
:class:`RustMmSpec` fields the Rust pipeline consumes, as the Rust
extension's own ``MmSpec`` class (same required-keyword contract as
:meth:`_build_server_args`; ``family`` / ``resample`` become the
:func:`_build_server_args`; ``family`` / ``resample`` become the
extension's ``MmFamily`` / ``MmResample`` enums)."""
from sglang.srt.rust_extensions import load_rust_extension
ext = load_rust_extension("sglang.srt.rust_extensions._server")
ext = cls._load_extension()
family = {"qwen_vl": ext.MmFamily.QwenVl}[spec.family]
resample = {"aten_u8": ext.MmResample.AtenU8, "pil": ext.MmResample.Pil}[
spec.resample
+16
View File
@@ -13,6 +13,22 @@
pub enum TensorData {
F32(Vec<f32>),
I64(Vec<i64>),
/// Raw BF16 bits, exposed to numpy as u16 without copying the allocation.
Bf16(Vec<u16>),
}
impl TensorData {
pub fn len(&self) -> usize {
match self {
Self::F32(data) => data.len(),
Self::I64(data) => data.len(),
Self::Bf16(data) => data.len(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
pub struct Tensor {
+1 -1
View File
@@ -1,7 +1,7 @@
//! OpenAI-compatible generation endpoints.
//!
//! The HTTP adapter stays deliberately thin: Dynamo owns the standard OpenAI
//! request and response primitives. Native [`ChunkEvent`] values remain the one
//! request and response primitives. Scheduler [`ChunkEvent`] values remain the one
//! backend output type for both unary and streaming responses.
use axum::{Router, http::StatusCode, response::Response};
@@ -56,7 +56,7 @@ pub(super) struct SubmittedChoice {
pub(super) struct ChoiceExtensions {
matched_stop: Option<serde_json::Value>,
/// Dynamo's enum covers the standard values. Python additionally exposes
/// `abort`, and native unknown finish types are preserved rather than lost.
/// `abort`; unrecognized scheduler finish types are preserved as well.
finish_reason_override: Option<String>,
}
@@ -445,7 +445,7 @@ fn completion_choice(
Matched::Token(id) => serde_json::json!(id),
Matched::Str(value) => serde_json::json!(value),
// Python's OpenAI schema supports an integer or string here, not a
// multi-token list. Preserve the native value rather than dropping it.
// multi-token list. Preserve the original token IDs.
Matched::Tokens(ids) => serde_json::json!(ids),
});
(
@@ -2,7 +2,7 @@
//
//! Hugging Face tokenizer configs contain Jinja templates. SGLang also accepts
//! legacy conversation JSON files and the names in Python's template registry.
//! Legacy definitions are rendered by a native port of Python's
//! Legacy definitions are rendered by a Rust implementation of Python's
//! `Conversation.get_prompt()` so there is exactly one implementation of the
//! per-style formatting logic (no Jinja translation to drift).
@@ -480,7 +480,7 @@ mod tests {
}
#[test]
fn json_legacy_template_is_rendered_natively() {
fn json_legacy_template_is_rendered() {
let base = std::env::temp_dir().join(format!(
"sglang-openai-template-base-{}-test.json",
std::process::id()
@@ -49,7 +49,7 @@ impl Default for LegacySpec {
}
}
/// Native port of Python `generate_chat_conv` + `Conversation.get_prompt()`:
/// Rust implementation of Python `generate_chat_conv` + `Conversation.get_prompt()`:
/// fold system messages into the system prompt, keep user/assistant messages in
/// order, always append the assistant opening, then render per `sep_style`.
#[derive(Clone)]
+172 -48
View File
@@ -5,7 +5,7 @@
//! (`_server`) and the classes exposed to the scheduler — the boot config
//! ([`ServerArgs`] and its parts, constructed by keyword from Python; their
//! `#[pyclass]`es and constructors live in `message::config`), [`Server`]
//! (boot, `recv_requests`/`wait_request`, `push_*`, MM handoff, shutdown),
//! (boot, `recv_requests`/`wait_request`, `push_*`, MM results, shutdown),
//! [`RequestBatch`] and [`MmEncodedResult`]. Everything behind that boundary —
//! receiving requests, encoding multimodal inputs, tokenizing, detokenizing,
//! SSE streaming, and so on — is implemented purely in Rust and never touches
@@ -17,23 +17,40 @@ mod multi_modality;
mod tokenizer_manager;
mod utils;
pub use message::config::{
DefaultSamplingParams, DisaggregationMode, MmFamily, MmResample, MmSpec, ModelConfig,
RustServerServerArgs, ServerArgs,
};
pub use message::multimodal::MmItem;
pub use message::request::{MmWorkItem, ProcessorExtensions};
pub use multi_modality::payload::{ResolvedMediaWork, resolve_media_work};
pub use multi_modality::result_store::{
ExternalMmEncodedEntry, ExternalMmItem, MmEncodedEntry, MmModality, MmTokenIds,
};
pub use multi_modality::worker::{MmProcessOutput, MmProcessor};
pub use sglang_mm::pipeline::{Tensor, TensorData};
pub use tokenizer_manager::tokenizer::TextTokenizer;
use std::collections::BTreeMap;
use std::sync::Arc;
use pyo3::prelude::*;
use pyo3::pybacked::PyBackedBytes;
use pyo3::types::PyBytes;
use crate::message::config::{
DefaultSamplingParams, DisaggregationMode, MmFamily, MmResample, MmSpec, ModelConfig,
RuntimeConfig, RustServerServerArgs, ServerArgs,
};
use crate::message::config::RuntimeConfig;
use crate::utils::startup::{listen_addr, value_error};
use crate::utils::{logging, runtime};
/// One drained MM result (see [`Server::take_mm_result`]), consumed by
/// `RustMmProcessor.wrap_encoded` to build the scheduler's
/// `MultimodalProcessorOutput`.
/// One drained MM result (see [`Server::take_mm_result`]).
///
/// Built-in results are consumed by `RustMmProcessor.wrap_encoded` to build
/// `MultimodalProcessorOutput`. External integrations consume `external_items`
/// and `external_token_ids` in their own Python wrappers; the built-in fields
/// are empty in that case.
#[pyclass(frozen, get_all)]
struct MmEncodedResult {
// General fields.
pub struct MmEncodedResult {
// General fields for the built-in processor path.
/// All items' `pixel_values` concatenated as flat `f32` with logical shape
/// `[sum(t*h*w), feature_dim]`; present on the inline (single-rank) path.
features: Option<Py<numpy::PyArray1<f32>>>,
@@ -57,12 +74,121 @@ struct MmEncodedResult {
/// M-RoPE delta, `max(mrope) + 1 - seq_len`, added to the plain sequence
/// position during decoding.
mrope_delta: i64,
// Fields for external processor integrations.
external_items: Vec<Py<ExternalMmItemResult>>,
external_token_ids: Option<MmTokenIds>,
}
/// One media item returned to an external Python integration. The NumPy view
/// owns the feature allocation.
#[pyclass(frozen, get_all)]
pub struct ExternalMmItemResult {
modality: MmModality,
feature: MmFeatureArray,
shape: Vec<usize>,
hash: u64,
offsets: Vec<(u32, u32)>,
model_specific_data: BTreeMap<String, i64>,
}
#[derive(IntoPyObjectRef)]
enum MmFeatureArray {
#[pyo3(transparent)]
F32(Py<numpy::PyArray1<f32>>),
#[pyo3(transparent)]
I64(Py<numpy::PyArray1<i64>>),
#[pyo3(transparent)]
Bf16(Py<numpy::PyArray1<u16>>),
}
#[pymethods]
impl ExternalMmItemResult {
/// BF16 is exposed as raw u16 bits for Python to reinterpret without a copy.
#[getter]
fn feature_is_bf16(&self) -> bool {
matches!(self.feature, MmFeatureArray::Bf16(_))
}
}
impl MmEncodedResult {
fn from_entry(py: Python<'_>, entry: MmEncodedEntry) -> PyResult<Self> {
use numpy::IntoPyArray;
match entry {
MmEncodedEntry::Qwen(entry) => {
let (features, shm_names) = match entry.features {
multi_modality::result_store::FeatureStore::Inline(v) => {
(Some(v.into_pyarray(py).unbind()), None)
}
// The segments — and the duty to unlink — move to Python here;
// `materialize()` unlinks after the post-broadcast clone on each rank.
multi_modality::result_store::FeatureStore::Shm(segments) => (
None,
Some(segments.into_iter().map(|s| s.into_name()).collect()),
),
};
Ok(Self {
features,
shm_names,
hashes: entry.hashes,
offsets: entry.offsets,
grids: entry.grids.iter().map(|g| (g[0], g[1], g[2])).collect(),
mrope: entry.mrope.into_pyarray(py).unbind(),
mrope_delta: entry.mrope_delta,
external_items: Vec::new(),
external_token_ids: None,
})
}
MmEncodedEntry::External(entry) => {
let external_items = entry
.items
.into_iter()
.map(|item| {
let feature = match item.feature.data {
TensorData::F32(data) => {
MmFeatureArray::F32(data.into_pyarray(py).unbind())
}
TensorData::I64(data) => {
MmFeatureArray::I64(data.into_pyarray(py).unbind())
}
TensorData::Bf16(data) => {
MmFeatureArray::Bf16(data.into_pyarray(py).unbind())
}
};
Py::new(
py,
ExternalMmItemResult {
modality: item.modality,
feature,
shape: item.feature.shape,
hash: item.hash,
offsets: item.offsets,
model_specific_data: item.model_specific_data,
},
)
})
.collect::<PyResult<_>>()?;
Ok(Self {
features: None,
shm_names: None,
hashes: Vec::new(),
offsets: Vec::new(),
grids: Vec::new(),
mrope: Vec::<i64>::new().into_pyarray(py).unbind(),
mrope_delta: 0,
external_items,
external_token_ids: Some(entry.token_ids),
})
}
}
}
}
/// Columnar request batch handed to Python by [`Server::recv_requests`].
/// `frozen`: immutable snapshot, so field access never contends on a borrow.
#[pyclass(frozen, get_all)]
struct RequestBatch {
pub struct RequestBatch {
/// One msgpack scalar header per request (`input_ids` omitted).
headers: Vec<Py<PyBytes>>,
/// The raw-data plane today just all requests' raw little-endian int64
@@ -75,7 +201,7 @@ struct RequestBatch {
/// Handle owned by the Python scheduler process. Construct once via
/// [`Server::start`], then poll it from the scheduler event loop.
#[pyclass]
struct Server {
pub struct Server {
rt: runtime::Runtime,
}
@@ -96,7 +222,7 @@ impl Server {
// pyo3 `#[new]` constructor: the wide arg list is the Python-facing boot
// surface (all optional overrides), not a call-site ergonomics problem.
#[allow(clippy::too_many_arguments)]
fn start(
pub fn start(
server_args: ServerArgs,
port_offset: Option<u16>, // DP rank; listen on server_args.port + offset
to_scheduler_cap: usize,
@@ -126,14 +252,14 @@ impl Server {
server_args: std::sync::Arc::new(server_args),
};
let rt = runtime::start(cfg).map_err(|e| value_error("runtime start failed", e))?;
Ok(Server { rt })
Ok(Self { rt })
}
/// Non-blocking drain of the to_scheduler channel, returned **columnar** as an
/// [`RequestBatch`] so the large `input_ids` tensor never goes through
/// msgpack (see the field docs for the layout).
#[pyo3(signature = (max = 256))]
fn recv_requests(&self, py: Python<'_>, max: usize) -> PyResult<RequestBatch> {
pub fn recv_requests(&self, py: Python<'_>, max: usize) -> PyResult<RequestBatch> {
let cols = self.rt.to_scheduler_rx.drain(max);
let headers = cols
.headers
@@ -154,7 +280,7 @@ impl Server {
/// Park up to `timeout_ms` for an incoming request so the idle scheduler loop
/// sleeps instead of spinning at 100% CPU.
#[pyo3(signature = (timeout_ms = 1000))]
fn wait_request(&self, py: Python<'_>, timeout_ms: u64) -> bool {
pub fn wait_request(&self, py: Python<'_>, timeout_ms: u64) -> bool {
py.detach(|| {
self.rt
.to_scheduler_rx
@@ -165,7 +291,7 @@ impl Server {
/// Push a whole decode batch as ONE frame: a columnar msgpack `header` plus
/// the raw `data_cols` (per-column `bytes`), concatenated here. Blocks for
/// backpressure; `False` only on shutdown.
fn push_decode_result_batch(
pub fn push_decode_result_batch(
&self,
py: Python<'_>,
header: &[u8],
@@ -180,7 +306,7 @@ impl Server {
/// Push a control-request result. Blocks for backpressure; `False` only on
/// shutdown.
fn push_control_result(&self, py: Python<'_>, rid: &str, payload: &[u8]) -> bool {
pub fn push_control_result(&self, py: Python<'_>, rid: &str, payload: &[u8]) -> bool {
self.push_frame(
py,
crate::message::response::frame_control_result(rid, payload),
@@ -189,7 +315,7 @@ impl Server {
/// Route a terminal failure back to request `rid`. Blocks for backpressure;
/// `False` only on shutdown.
fn push_error(&self, py: Python<'_>, rid: &str, message: &str) -> bool {
pub fn push_error(&self, py: Python<'_>, rid: &str, message: &str) -> bool {
self.push_frame(py, crate::message::response::frame_error(rid, message))
}
@@ -198,7 +324,7 @@ impl Server {
/// `RustServer._build_mm_spec`). Image-only requests are processed entirely
/// in Rust and parked for [`Server::take_mm_result`]; anything the pipeline
/// cannot serve is rejected back to the client — there is no Python fallback.
fn start_mm_workers(&self, spec: MmSpec, workers: usize) -> PyResult<()> {
pub fn start_mm_workers(&self, spec: MmSpec, workers: usize) -> PyResult<()> {
self.rt
.start_mm_workers(spec, workers)
.map_err(|e| value_error("mm spec", e))
@@ -212,39 +338,27 @@ impl Server {
/// Runs on the scheduler loop between decode steps, so any per-byte work
/// here — memcpy or hashing, tens of MB per image-heavy request — would
/// stall every running request's ITL. Hence the worker-precomputed `hashes`.
fn take_mm_result(&self, py: Python<'_>, rid: &str) -> Option<MmEncodedResult> {
use numpy::IntoPyArray;
let res = self.rt.mm_results.take(rid)?;
let (features, shm_names) = match res.features {
multi_modality::result_store::FeatureStore::Inline(v) => {
(Some(v.into_pyarray(py).unbind()), None)
}
// The segments — and the duty to unlink — move to Python here;
// `materialize()` unlinks after the post-broadcast clone on each rank.
multi_modality::result_store::FeatureStore::Shm(segments) => (
None,
Some(segments.into_iter().map(|s| s.into_name()).collect()),
),
};
Some(MmEncodedResult {
features,
shm_names,
grids: res.grids.iter().map(|g| (g[0], g[1], g[2])).collect(),
hashes: res.hashes,
offsets: res.offsets,
mrope: res.mrope.into_pyarray(py).unbind(),
mrope_delta: res.mrope_delta,
})
pub fn take_mm_result(&self, py: Python<'_>, rid: &str) -> PyResult<Option<MmEncodedResult>> {
self.rt
.mm_results
.take(rid)
.map(|entry| MmEncodedResult::from_entry(py, entry))
.transpose()
}
/// Signal all threads to stop (best effort).
fn shutdown(&self) {
pub fn shutdown(&self) {
self.rt.request_shutdown();
}
}
impl Server {
/// Start the shared worker pool with a processor supplied by an external
/// model package. The default Python API retains the built-in Qwen path.
pub fn start_mm_workers_with_processor(&self, processor: Arc<dyn MmProcessor>, workers: usize) {
self.rt.start_mm_workers_with_processor(processor, workers);
}
/// Hand one already-framed message to the ring. Shared by every push path —
/// they differ solely in how the frame is built. `false` only on shutdown.
#[inline]
@@ -259,8 +373,9 @@ impl Server {
}
}
#[pymodule]
fn _server(m: &Bound<'_, PyModule>) -> PyResult<()> {
/// Register all Python boundary types used by [`Server`]. External
/// model-package modules call this before exposing their wrapper server.
pub fn register_boundary_types(m: &Bound<'_, PyModule>) -> PyResult<()> {
logging::init_tracing();
m.add_class::<DisaggregationMode>()?;
m.add_class::<DefaultSamplingParams>()?;
@@ -269,8 +384,17 @@ fn _server(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<MmFamily>()?;
m.add_class::<MmResample>()?;
m.add_class::<MmSpec>()?;
m.add_class::<Server>()?;
m.add_class::<RequestBatch>()?;
m.add_class::<MmEncodedResult>()?;
m.add_class::<ExternalMmItemResult>()?;
m.add_class::<MmModality>()?;
m.add_class::<MmTokenIds>()?;
Ok(())
}
#[pymodule]
fn _server(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Server>()?;
register_boundary_types(m)?;
Ok(())
}
@@ -66,8 +66,8 @@ pub enum FinishReason {
/// This arm is why the outer enum is untagged: a finish reason added Python-side
/// must not fail the header decode, which rejects the whole frame — every
/// request in the batch, not just the one that carried it.
// Keep the native frame compact even when HTTP/rendering dependencies turn
// on serde_json's large `preserve_order` map representation.
// Keep the scheduler response compact even when HTTP/rendering dependencies
// turn on serde_json's large `preserve_order` map representation.
Unknown(Box<serde_json::Map<String, serde_json::Value>>),
}
+214 -17
View File
@@ -1,12 +1,12 @@
//! The `/generate` request path: the HTTP body and its per-request fan-out
//! ([`GenerateBody`] → [`GenerateRequest`]s).
use std::collections::HashSet;
use std::collections::{BTreeMap, HashSet};
use std::sync::LazyLock;
use bytes::Bytes;
use itertools::izip;
use serde::Deserialize;
use serde::{Deserialize, de::DeserializeOwned};
use super::io_struct::{ControlRequest, TokenizedGenerateReqInput};
use super::multimodal::{self, MmDataInput, MmItem};
@@ -44,18 +44,53 @@ const MAX_BROADCAST_CLONE_BYTES: usize = 64 << 20;
/// the wire form does not); 8 is the ceiling of that range, not a worst case.
const JSON_TO_HEAP_FACTOR: usize = 8;
/// Top-level fields in this namespace belong to the selected multimodal
/// processor. Everything else unknown to [`GenerateBody`] keeps Python's
/// accepted-but-ignored behavior.
const PROCESSOR_EXTENSION_PREFIX: &str = "multimodal_";
/// Model-owned request fields. The shared server preserves and batches their
/// MessagePack value representation; the selected processor deserializes that
/// map into its own concrete schema.
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(transparent)]
pub struct ProcessorExtensions(BTreeMap<String, rmpv::Value>);
impl ProcessorExtensions {
/// Deserialize the model-agnostic value tree directly into the selected
/// processor's schema. This does not encode or decode MessagePack bytes.
pub fn deserialize<T: DeserializeOwned>(self) -> Result<T, String> {
let fields = self
.0
.into_iter()
.map(|(name, value)| (rmpv::Value::from(name), value))
.collect();
rmpv::ext::from_value(rmpv::Value::Map(fields))
.map_err(|error| format!("invalid processor extensions: {error}"))
}
pub(crate) fn is_empty(&self) -> bool {
self.0.is_empty()
}
fn values(&self) -> impl Iterator<Item = &rmpv::Value> {
self.0.values()
}
}
impl FromIterator<(String, rmpv::Value)> for ProcessorExtensions {
fn from_iter<T: IntoIterator<Item = (String, rmpv::Value)>>(iter: T) -> Self {
Self(iter.into_iter().collect())
}
}
/// The `/generate` wire body before batch splitting: `text`/`input_ids`/`sampling_params`
/// each scalar-or-list, fanned into per-request [`GenerateRequest`]s by
/// [`into_requests`](GenerateBody::into_requests).
///
/// Unknown keys are IGNORED, matching Python: FastAPI builds `GenerateReqInput`
/// as a pydantic dataclass, which drops extras. `deny_unknown_fields` here turned
/// every `GenerateReqInput` field this server has not ported — `priority`,
/// `extra_key`, `session_id`, `session_params`, `return_sampling_mask`,
/// `custom_logit_processor`, and ~40 more — into a 400, so a client that worked
/// against the Python server broke against this one. The cost of dropping it is
/// that a typo (`temperature`) is silently ignored rather than reported; that is
/// the same trade Python already makes.
/// Unknown keys are ignored, matching Python, except `multimodal_*` fields. Those
/// are opaque processor extensions: this layer only fans them out with the
/// request batch and passes them to the selected multimodal processor.
#[derive(Debug, Clone, Default, Deserialize)]
pub struct GenerateBody {
/// Optional client-supplied request id(s): a single string (a batch fans it
@@ -105,6 +140,10 @@ pub struct GenerateBody {
pub mm_hashes: Option<OneOrMany<Vec<String>>>,
pub video_data: Option<MmDataInput>,
pub audio_data: Option<MmDataInput>,
/// Model-specific multimodal fields, retained without teaching the shared
/// request schema their contents. Other unknown fields remain ignored.
#[serde(flatten)]
processor_extensions: ProcessorExtensions,
}
impl GenerateBody {
@@ -150,9 +189,7 @@ impl GenerateBody {
video_data,
audio_data,
mm_hashes,
// Unported `GenerateReqInput` fields land here and are dropped, as they
// are on the Python path.
..
processor_extensions,
} = self;
// Cap the batch BEFORE the columns below allocate anything. Reading the
@@ -359,6 +396,7 @@ impl GenerateBody {
let images = multimodal::fan_out(image_data, n, is_batch, "image_data")?;
let videos = multimodal::fan_out(video_data, n, is_batch, "video_data")?;
let audios = multimodal::fan_out(audio_data, n, is_batch, "audio_data")?;
let processor_extensions = split_extension_columns(processor_extensions, n, is_batch)?;
// Every column above is exactly `n` long, so zip them by value: each
// request takes ownership of its cell, with no indexing or bounds checks.
@@ -380,6 +418,7 @@ impl GenerateBody {
images,
videos,
audios,
processor_extensions,
)
.map(
|(
@@ -400,11 +439,12 @@ impl GenerateBody {
image_data,
video_data,
audio_data,
processor_extensions,
)| GenerateRequest {
rid,
text,
input_ids,
// Native text prompts keep the post-processor specials; the
// Plain text prompts keep the post-processor specials; the
// chat flow sets this explicitly.
skip_special_tokens: false,
sampling_params,
@@ -426,7 +466,7 @@ impl GenerateBody {
decode_tp_size,
routed_dp_rank,
disagg_prefill_dp_rank,
mm: pack_mm(image_data, video_data, audio_data),
mm: pack_mm(image_data, video_data, audio_data, processor_extensions),
},
)
.collect();
@@ -445,18 +485,72 @@ fn pack_mm(
image_data: Vec<MmItem>,
video_data: Vec<MmItem>,
audio_data: Vec<MmItem>,
processor_extensions: ProcessorExtensions,
) -> Option<Box<MmData>> {
if image_data.is_empty() && video_data.is_empty() && audio_data.is_empty() {
if image_data.is_empty()
&& video_data.is_empty()
&& audio_data.is_empty()
&& processor_extensions.is_empty()
{
return None;
}
Some(Box::new(MmData {
image_data,
video_data,
audio_data,
processor_extensions,
..Default::default()
}))
}
fn split_extension_columns(
fields: ProcessorExtensions,
n: usize,
is_batch: bool,
) -> Result<Vec<ProcessorExtensions>, Error> {
let mut requests = vec![ProcessorExtensions::default(); n];
for (name, value) in fields.0 {
if !name.starts_with(PROCESSOR_EXTENSION_PREFIX) || value.is_nil() {
continue;
}
if !is_batch {
requests[0].0.insert(name, value);
continue;
}
let rmpv::Value::Array(values) = value else {
return Err(Error::Validation(format!(
"{name} must be a list for batch processing"
)));
};
if values.is_empty() {
for request in &mut requests {
request
.0
.insert(name.clone(), rmpv::Value::Array(Vec::new()));
}
continue;
}
if values.len() != n {
return Err(Error::Validation(format!(
"{name} list length {} does not match batch size {n}",
values.len()
)));
}
for (request, value) in requests.iter_mut().zip(values) {
request.0.insert(name.clone(), value);
}
}
Ok(requests)
}
fn extension_value_present(value: &rmpv::Value) -> bool {
match value {
rmpv::Value::Nil => false,
rmpv::Value::Array(values) => values.iter().any(extension_value_present),
_ => true,
}
}
/// One request handed to the MM worker pool: the rid to correlate the result,
/// plus the owned inputs from [`GenerateRequest::take_mm_work`].
#[derive(Debug)]
@@ -474,6 +568,7 @@ pub struct MmWorkItem {
pub image_data: Vec<MmItem>,
pub video_data: Vec<MmItem>,
pub audio_data: Vec<MmItem>,
pub processor_extensions: ProcessorExtensions,
/// See [`MmData::prefetched`].
pub prefetched: Vec<Bytes>,
/// See [`GenerateBody::mm_hashes`].
@@ -606,6 +701,7 @@ pub struct MmData {
pub image_data: Vec<MmItem>,
pub video_data: Vec<MmItem>,
pub audio_data: Vec<MmItem>,
pub processor_extensions: ProcessorExtensions,
/// Bytes of `image_data`'s I/O-backed sources, resolved by
/// `api_server::prefetch` in `payload::io_sources` order so MM workers
/// never block on I/O. Out-of-band: the values above stay as the client
@@ -625,7 +721,13 @@ impl GenerateRequest {
/// Python `GenerateReqInput.contains_mm_input()`.
pub fn has_multimodal(&self) -> bool {
self.mm.as_ref().is_some_and(|mm| {
!mm.image_data.is_empty() || !mm.video_data.is_empty() || !mm.audio_data.is_empty()
!mm.image_data.is_empty()
|| !mm.video_data.is_empty()
|| !mm.audio_data.is_empty()
|| mm
.processor_extensions
.values()
.any(extension_value_present)
})
}
@@ -642,6 +744,7 @@ impl GenerateRequest {
work.image_data = std::mem::take(&mut m.image_data);
work.video_data = std::mem::take(&mut m.video_data);
work.audio_data = std::mem::take(&mut m.audio_data);
work.processor_extensions = std::mem::take(&mut m.processor_extensions);
work.prefetched = std::mem::take(&mut m.prefetched);
work.mm_hashes = std::mem::take(&mut m.mm_hashes);
}
@@ -751,6 +854,17 @@ fn fan_out<T: OneOrManyItem + Clone + HeapBytes>(
mod tests {
use super::*;
#[derive(Debug, Deserialize, PartialEq)]
struct TestProcessorExtensions {
multimodal_custom: TestProcessorExtension,
}
#[derive(Debug, Deserialize, PartialEq)]
#[serde(deny_unknown_fields)]
struct TestProcessorExtension {
value: i64,
}
/// Vocab size for tests that aren't about the vocab bound (see
/// `sampling::tests::TEST_VOCAB`).
const TEST_VOCAB: u64 = 1000;
@@ -972,6 +1086,89 @@ mod tests {
assert!(ps[1].has_multimodal());
}
#[test]
fn multimodal_extensions_follow_request_batch_shape() {
let single = r#"{"input_ids":[9],"image_data":"u","multimodal_placeholders":[{"type":"image","token_index":0,"item_index":0}]}"#;
let (reqs, is_batch) = requests(single).unwrap();
assert!(!is_batch);
let value = reqs[0]
.mm
.as_ref()
.unwrap()
.processor_extensions
.0
.get("multimodal_placeholders")
.unwrap();
assert_eq!(value.as_array().unwrap().len(), 1);
let batched = r#"{"input_ids":[[9],[8]],"image_data":["u","v"],"multimodal_placeholders":[[{"type":"image","token_index":0,"item_index":0}],[{"type":"image","token_index":0,"item_index":0}]]}"#;
let (reqs, is_batch) = requests(batched).unwrap();
assert!(is_batch);
assert_eq!(reqs.len(), 2);
assert!(reqs.iter().all(GenerateRequest::has_multimodal));
assert!(reqs.iter().all(|request| {
request
.mm
.as_ref()
.and_then(|mm| mm.processor_extensions.0.get("multimodal_placeholders"))
.and_then(rmpv::Value::as_array)
.is_some_and(|placeholders| placeholders.len() == 1)
}));
let invalid = r#"{"input_ids":[[9],[8]],"image_data":["u","v"],"multimodal_placeholders":[{"type":"image","token_index":0,"item_index":0}]}"#;
assert!(requests(invalid).is_err());
let generic = r#"{"input_ids":[[9],[8]],"image_data":["u","v"],"multimodal_custom":[{"value":1},{"value":2}]}"#;
let (reqs, _) = requests(generic).unwrap();
assert_eq!(
reqs[1]
.mm
.as_ref()
.unwrap()
.processor_extensions
.0
.get("multimodal_custom")
.unwrap()
.as_map()
.unwrap()[0]
.1
.as_i64(),
Some(2)
);
let extensions: TestProcessorExtensions =
requests(r#"{"input_ids":[9],"multimodal_custom":{"value":3}}"#)
.unwrap()
.0
.pop()
.unwrap()
.mm
.unwrap()
.processor_extensions
.deserialize()
.unwrap();
assert_eq!(extensions.multimodal_custom.value, 3);
for fields in [
r#"{"multimodal_custom":{"value":true}}"#,
r#"{"multimodal_custom":{"value":"3"}}"#,
r#"{"multimodal_custom":{"value":3,"unknown":0}}"#,
r#"{"multimodal_custom":{}}"#,
] {
let extensions: ProcessorExtensions = serde_json::from_str(fields).unwrap();
assert!(
extensions.deserialize::<TestProcessorExtensions>().is_err(),
"{fields}"
);
}
let (reqs, _) = requests(r#"{"text":"hi","totally_made_up":1}"#).unwrap();
assert!(reqs[0].mm.is_none());
let (reqs, _) = requests(r#"{"input_ids":[9],"multimodal_custom":null}"#).unwrap();
assert!(!reqs[0].has_multimodal());
}
/// A scalar broadcast is budget-checked before the deep clones (16 MiB ×
/// 4096 prompts would be 64 GiB and an abort); per-item lists clone nothing
/// and are never charged.
@@ -6,10 +6,89 @@
//! precomputed features, …).
use bytes::Bytes;
use sglang_mm::common::fetch::{ByteBudget, fetch_bytes_budgeted};
use sglang_mm::driver::{ImageSource, MmInput};
use crate::message::multimodal::MmItem;
use crate::message::request::MmWorkItem;
use crate::message::request::{MmWorkItem, ProcessorExtensions};
/// Fully resolved media for a multimodal processor. I/O sources were
/// prefetched on the async API layer; data URLs and bare base64 are decoded on
/// the MM worker.
pub struct ResolvedMediaWork {
pub text: Option<String>,
pub input_ids: Option<Vec<i32>>,
pub images: Vec<Bytes>,
pub videos: Vec<Bytes>,
pub audios: Vec<Bytes>,
/// Request fields owned by the selected processor rather than this shared
/// payload layer.
pub processor_extensions: ProcessorExtensions,
}
/// Resolve all modality fields in the fixed image/video/audio prefetch order.
pub fn resolve_media_work(work: MmWorkItem) -> Result<ResolvedMediaWork, String> {
resolve_media_work_with_budget(work, sglang_mm::driver::MAX_REQUEST_BYTES)
}
fn resolve_media_work_with_budget(
work: MmWorkItem,
max_request_bytes: u64,
) -> Result<ResolvedMediaWork, String> {
let MmWorkItem {
text,
input_ids,
image_data,
video_data,
audio_data,
processor_extensions,
prefetched,
mm_hashes: _,
} = work;
let mut prefetched = prefetched.into_iter();
let budget = ByteBudget::new(max_request_bytes);
let images = collect_media(image_data, &mut prefetched, "image_data", &budget)?;
let videos = collect_media(video_data, &mut prefetched, "video_data", &budget)?;
let audios = collect_media(audio_data, &mut prefetched, "audio_data", &budget)?;
if prefetched.next().is_some() {
return Err("media prefetch produced more payloads than the request consumes".into());
}
Ok(ResolvedMediaWork {
text,
input_ids,
images,
videos,
audios,
processor_extensions,
})
}
fn collect_media(
items: Vec<MmItem>,
prefetched: &mut std::vec::IntoIter<Bytes>,
field: &str,
budget: &ByteBudget,
) -> Result<Vec<Bytes>, String> {
items
.into_iter()
.map(|item| match item {
MmItem::Source(source) | MmItem::Ref { url: source } => {
if is_io_source(&source) {
let bytes = prefetched
.next()
.ok_or_else(|| format!("I/O-backed {field} source was not prefetched"))?;
budget.charge_existing(bytes.len(), field)?;
Ok(bytes)
} else {
fetch_bytes_budgeted(&source, budget).map(Bytes::from)
}
}
MmItem::Preprocessed { format } => Err(format!(
"unsupported {field} item: preprocessed `{format}` input"
)),
})
.collect()
}
/// True for sources the API layer must resolve before MM dispatch: I/O — network
/// *or* disk, since a network mount can hang past any HTTP timeout — never runs
@@ -43,12 +122,16 @@ pub fn to_mm_input(work: MmWorkItem) -> Result<MmInput, String> {
image_data,
video_data,
audio_data,
processor_extensions,
prefetched,
mm_hashes: _,
} = work;
if !video_data.is_empty() || !audio_data.is_empty() {
return Err("unsupported modality: video/audio input".into());
}
if !processor_extensions.is_empty() {
return Err("unsupported generate extensions for this processor".into());
}
let mut prefetched = prefetched.iter();
let images = image_data
.into_iter()
@@ -124,6 +207,21 @@ mod tests {
.err()
.unwrap();
assert!(err.contains("preprocessed `processor_output`"), "{err}");
let extension = MmWorkItem {
processor_extensions: std::iter::once((
"multimodal_custom".to_owned(),
rmpv::Value::Boolean(true),
))
.collect(),
..Default::default()
};
assert!(
to_mm_input(extension)
.err()
.unwrap()
.contains("unsupported generate extensions")
);
}
/// I/O-backed sources (URLs, file paths) take their prefetched bytes in walk
@@ -163,4 +261,59 @@ mod tests {
.contains("no raw image sources")
);
}
#[test]
fn resolved_media_shares_one_byte_budget_across_source_forms() {
let work = MmWorkItem {
image_data: vec![src("YWJjZA=="), src("ZWZnaA==")],
..Default::default()
};
let err = resolve_media_work_with_budget(work, 7).err().unwrap();
assert!(err.contains("request media byte budget"), "{err}");
let work = MmWorkItem {
image_data: vec![src("YWJjZA==")],
video_data: vec![src("ZWZnaA==")],
..Default::default()
};
let err = resolve_media_work_with_budget(work, 7).err().unwrap();
assert!(err.contains("request media byte budget"), "{err}");
}
#[test]
fn resolved_media_preserves_order_and_prefetched_allocations() {
let image = Bytes::from(vec![1, 2]);
let video = Bytes::from(vec![3, 4]);
let image_ptr = image.as_ptr();
let video_ptr = video.as_ptr();
let work = MmWorkItem {
image_data: vec![src("/image"), src("BQY=")],
video_data: vec![src("https://example.test/video")],
audio_data: vec![src("Bwg=")],
prefetched: vec![image, video],
..Default::default()
};
let resolved = resolve_media_work_with_budget(work, 8).unwrap();
assert_eq!(resolved.images[0].as_ptr(), image_ptr);
assert_eq!(resolved.videos[0].as_ptr(), video_ptr);
assert_eq!(resolved.images[1].as_ref(), [5, 6]);
assert_eq!(resolved.audios[0].as_ref(), [7, 8]);
for work in [
image_work(vec![src("/missing")]),
MmWorkItem {
prefetched: vec![Bytes::from_static(b"extra")],
..Default::default()
},
] {
assert!(resolve_media_work(work).is_err());
}
let work = MmWorkItem {
image_data: vec![src("/image")],
audio_data: vec![src("Bwg=")],
prefetched: vec![Bytes::from_static(b"1234")],
..Default::default()
};
assert!(resolve_media_work_with_budget(work, 5).is_err());
}
}
@@ -1,21 +1,16 @@
//! Rid-keyed parking of finished results between an MM worker and the
//! scheduler drain.
use std::collections::HashMap;
use std::collections::{BTreeMap, HashMap};
use std::sync::{Arc, Mutex};
use pyo3::prelude::*;
use sglang_mm::pipeline::Tensor;
use super::shm::{ShmSegment, shm_name};
/// One parked result: the buffers the drain-time Python adapter needs (the
/// expanded `input_ids` travel separately, via `TmEvent::MmEncoded`).
///
/// TODO(mm-families): these fields are the shape the only current family
/// (qwen_vl) produces; generalize to a named-tensor handoff when a family
/// needs a different one.
///
/// Constructed from outside the module only by tests; the worker parks every
/// real entry itself.
pub struct MmEncodedEntry {
/// The built-in Qwen drain shape.
pub struct QwenMmEncodedEntry {
pub features: FeatureStore,
/// Per item `[t, h, w]` patch grid.
pub grids: Vec<[u32; 3]>,
@@ -27,6 +22,78 @@ pub struct MmEncodedEntry {
pub mrope_delta: i64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
#[pyclass(frozen, eq, hash, skip_from_py_object)]
pub enum MmModality {
Image,
Video,
Audio,
}
/// Placeholder and boundary tokens consumed by `MultimodalProcessorOutput`.
#[derive(Clone, Debug, Default, PartialEq, Eq)]
#[pyclass(frozen, get_all, skip_from_py_object)]
pub struct MmTokenIds {
pub im_token_id: Option<i32>,
pub im_start_id: Option<i32>,
pub im_end_id: Option<i32>,
pub video_token_id: Option<i32>,
pub audio_token_id: Option<i32>,
pub audio_start_id: Option<i32>,
pub audio_end_id: Option<i32>,
}
/// One media item produced by an external processor, including its features
/// and inclusive token spans in the expanded prompt.
pub struct ExternalMmItem {
pub modality: MmModality,
pub feature: Tensor,
pub hash: u64,
pub offsets: Vec<(u32, u32)>,
/// Processor-owned integer attributes, such as a clip index and count.
pub model_specific_data: BTreeMap<String, i64>,
}
/// Encoded data produced by an external processor implementation and consumed
/// by its Python integration.
pub struct ExternalMmEncodedEntry {
pub items: Vec<ExternalMmItem>,
pub token_ids: MmTokenIds,
}
impl ExternalMmEncodedEntry {
pub(super) fn validate(&self, input_len: usize) -> Result<(), String> {
for (index, item) in self.items.iter().enumerate() {
let elements = item
.feature
.shape
.iter()
.try_fold(1usize, |size, &dim| size.checked_mul(dim));
if elements != Some(item.feature.data.len()) {
return Err(format!(
"multimodal item {index}: feature shape does not match its data"
));
}
if item
.offsets
.iter()
.any(|&(start, end)| start > end || end as usize >= input_len)
{
return Err(format!(
"multimodal item {index}: token offsets are outside the expanded prompt"
));
}
}
Ok(())
}
}
/// Encoded data parked between a multimodal worker and the scheduler drain.
pub enum MmEncodedEntry {
Qwen(QwenMmEncodedEntry),
External(ExternalMmEncodedEntry),
}
/// Where a result's feature buffers live between worker and drain.
pub enum FeatureStore {
/// In-process; the drain wraps them zero-copy. Single-rank serving, or the
+184 -50
View File
@@ -3,10 +3,12 @@
use std::sync::Arc;
use super::result_store::{FeatureStore, MmEncodedEntry, MmResultStore, park_features_in_shm};
use super::result_store::{
FeatureStore, MmEncodedEntry, MmResultStore, QwenMmEncodedEntry, park_features_in_shm,
};
use crate::message::config::MmSpec;
use crate::message::ids::Rid;
use crate::message::request::MmRequest;
use crate::message::request::{MmRequest, MmWorkItem};
use crate::tokenizer_manager::tokenizer::TextTokenizer;
use crate::tokenizer_manager::wiring::TmEvent;
use crate::utils::runtime::Runnable;
@@ -14,7 +16,7 @@ use crate::utils::runtime::Runnable;
/// Python parity: caller hashes override the computed ones so an external
/// router's keys align with the prefix cache. A length mismatch or malformed
/// entry warns and keeps the computed hash — never blocks the request.
fn apply_caller_hashes(hashes: &mut [u64], caller: &[String]) {
fn apply_caller_hashes<'a>(hashes: impl ExactSizeIterator<Item = &'a mut u64>, caller: &[String]) {
if caller.is_empty() {
return;
}
@@ -26,7 +28,7 @@ fn apply_caller_hashes(hashes: &mut [u64], caller: &[String]) {
);
return;
}
for (hash, entry) in hashes.iter_mut().zip(caller) {
for (hash, entry) in hashes.zip(caller) {
match parse_caller_hash(entry) {
Some(v) => *hash = v,
None => tracing::warn!(%entry, "malformed mm_hashes entry; keeping computed hash"),
@@ -45,16 +47,76 @@ fn parse_caller_hash(entry: &str) -> Option<u64> {
u64::from_str_radix(&hex[hex.len().saturating_sub(16)..], 16).ok()
}
/// Shared state of the mm path, built once at `start_mm_workers`.
/// Complete result of one multimodal processor invocation.
pub struct MmProcessOutput {
pub input_ids: Vec<i32>,
pub result: MmEncodedEntry,
}
/// Multimodal processor shared by built-in and external implementations.
/// Implementations run on the fixed Rust worker pool and must not retain
/// request-scoped Python objects.
pub trait MmProcessor: Send + Sync {
fn process(
&self,
work: MmWorkItem,
tokenizer: Option<&dyn TextTokenizer>,
) -> Result<MmProcessOutput, String>;
}
struct QwenMmProcessor {
family: Box<dyn sglang_mm::pipeline::MmFamilyProcessor>,
feature_shm: bool,
}
impl QwenMmProcessor {
fn new(spec: MmSpec) -> Result<Self, String> {
Ok(Self {
family: sglang_mm::registry::build_pipeline(spec.pipeline)?,
feature_shm: spec.feature_shm,
})
}
}
impl MmProcessor for QwenMmProcessor {
fn process(
&self,
work: MmWorkItem,
tokenizer: Option<&dyn TextTokenizer>,
) -> Result<MmProcessOutput, String> {
let input = super::payload::to_mm_input(work)?;
let output = sglang_mm::driver::process(self.family.as_ref(), input, |text| {
let tokenizer = tokenizer.ok_or_else(|| {
"skip_tokenizer_init is set: multimodal text prompts require input_ids".to_string()
})?;
tokenizer.encode(text).map_err(|error| error.to_string())
})?;
let drain = sglang_mm::qwen_vl::pack_output(output)?;
let features = if self.feature_shm {
park_features_in_shm(&drain.features, &drain.grids)
} else {
FeatureStore::Inline(drain.features)
};
Ok(MmProcessOutput {
input_ids: drain.input_ids,
result: MmEncodedEntry::Qwen(QwenMmEncodedEntry {
features,
grids: drain.grids,
hashes: drain.hashes,
offsets: drain.offsets,
mrope: drain.mrope,
mrope_delta: drain.mrope_delta,
}),
})
}
}
/// Shared state of the multimodal path, built once at worker startup.
pub struct MmContext {
pub family: Box<dyn sglang_mm::pipeline::MmFamilyProcessor>,
pub processor: Arc<dyn MmProcessor>,
/// `None` under `skip_tokenizer_init` (requests must carry `input_ids`).
pub tokenizer: Option<Arc<dyn TextTokenizer>>,
pub results: MmResultStore,
/// Park feature buffers in POSIX shm. Set by the Python launcher
/// (`RustMmProcessor._use_feature_shm`) exactly when the scheduler broadcasts
/// across TP ranks and will unwrap `ShmPointerMMData`.
pub feature_shm: bool,
}
impl MmContext {
@@ -64,51 +126,42 @@ impl MmContext {
results: MmResultStore,
) -> Result<Self, String> {
Ok(Self {
family: sglang_mm::registry::build_pipeline(spec.pipeline)?,
processor: Arc::new(QwenMmProcessor::new(spec)?),
tokenizer,
results,
feature_shm: spec.feature_shm,
})
}
pub fn with_processor(
processor: Arc<dyn MmProcessor>,
tokenizer: Option<Arc<dyn TextTokenizer>>,
results: MmResultStore,
) -> Self {
Self {
processor,
tokenizer,
results,
}
}
}
/// Run the pipeline for one request. `Ok` returns the final expanded ids, the
/// buffers already parked; `Err` rejects the request back to the client.
fn process(
ctx: &MmContext,
rid: &Rid,
mut work: crate::message::request::MmWorkItem,
) -> Result<Vec<i32>, String> {
fn process(ctx: &MmContext, rid: &Rid, mut work: MmWorkItem) -> Result<Vec<i32>, String> {
let caller_hashes = std::mem::take(&mut work.mm_hashes);
let input = super::payload::to_mm_input(work)?;
let output = sglang_mm::driver::process(ctx.family.as_ref(), input, |text| {
let tokenizer = ctx.tokenizer.as_ref().ok_or_else(|| {
"skip_tokenizer_init is set: multimodal text prompts require input_ids".to_string()
})?;
tokenizer.encode(text).map_err(|error| error.to_string())
})?;
// TODO(mm-families): the one family-specific call in this worker — dispatch
// on the spec's `family` (as `registry::build_pipeline` does) once a
// second family lands.
let mut packed = sglang_mm::qwen_vl::pack_output(output)?;
apply_caller_hashes(&mut packed.hashes, &caller_hashes);
let features = if ctx.feature_shm {
park_features_in_shm(&packed.features, &packed.grids)
} else {
FeatureStore::Inline(packed.features)
};
ctx.results.park(
rid.as_str().to_owned(),
MmEncodedEntry {
features,
grids: packed.grids,
hashes: packed.hashes,
offsets: packed.offsets,
mrope: packed.mrope,
mrope_delta: packed.mrope_delta,
},
);
Ok(packed.input_ids)
let mut output = ctx.processor.process(work, ctx.tokenizer.as_deref())?;
match &mut output.result {
MmEncodedEntry::Qwen(entry) => apply_caller_hashes(entry.hashes.iter_mut(), &caller_hashes),
MmEncodedEntry::External(entry) => {
entry.validate(output.input_ids.len())?;
apply_caller_hashes(
entry.items.iter_mut().map(|item| &mut item.hash),
&caller_hashes,
);
}
}
ctx.results.park(rid.as_str().to_owned(), output.result);
Ok(output.input_ids)
}
/// Boot-time wiring of the MM path, held privately by the `Runtime` for the
@@ -171,19 +224,100 @@ impl Runnable for MmWorker {
#[cfg(test)]
mod tests {
use super::*;
use crate::{
ExternalMmEncodedEntry, ExternalMmItem, MmModality, MmTokenIds, Tensor, TensorData,
};
struct ExternalProcessor {
shape: Vec<usize>,
offsets: Vec<(u32, u32)>,
}
impl MmProcessor for ExternalProcessor {
fn process(
&self,
work: MmWorkItem,
tokenizer: Option<&dyn TextTokenizer>,
) -> Result<MmProcessOutput, String> {
assert!(tokenizer.is_none());
Ok(MmProcessOutput {
input_ids: work.input_ids.unwrap_or_default(),
result: MmEncodedEntry::External(ExternalMmEncodedEntry {
items: vec![ExternalMmItem {
modality: MmModality::Image,
feature: Tensor {
shape: self.shape.clone(),
data: TensorData::F32(vec![1.0]),
},
hash: 7,
offsets: self.offsets.clone(),
model_specific_data: Default::default(),
}],
token_ids: MmTokenIds::default(),
}),
})
}
}
#[test]
fn external_processor_result_reaches_store() {
let results = MmResultStore::default();
let processor = ExternalProcessor {
shape: vec![1],
offsets: vec![(1, 1)],
};
let ctx = MmContext::with_processor(Arc::new(processor), None, results.clone());
let rid = Rid::from_client("external");
let work = MmWorkItem {
input_ids: Some(vec![1, 2]),
mm_hashes: vec!["2a".to_owned()],
..Default::default()
};
assert_eq!(process(&ctx, &rid, work).unwrap(), [1, 2]);
let Some(MmEncodedEntry::External(entry)) = results.take(rid.as_str()) else {
panic!("external processor must park an external entry")
};
assert_eq!(entry.items.len(), 1);
assert_eq!(entry.items[0].hash, 0x2a);
}
#[test]
fn malformed_processor_results_are_rejected_before_parking() {
for (shape, offsets) in [
(vec![2], vec![(1, 1)]),
(vec![usize::MAX, 2], vec![(1, 1)]),
(vec![1], vec![(2, 1)]),
(vec![1], vec![(1, 2)]),
] {
let results = MmResultStore::default();
let processor = ExternalProcessor { shape, offsets };
let ctx = MmContext::with_processor(Arc::new(processor), None, results.clone());
let rid = Rid::from_client("invalid");
let work = MmWorkItem {
input_ids: Some(vec![1, 2]),
..Default::default()
};
assert!(process(&ctx, &rid, work).is_err());
assert!(results.take(rid.as_str()).is_none());
}
}
/// Caller hashes override computed ones; mismatched lengths and malformed
/// entries fall back per item, never reject (Python parity).
#[test]
fn caller_hashes_override_with_fallback() {
let mut hashes = vec![1, 2, 3];
apply_caller_hashes(&mut hashes, &[]);
apply_caller_hashes(hashes.iter_mut(), &[]);
assert_eq!(hashes, [1, 2, 3]);
apply_caller_hashes(&mut hashes, &["ff".into()]); // length mismatch
apply_caller_hashes(hashes.iter_mut(), &["ff".into()]); // length mismatch
assert_eq!(hashes, [1, 2, 3]);
apply_caller_hashes(&mut hashes, &["ff".into(), "not-hex".into(), "0x10".into()]);
apply_caller_hashes(
hashes.iter_mut(),
&["ff".into(), "not-hex".into(), "0x10".into()],
);
assert_eq!(hashes, [0xff, 2, 0x10]);
}
@@ -193,7 +327,7 @@ mod tests {
fn caller_hashes_accept_arbitrary_width() {
let sha256 = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
let mut hashes = vec![1];
apply_caller_hashes(&mut hashes, &[sha256.into()]);
apply_caller_hashes(hashes.iter_mut(), &[sha256.into()]);
assert_eq!(hashes, [0xa495991b7852b855]);
assert_eq!(hashes[0] % (1 << 30), 944_945_237); // int(sha256, 16) % (1 << 30)
@@ -16,7 +16,7 @@ use crate::tokenizer_manager::channel::FromSchedulerRx;
use crate::tokenizer_manager::wiring::{Senders, recv};
/// A monotonic counter bumped once per from_scheduler frame the dispatcher drains.
/// It's the rust-native equivalent of the Python `TokenizerManager`'s
/// Equivalent to the Python `TokenizerManager`'s
/// `last_receive_tstamp`: `/health_generate` watches it advance to confirm the
/// scheduler → detok path is alive (the value itself is meaningless).
pub type ActivityCounter = Arc<AtomicU64>;
@@ -795,14 +795,16 @@ fn abort_cancels_parked_mm_request() {
// The worker parks its result, as it always does before MmEncoded.
intake.mm.results.park(
"mm-gone".into(),
crate::multi_modality::result_store::MmEncodedEntry {
features: crate::multi_modality::result_store::FeatureStore::Inline(vec![]),
grids: vec![],
hashes: vec![],
offsets: vec![],
mrope: vec![],
mrope_delta: 0,
},
crate::multi_modality::result_store::MmEncodedEntry::Qwen(
crate::multi_modality::result_store::QwenMmEncodedEntry {
features: crate::multi_modality::result_store::FeatureStore::Inline(vec![]),
grids: vec![],
hashes: vec![],
offsets: vec![],
mrope: vec![],
mrope_delta: 0,
},
),
);
intake.on_abort(AbortSource::Guard("mm-gone".to_string().into()));
assert_eq!(consumer.drain(16).headers.len(), 1, "only the AbortReq");
@@ -158,7 +158,7 @@ fn strip_auto_specials(mut ids: Vec<i32>, auto_specials: &[i32]) -> Vec<i32> {
/// The `auto_specials` prefix (probed once at construction, Python's
/// `encode("")` probe) is stripped from template-rendered prompts —
/// [`GenerateRequest`]'s `skip_special_tokens` — so chat prompts gain no
/// extra BOS/EOS while native text keeps the post-processor specials.
/// extra BOS/EOS while plain text keeps the post-processor specials.
pub struct TokenizerWorker {
rx: flume::Receiver<Request>,
tm: flume::Sender<TmEvent>,
@@ -347,7 +347,7 @@ mod tests {
};
g.input_ids.clone().expect("tokenized")
};
assert_eq!(run(false), vec![0, 2], "native prompts keep specials");
assert_eq!(run(false), vec![0, 2], "plain text prompts keep specials");
assert_eq!(run(true), vec![2], "rendered prompts lose the auto BOS");
}
}
+1 -1
View File
@@ -39,7 +39,7 @@ pub enum RequestState {
/// Outcome of validation.
#[derive(Debug, Clone, Copy)]
pub enum ValidationOutcome {
/// Has multimodal inputs → Encoding, where an MM worker runs the native
/// Has multimodal inputs → Encoding, where an MM worker runs the multimodal
/// pipeline and returns the final expanded `input_ids`.
HasMultimodal,
/// Plain text → Tokenizing.
+18 -1
View File
@@ -76,6 +76,24 @@ impl Runtime {
self.mm_wiring.tokenizer.clone(),
self.mm_results.clone(),
)?);
self.spawn_mm_pool(workers, ctx);
Ok(())
}
pub fn start_mm_workers_with_processor(
&self,
processor: Arc<dyn crate::multi_modality::worker::MmProcessor>,
workers: usize,
) {
let ctx = Arc::new(crate::multi_modality::worker::MmContext::with_processor(
processor,
self.mm_wiring.tokenizer.clone(),
self.mm_results.clone(),
));
self.spawn_mm_pool(workers, ctx);
}
fn spawn_mm_pool(&self, workers: usize, ctx: Arc<crate::multi_modality::worker::MmContext>) {
let mut threads = self.threads.lock().unwrap();
spawn_pool("mm-worker", None, workers.max(1), &mut threads, |_| {
crate::multi_modality::worker::MmWorker::new(
@@ -84,7 +102,6 @@ impl Runtime {
ctx.clone(),
)
});
Ok(())
}
/// Stop the runtime and join every worker thread (with a bounded wait).
@@ -0,0 +1,82 @@
"""Model packages use one extension for server arguments and worker startup."""
import unittest
from types import ModuleType, SimpleNamespace
from unittest.mock import Mock, patch, sentinel
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.rust_server import server as server_module # noqa: E402
from sglang.srt.rust_server.server import RustServer # noqa: E402
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestRustServerExtension(CustomTestCase):
def test_launch_uses_the_model_extension_and_instance_worker_state(self):
extension = ModuleType("model_server")
extension.Server = Mock()
class ModelServer(RustServer):
@classmethod
def _load_extension(cls):
return extension
def _start_multimodal(self, scheduler):
self.server.start_mm_workers(sentinel.spec, 8)
scheduler = SimpleNamespace(
ps=SimpleNamespace(
dp_size=2,
attn_dp_rank=1,
tp_size=2,
tp_rank=1,
pp_size=1,
attn_tp_size=1,
attn_cp_size=1,
),
model_config=SimpleNamespace(is_multimodal=True),
)
with (
patch.object(
server_module,
"get_exec",
return_value=SimpleNamespace(
moe=SimpleNamespace(is_ep_scale_joiner=False)
),
),
patch.object(
server_module, "get_parallel", return_value=SimpleNamespace(nnodes=1)
),
patch.object(ModelServer, "_partition_cores", return_value=(None, None)),
patch.object(
server_module,
"get_mm",
return_value=SimpleNamespace(mm_processor_worker_num=8),
),
patch.object(
server_module,
"get_serving",
return_value=SimpleNamespace(host="::", port=30000),
),
patch.object(
server_module, "_build_server_args", return_value=sentinel.args
) as build_args,
):
instance = ModelServer.launch(scheduler)
build_args.assert_called_once_with(scheduler, extension=extension)
extension.Server.assert_called_once_with(
sentinel.args, cores=None, port_offset=1
)
instance.server.start_mm_workers.assert_called_once_with(sentinel.spec, 8)
self.assertIsInstance(instance, ModelServer)
self.assertEqual(instance.http_port, 30001)
self.assertTrue(instance._multimodal_enabled)
if __name__ == "__main__":
unittest.main()