Add CUDA VMM multimodal feature transport (#33899)

Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
Oguz Ulgen
2026-08-07 13:39:54 -07:00
committed by GitHub
co-authored by Yinghai Lu
parent 3c51e29deb
commit 7f6b4cb94b
10 changed files with 2490 additions and 78 deletions
+44 -34
View File
@@ -1191,26 +1191,32 @@ class Engine(EngineScoreMixin, EngineBase):
tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
template_manager = None
# Wait for the model to finish loading
scheduler_init_result.wait_for_ready()
startup_complete = False
try:
# Wait for the model to finish loading
scheduler_init_result.wait_for_ready()
cls._set_startup_time(tokenizer_manager, scheduler_init_result, startup_tic)
cls._set_startup_time(tokenizer_manager, scheduler_init_result, startup_tic)
# Get back some info from scheduler to tokenizer_manager
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[0][
"max_req_input_len"
]
# Get back some info from scheduler to tokenizer_manager
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[
0
]["max_req_input_len"]
# Set up subprocess liveness watchdog to detect crashes
# Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
processes = list(scheduler_procs or [])
names = [f"scheduler_{i}" for i in range(len(processes))]
processes.extend(detoken_procs)
names.extend(detoken_names)
subprocess_watchdog = SubprocessWatchdog(
processes=processes, process_names=names
)
subprocess_watchdog.start()
# Set up subprocess liveness watchdog to detect crashes
# Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
processes = list(scheduler_procs or [])
names = [f"scheduler_{i}" for i in range(len(processes))]
processes.extend(detoken_procs)
names.extend(detoken_names)
subprocess_watchdog = SubprocessWatchdog(
processes=processes, process_names=names
)
subprocess_watchdog.start()
startup_complete = True
finally:
if not startup_complete and isinstance(tokenizer_manager, TokenizerManager):
tokenizer_manager.cuda_vmm_feature_transport.shutdown()
return (
tokenizer_manager,
@@ -1225,26 +1231,30 @@ class Engine(EngineScoreMixin, EngineBase):
"""Shutdown the engine; block until the scheduler subprocess releases
its GPU context so the caller can immediately reallocate on the same
device."""
if (
self.tokenizer_manager is not None
and self.tokenizer_manager._subprocess_watchdog is not None
):
self.tokenizer_manager._subprocess_watchdog.stop()
try:
if (
self.tokenizer_manager is not None
and self.tokenizer_manager._subprocess_watchdog is not None
):
self.tokenizer_manager._subprocess_watchdog.stop()
send_to_rpc = getattr(self, "send_to_rpc", None)
if send_to_rpc is not None:
send_to_rpc.close(linger=0)
self.send_to_rpc = None
send_to_rpc = getattr(self, "send_to_rpc", None)
if send_to_rpc is not None:
send_to_rpc.close(linger=0)
self.send_to_rpc = None
# Gracefully stop weight cache daemons *before* the blanket
# kill_process_tree below, so their SIGTERM handlers can unlink the
# .sock/.ready files instead of being SIGKILLed and leaving stale state.
daemon_procs = getattr(self, "_weight_cache_daemon_procs", None)
if daemon_procs:
self._terminate_weight_cache_daemons(daemon_procs)
self._weight_cache_daemon_procs = []
# Gracefully stop weight cache daemons *before* the blanket
# kill_process_tree below, so their SIGTERM handlers can unlink the
# .sock/.ready files instead of being SIGKILLed and leaving stale state.
daemon_procs = getattr(self, "_weight_cache_daemon_procs", None)
if daemon_procs:
self._terminate_weight_cache_daemons(daemon_procs)
self._weight_cache_daemon_procs = []
kill_process_tree(os.getpid(), include_parent=False, wait_timeout=60)
kill_process_tree(os.getpid(), include_parent=False, wait_timeout=60)
finally:
if isinstance(self.tokenizer_manager, TokenizerManager):
self.tokenizer_manager.cuda_vmm_feature_transport.shutdown()
def __enter__(self):
return self
+32 -4
View File
@@ -1834,6 +1834,10 @@ class Scheduler(
def process_input_requests(self, recv_reqs: List):
now = time.monotonic()
self.session_controller.maybe_reap(now)
if self.server_args.mm_feature_transport == "cuda_vmm":
for recv_req in recv_reqs:
self._materialize_cuda_vmm_inputs(recv_req)
for recv_req in recv_reqs:
# Skip health check when server is busy — ongoing requests already carry health info.
if is_health_check_generate_req(recv_req) and not self.is_fully_idle(
@@ -1861,6 +1865,28 @@ class Scheduler(
if self.external_corpus_manager is not None:
self.external_corpus_manager.check_pending_load()
def _materialize_cuda_vmm_inputs(self, recv_req):
"""Release VMM slices before request handling can reject the request."""
if isinstance(
recv_req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)
):
tokenized_reqs = (recv_req,)
elif isinstance(
recv_req,
(BatchTokenizedGenerateReqInput, BatchTokenizedEmbeddingReqInput),
):
tokenized_reqs = recv_req
else:
return
for tokenized_req in tokenized_reqs:
if tokenized_req.mm_inputs is not None and not isinstance(
tokenized_req.mm_inputs, MultimodalInputs
):
tokenized_req.mm_inputs = MultimodalInputs.from_processor_output(
tokenized_req.mm_inputs
)
def init_profiler(self) -> None:
self.profiler_manager = SchedulerProfilerManager(
ps=self.ps,
@@ -2210,11 +2236,13 @@ class Scheduler(
return image_inputs
def _get_multimodal_inputs(self, mm_inputs_dict):
def _get_multimodal_inputs(self, mm_inputs):
if isinstance(mm_inputs, MultimodalInputs):
return mm_inputs
if get_mm().enable_broadcast_mm_inputs_process:
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
else:
return MultimodalInputs.from_processor_output(mm_inputs_dict)
return self._process_and_broadcast_mm_inputs(mm_inputs)
return MultimodalInputs.from_processor_output(mm_inputs)
@staticmethod
def _try_apply_padded_mm_input_ids(recv_req, req, image_inputs) -> bool:
+68 -26
View File
@@ -135,6 +135,7 @@ from sglang.srt.utils import (
kill_process_tree,
)
from sglang.srt.utils.aio_rwlock import RWLock
from sglang.srt.utils.cuda_vmm_transport_utils import CudaVmmFeatureTransport
from sglang.srt.utils.cudacore_pyspy_dump_utils import (
collect_scheduler_processes,
pyspy_dump_schedulers,
@@ -416,6 +417,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Init model config
self.init_model_config()
self._validate_cuda_vmm_feature_transport_support()
# Initialize tokenizer and multimodalprocessor
self.init_tokenizer_and_processor()
@@ -444,6 +446,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Init request dispatcher
self.init_request_dispatcher()
# Construct this last so later initialization failures cannot orphan
# the transport's recycler thread.
self.cuda_vmm_feature_transport = CudaVmmFeatureTransport(
self.server_args, self.mm_processor
)
def init_model_config(self):
server_args = self.server_args
model_config_class = getattr(self, "model_config_class", ModelConfig)
@@ -516,6 +524,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else:
self.async_dynamic_batch_tokenizer = None
def _validate_cuda_vmm_feature_transport_support(self) -> None:
if self.server_args.mm_feature_transport != "cuda_vmm":
return
from sglang.srt.model_loader.utils import get_model_architecture
model_class, _ = get_model_architecture(self.model_config)
if not getattr(model_class, "supports_cuda_vmm_feature_transport", False):
raise ValueError(
"--mm-feature-transport=cuda_vmm is not supported by model class "
f"{model_class.__name__}"
)
def init_ipc_channels(self, port_args: PortArgs):
context = zmq.asyncio.Context(2)
self.recv_from_detokenizer = get_zmq_socket(
@@ -1541,16 +1562,26 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self,
tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
):
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
time_stats = tokenized_obj.time_stats
tokenized_obj.wrap_pickle_fields()
self._dispatch_to_scheduler(tokenized_obj)
state = self.rid_to_state.get(tokenized_obj.rid)
if state is not None:
state.dispatched = True
tokenized_obj.time_stats = time_stats
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
prepared_mm_items = []
dispatched = False
try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
(tokenized_obj.mm_inputs,)
)
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
time_stats = tokenized_obj.time_stats
tokenized_obj.wrap_pickle_fields()
self._dispatch_to_scheduler(tokenized_obj)
dispatched = True
state = self.rid_to_state.get(tokenized_obj.rid)
if state is not None:
state.dispatched = True
tokenized_obj.time_stats = time_stats
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
finally:
if not dispatched:
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
def _send_batch_request(
self,
@@ -1559,24 +1590,35 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
],
):
"""Send a batch of tokenized requests as a single batched request to the scheduler."""
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
for tokenized_obj in tokenized_objs:
tokenized_obj.wrap_pickle_fields()
prepared_mm_items = []
dispatched = False
try:
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
tokenized_obj.mm_inputs for tokenized_obj in tokenized_objs
)
if isinstance(tokenized_objs[0], TokenizedGenerateReqInput):
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
for tokenized_obj in tokenized_objs:
tokenized_obj.wrap_pickle_fields()
self._dispatch_to_scheduler(batch_req)
for tokenized_obj in tokenized_objs:
state = self.rid_to_state.get(tokenized_obj.rid)
if state is not None:
state.dispatched = True
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
tokenized_obj.time_stats = time_stat
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
if isinstance(tokenized_objs[0], TokenizedGenerateReqInput):
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
self._dispatch_to_scheduler(batch_req)
dispatched = True
for tokenized_obj in tokenized_objs:
state = self.rid_to_state.get(tokenized_obj.rid)
if state is not None:
state.dispatched = True
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
tokenized_obj.time_stats = time_stat
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
finally:
if not dispatched:
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
def _coalesce_streaming_chunks(
self,
@@ -200,7 +200,7 @@ class BaseMultimodalProcessor(ABC):
)
self.mm_feature_transport = (
configured_mm_feature_transport
if configured_mm_feature_transport in ("cpu", "cuda_ipc")
if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm")
else "cpu"
)
self.use_cuda_ipc = self.mm_feature_transport == "cuda_ipc"
@@ -289,8 +289,11 @@ class BaseMultimodalProcessor(ABC):
self.mm_processor_worker_num,
"auto" if requested_mm_processor_worker_num == 0 else "explicit",
)
cpu_worker_start_method = (
"spawn" if self.mm_feature_transport == "cuda_vmm" else "fork"
)
self.cpu_executor = concurrent.futures.ProcessPoolExecutor(
mp_context=mp.get_context("fork"),
mp_context=mp.get_context(cpu_worker_start_method),
max_workers=int(os.environ.get("SGLANG_CPU_WORKERS", os.cpu_count())),
)
@@ -363,6 +366,10 @@ class BaseMultimodalProcessor(ABC):
self.server_args.base_gpu_id,
)
@property
def keep_mm_features_on_device(self) -> bool:
return self.mm_feature_transport in ("cuda_ipc", "cuda_vmm")
def compute_mrope_positions(self, input_ids, mm_items):
"""Compute M-RoPE positions from expanded input_ids and multimodal items.
@@ -588,7 +595,10 @@ class BaseMultimodalProcessor(ABC):
)
# Deferred: the hash is computed on the GPU tensor first, and
# _precompute_hashes_before_cpu_transfer moves it down afterwards.
if not self.use_cuda_ipc and not self.precompute_hash_before_cpu_transfer:
if (
not self.keep_mm_features_on_device
and not self.precompute_hash_before_cpu_transfer
):
# move feature tensors to cpu
for feature_name in self.FEATURE_NAMES:
if feature_name in result and isinstance(
@@ -1395,7 +1405,7 @@ class BaseMultimodalProcessor(ABC):
for item in mm_items:
item.set_pad_value()
if not self.use_cuda_ipc:
if not self.keep_mm_features_on_device:
item.feature = self._move_feature_to_cpu(item.feature)
item.precomputed_embeddings = self._move_feature_to_cpu(
item.precomputed_embeddings
+29 -10
View File
@@ -2746,13 +2746,15 @@ class ServerArgs:
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
] = False
mm_feature_transport: A[
Optional[Literal["cpu", "cuda_ipc"]],
"Transport multimodal features through CPU memory or a bounded CUDA IPC pool. "
Optional[Literal["cpu", "cuda_ipc", "cuda_vmm"]],
"Transport multimodal features through CPU memory, a bounded CUDA IPC "
"pool, or a bounded CUDA VMM pool. CUDA VMM must be selected explicitly "
"and is available only to models that opt in. "
"Unset resolves automatically: multimodal models on single-node CUDA "
"deployments (without disaggregation) use cuda_ipc, everything else uses "
"cpu. CUDA IPC reserves SGLANG_MM_FEATURE_CACHE_MB (default 1024 MiB) on "
"the base GPU and falls back to CPU transport per tensor when the pool is "
"full.",
"cpu. Both CUDA transports reserve SGLANG_MM_FEATURE_CACHE_MB (default "
"1024 MiB) on the base GPU across tokenizer workers and fall back to CPU "
"transport per tensor when full.",
NS("mm"),
] = None
keep_mm_feature_on_device: A[
@@ -7582,10 +7584,10 @@ class ServerArgs:
legacy_ipc_enabled = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
if self.keep_mm_feature_on_device:
if requested_transport == "cpu":
if requested_transport not in (None, "cuda_ipc"):
raise ValueError(
"--keep-mm-feature-on-device conflicts with "
"--mm-feature-transport=cpu. Use only "
f"--mm-feature-transport={requested_transport}. Use only "
"--mm-feature-transport=cuda_ipc."
)
requested_transport = "cuda_ipc"
@@ -7638,14 +7640,31 @@ class ServerArgs:
int(legacy_ipc_enabled),
)
if self.encoder_only and requested_transport == "cuda_ipc":
if self.encoder_only and requested_transport in ("cuda_ipc", "cuda_vmm"):
logger.warning(
"--mm-feature-transport=cuda_ipc does not control encoder-only "
"--mm-feature-transport=%s does not control encoder-only "
"output transfer; using cpu for this inactive transport. Select "
"--encoder-transfer-backend for encoder outputs."
"--encoder-transfer-backend for encoder outputs.",
requested_transport,
)
requested_transport = "cpu"
if requested_transport == "cuda_vmm":
if not is_cuda():
raise ValueError(
"--mm-feature-transport=cuda_vmm requires NVIDIA CUDA."
)
if self.pp_size != 1:
raise ValueError(
"--mm-feature-transport=cuda_vmm does not support pipeline "
"parallelism."
)
if envs.SGLANG_RUST_SERVER.get():
raise ValueError(
"--mm-feature-transport=cuda_vmm is not supported with "
"SGLANG_RUST_SERVER."
)
if requested_transport == "cuda_ipc":
if not is_cuda():
raise ValueError(
File diff suppressed because it is too large Load Diff