Add CUDA VMM multimodal feature transport (#33899)
Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user