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)
|
tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
|
||||||
template_manager = None
|
template_manager = None
|
||||||
|
|
||||||
# Wait for the model to finish loading
|
startup_complete = False
|
||||||
scheduler_init_result.wait_for_ready()
|
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
|
# Get back some info from scheduler to tokenizer_manager
|
||||||
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[0][
|
tokenizer_manager.max_req_input_len = scheduler_init_result.scheduler_infos[
|
||||||
"max_req_input_len"
|
0
|
||||||
]
|
]["max_req_input_len"]
|
||||||
|
|
||||||
# Set up subprocess liveness watchdog to detect crashes
|
# Set up subprocess liveness watchdog to detect crashes
|
||||||
# Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
|
# Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
|
||||||
processes = list(scheduler_procs or [])
|
processes = list(scheduler_procs or [])
|
||||||
names = [f"scheduler_{i}" for i in range(len(processes))]
|
names = [f"scheduler_{i}" for i in range(len(processes))]
|
||||||
processes.extend(detoken_procs)
|
processes.extend(detoken_procs)
|
||||||
names.extend(detoken_names)
|
names.extend(detoken_names)
|
||||||
subprocess_watchdog = SubprocessWatchdog(
|
subprocess_watchdog = SubprocessWatchdog(
|
||||||
processes=processes, process_names=names
|
processes=processes, process_names=names
|
||||||
)
|
)
|
||||||
subprocess_watchdog.start()
|
subprocess_watchdog.start()
|
||||||
|
startup_complete = True
|
||||||
|
finally:
|
||||||
|
if not startup_complete and isinstance(tokenizer_manager, TokenizerManager):
|
||||||
|
tokenizer_manager.cuda_vmm_feature_transport.shutdown()
|
||||||
|
|
||||||
return (
|
return (
|
||||||
tokenizer_manager,
|
tokenizer_manager,
|
||||||
@@ -1225,26 +1231,30 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
"""Shutdown the engine; block until the scheduler subprocess releases
|
"""Shutdown the engine; block until the scheduler subprocess releases
|
||||||
its GPU context so the caller can immediately reallocate on the same
|
its GPU context so the caller can immediately reallocate on the same
|
||||||
device."""
|
device."""
|
||||||
if (
|
try:
|
||||||
self.tokenizer_manager is not None
|
if (
|
||||||
and self.tokenizer_manager._subprocess_watchdog is not None
|
self.tokenizer_manager is not None
|
||||||
):
|
and self.tokenizer_manager._subprocess_watchdog is not None
|
||||||
self.tokenizer_manager._subprocess_watchdog.stop()
|
):
|
||||||
|
self.tokenizer_manager._subprocess_watchdog.stop()
|
||||||
|
|
||||||
send_to_rpc = getattr(self, "send_to_rpc", None)
|
send_to_rpc = getattr(self, "send_to_rpc", None)
|
||||||
if send_to_rpc is not None:
|
if send_to_rpc is not None:
|
||||||
send_to_rpc.close(linger=0)
|
send_to_rpc.close(linger=0)
|
||||||
self.send_to_rpc = None
|
self.send_to_rpc = None
|
||||||
|
|
||||||
# Gracefully stop weight cache daemons *before* the blanket
|
# Gracefully stop weight cache daemons *before* the blanket
|
||||||
# kill_process_tree below, so their SIGTERM handlers can unlink the
|
# kill_process_tree below, so their SIGTERM handlers can unlink the
|
||||||
# .sock/.ready files instead of being SIGKILLed and leaving stale state.
|
# .sock/.ready files instead of being SIGKILLed and leaving stale state.
|
||||||
daemon_procs = getattr(self, "_weight_cache_daemon_procs", None)
|
daemon_procs = getattr(self, "_weight_cache_daemon_procs", None)
|
||||||
if daemon_procs:
|
if daemon_procs:
|
||||||
self._terminate_weight_cache_daemons(daemon_procs)
|
self._terminate_weight_cache_daemons(daemon_procs)
|
||||||
self._weight_cache_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):
|
def __enter__(self):
|
||||||
return self
|
return self
|
||||||
|
|||||||
@@ -1834,6 +1834,10 @@ class Scheduler(
|
|||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
self.session_controller.maybe_reap(now)
|
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:
|
for recv_req in recv_reqs:
|
||||||
# Skip health check when server is busy — ongoing requests already carry health info.
|
# 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(
|
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:
|
if self.external_corpus_manager is not None:
|
||||||
self.external_corpus_manager.check_pending_load()
|
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:
|
def init_profiler(self) -> None:
|
||||||
self.profiler_manager = SchedulerProfilerManager(
|
self.profiler_manager = SchedulerProfilerManager(
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
@@ -2210,11 +2236,13 @@ class Scheduler(
|
|||||||
|
|
||||||
return image_inputs
|
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:
|
if get_mm().enable_broadcast_mm_inputs_process:
|
||||||
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
|
return self._process_and_broadcast_mm_inputs(mm_inputs)
|
||||||
else:
|
return MultimodalInputs.from_processor_output(mm_inputs)
|
||||||
return MultimodalInputs.from_processor_output(mm_inputs_dict)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _try_apply_padded_mm_input_ids(recv_req, req, image_inputs) -> bool:
|
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,
|
kill_process_tree,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.aio_rwlock import RWLock
|
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 (
|
from sglang.srt.utils.cudacore_pyspy_dump_utils import (
|
||||||
collect_scheduler_processes,
|
collect_scheduler_processes,
|
||||||
pyspy_dump_schedulers,
|
pyspy_dump_schedulers,
|
||||||
@@ -416,6 +417,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
|
|
||||||
# Init model config
|
# Init model config
|
||||||
self.init_model_config()
|
self.init_model_config()
|
||||||
|
self._validate_cuda_vmm_feature_transport_support()
|
||||||
|
|
||||||
# Initialize tokenizer and multimodalprocessor
|
# Initialize tokenizer and multimodalprocessor
|
||||||
self.init_tokenizer_and_processor()
|
self.init_tokenizer_and_processor()
|
||||||
@@ -444,6 +446,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Init request dispatcher
|
# Init request dispatcher
|
||||||
self.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):
|
def init_model_config(self):
|
||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
model_config_class = getattr(self, "model_config_class", ModelConfig)
|
model_config_class = getattr(self, "model_config_class", ModelConfig)
|
||||||
@@ -516,6 +524,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
else:
|
else:
|
||||||
self.async_dynamic_batch_tokenizer = None
|
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):
|
def init_ipc_channels(self, port_args: PortArgs):
|
||||||
context = zmq.asyncio.Context(2)
|
context = zmq.asyncio.Context(2)
|
||||||
self.recv_from_detokenizer = get_zmq_socket(
|
self.recv_from_detokenizer = get_zmq_socket(
|
||||||
@@ -1541,16 +1562,26 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self,
|
self,
|
||||||
tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
|
tokenized_obj: Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput],
|
||||||
):
|
):
|
||||||
tokenized_obj.time_stats.set_api_server_dispatch_time()
|
prepared_mm_items = []
|
||||||
tokenized_obj = wrap_shm_features(tokenized_obj)
|
dispatched = False
|
||||||
time_stats = tokenized_obj.time_stats
|
try:
|
||||||
tokenized_obj.wrap_pickle_fields()
|
prepared_mm_items = self.cuda_vmm_feature_transport.prepare_for_dispatch(
|
||||||
self._dispatch_to_scheduler(tokenized_obj)
|
(tokenized_obj.mm_inputs,)
|
||||||
state = self.rid_to_state.get(tokenized_obj.rid)
|
)
|
||||||
if state is not None:
|
tokenized_obj.time_stats.set_api_server_dispatch_time()
|
||||||
state.dispatched = True
|
tokenized_obj = wrap_shm_features(tokenized_obj)
|
||||||
tokenized_obj.time_stats = time_stats
|
time_stats = tokenized_obj.time_stats
|
||||||
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
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(
|
def _send_batch_request(
|
||||||
self,
|
self,
|
||||||
@@ -1559,24 +1590,35 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
],
|
],
|
||||||
):
|
):
|
||||||
"""Send a batch of tokenized requests as a single batched request to the scheduler."""
|
"""Send a batch of tokenized requests as a single batched request to the scheduler."""
|
||||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
|
prepared_mm_items = []
|
||||||
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
|
dispatched = False
|
||||||
for tokenized_obj in tokenized_objs:
|
try:
|
||||||
tokenized_obj.wrap_pickle_fields()
|
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):
|
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
|
||||||
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
|
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
|
||||||
else:
|
for tokenized_obj in tokenized_objs:
|
||||||
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
tokenized_obj.wrap_pickle_fields()
|
||||||
|
|
||||||
self._dispatch_to_scheduler(batch_req)
|
if isinstance(tokenized_objs[0], TokenizedGenerateReqInput):
|
||||||
for tokenized_obj in tokenized_objs:
|
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
|
||||||
state = self.rid_to_state.get(tokenized_obj.rid)
|
else:
|
||||||
if state is not None:
|
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
||||||
state.dispatched = True
|
|
||||||
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
|
self._dispatch_to_scheduler(batch_req)
|
||||||
tokenized_obj.time_stats = time_stat
|
dispatched = True
|
||||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
|
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(
|
def _coalesce_streaming_chunks(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -200,7 +200,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
)
|
)
|
||||||
self.mm_feature_transport = (
|
self.mm_feature_transport = (
|
||||||
configured_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"
|
else "cpu"
|
||||||
)
|
)
|
||||||
self.use_cuda_ipc = self.mm_feature_transport == "cuda_ipc"
|
self.use_cuda_ipc = self.mm_feature_transport == "cuda_ipc"
|
||||||
@@ -289,8 +289,11 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.mm_processor_worker_num,
|
self.mm_processor_worker_num,
|
||||||
"auto" if requested_mm_processor_worker_num == 0 else "explicit",
|
"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(
|
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())),
|
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,
|
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):
|
def compute_mrope_positions(self, input_ids, mm_items):
|
||||||
"""Compute M-RoPE positions from expanded input_ids and multimodal 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
|
# Deferred: the hash is computed on the GPU tensor first, and
|
||||||
# _precompute_hashes_before_cpu_transfer moves it down afterwards.
|
# _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
|
# move feature tensors to cpu
|
||||||
for feature_name in self.FEATURE_NAMES:
|
for feature_name in self.FEATURE_NAMES:
|
||||||
if feature_name in result and isinstance(
|
if feature_name in result and isinstance(
|
||||||
@@ -1395,7 +1405,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
|
|
||||||
for item in mm_items:
|
for item in mm_items:
|
||||||
item.set_pad_value()
|
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.feature = self._move_feature_to_cpu(item.feature)
|
||||||
item.precomputed_embeddings = self._move_feature_to_cpu(
|
item.precomputed_embeddings = self._move_feature_to_cpu(
|
||||||
item.precomputed_embeddings
|
item.precomputed_embeddings
|
||||||
|
|||||||
@@ -2746,13 +2746,15 @@ class ServerArgs:
|
|||||||
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
|
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
|
||||||
] = False
|
] = False
|
||||||
mm_feature_transport: A[
|
mm_feature_transport: A[
|
||||||
Optional[Literal["cpu", "cuda_ipc"]],
|
Optional[Literal["cpu", "cuda_ipc", "cuda_vmm"]],
|
||||||
"Transport multimodal features through CPU memory or a bounded CUDA IPC pool. "
|
"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 "
|
"Unset resolves automatically: multimodal models on single-node CUDA "
|
||||||
"deployments (without disaggregation) use cuda_ipc, everything else uses "
|
"deployments (without disaggregation) use cuda_ipc, everything else uses "
|
||||||
"cpu. CUDA IPC reserves SGLANG_MM_FEATURE_CACHE_MB (default 1024 MiB) on "
|
"cpu. Both CUDA transports reserve SGLANG_MM_FEATURE_CACHE_MB (default "
|
||||||
"the base GPU and falls back to CPU transport per tensor when the pool is "
|
"1024 MiB) on the base GPU across tokenizer workers and fall back to CPU "
|
||||||
"full.",
|
"transport per tensor when full.",
|
||||||
NS("mm"),
|
NS("mm"),
|
||||||
] = None
|
] = None
|
||||||
keep_mm_feature_on_device: A[
|
keep_mm_feature_on_device: A[
|
||||||
@@ -7582,10 +7584,10 @@ class ServerArgs:
|
|||||||
legacy_ipc_enabled = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
|
legacy_ipc_enabled = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
|
||||||
|
|
||||||
if self.keep_mm_feature_on_device:
|
if self.keep_mm_feature_on_device:
|
||||||
if requested_transport == "cpu":
|
if requested_transport not in (None, "cuda_ipc"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--keep-mm-feature-on-device conflicts with "
|
"--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."
|
"--mm-feature-transport=cuda_ipc."
|
||||||
)
|
)
|
||||||
requested_transport = "cuda_ipc"
|
requested_transport = "cuda_ipc"
|
||||||
@@ -7638,14 +7640,31 @@ class ServerArgs:
|
|||||||
int(legacy_ipc_enabled),
|
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(
|
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 "
|
"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"
|
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 requested_transport == "cuda_ipc":
|
||||||
if not is_cuda():
|
if not is_cuda():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -237,6 +237,31 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
|
|||||||
self.assertFalse(processor.use_ipc_pool_handle_cache)
|
self.assertFalse(processor.use_ipc_pool_handle_cache)
|
||||||
memory_pool.assert_not_called()
|
memory_pool.assert_not_called()
|
||||||
|
|
||||||
|
def test_cuda_vmm_keeps_features_on_device_without_ipc_pool(self):
|
||||||
|
from sglang.srt.multimodal.processors import base_processor
|
||||||
|
|
||||||
|
hf_processor = self._processor()
|
||||||
|
feature = torch.empty(1, device="meta")
|
||||||
|
hf_processor.return_value = {"pixel_values": feature}
|
||||||
|
with patch.object(
|
||||||
|
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
|
||||||
|
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
|
||||||
|
processor = base_processor.BaseMultimodalProcessor(
|
||||||
|
hf_config=MagicMock(),
|
||||||
|
server_args=self._server_args("cuda_vmm"),
|
||||||
|
_processor=hf_processor,
|
||||||
|
transport_mode=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = processor.process_mm_data("test")
|
||||||
|
|
||||||
|
self.assertEqual(processor.mm_feature_transport, "cuda_vmm")
|
||||||
|
self.assertFalse(processor.use_cuda_ipc)
|
||||||
|
self.assertTrue(processor.keep_mm_features_on_device)
|
||||||
|
self.assertEqual(processor.cpu_executor._mp_context.get_start_method(), "spawn")
|
||||||
|
self.assertIs(result["pixel_values"], feature)
|
||||||
|
memory_pool.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
|
class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -251,6 +276,7 @@ class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
|
|||||||
processor = BaseMultimodalProcessor()
|
processor = BaseMultimodalProcessor()
|
||||||
processor.precompute_hash_before_cpu_transfer = enabled
|
processor.precompute_hash_before_cpu_transfer = enabled
|
||||||
processor.use_cuda_ipc = False
|
processor.use_cuda_ipc = False
|
||||||
|
processor.mm_feature_transport = "cpu"
|
||||||
return processor
|
return processor
|
||||||
|
|
||||||
def test_enabled_path_sets_hash_and_pad_value(self):
|
def test_enabled_path_sets_hash_and_pad_value(self):
|
||||||
|
|||||||
@@ -0,0 +1,490 @@
|
|||||||
|
"""CUDA VMM multimodal feature transport regression tests."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import multiprocessing as mp
|
||||||
|
import os
|
||||||
|
import pickle
|
||||||
|
import queue
|
||||||
|
import threading
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmMemoryPool,
|
||||||
|
CudaVmmPackedTensorTransportProxy,
|
||||||
|
_imported_pool_cache_clear,
|
||||||
|
_PosixFdBroker,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=60, stage="base-c", runner_config="4-gpu-gb300")
|
||||||
|
|
||||||
|
|
||||||
|
class _FabricUnavailableCudaVmmMemoryPool(CudaVmmMemoryPool):
|
||||||
|
def _allocate(self, memory_size: int) -> None:
|
||||||
|
if self.use_fabric:
|
||||||
|
raise RuntimeError("forced FABRIC allocation failure")
|
||||||
|
super()._allocate(memory_size)
|
||||||
|
|
||||||
|
|
||||||
|
def _produce_vmm_tensor(proxy_queue, consumer_done, result_queue, mode):
|
||||||
|
pool = source = proxy = None
|
||||||
|
try:
|
||||||
|
torch.cuda.set_device(0)
|
||||||
|
pool_cls = (
|
||||||
|
_FabricUnavailableCudaVmmMemoryPool
|
||||||
|
if mode == "posix_fallback"
|
||||||
|
else CudaVmmMemoryPool
|
||||||
|
)
|
||||||
|
pool = pool_cls(
|
||||||
|
memory_size=4 << 20,
|
||||||
|
recycle_interval=60,
|
||||||
|
base_gpu_id=0,
|
||||||
|
consumer_count=2,
|
||||||
|
allow_posix_fallback=True,
|
||||||
|
)
|
||||||
|
source = torch.arange(35, dtype=torch.float32, device="cuda").reshape(5, 7)
|
||||||
|
expected = source.cpu().tolist()
|
||||||
|
proxy = pool.wrap_tensor(source)
|
||||||
|
proxy_queue.put((proxy, expected))
|
||||||
|
if not consumer_done.wait(timeout=60):
|
||||||
|
raise TimeoutError("consumers did not release the CUDA VMM tensor")
|
||||||
|
with pool._lock:
|
||||||
|
pool._recycle_chunks()
|
||||||
|
pool._merge_chunks()
|
||||||
|
if pool.occupied_chunks:
|
||||||
|
raise RuntimeError(
|
||||||
|
"consumer acknowledgements did not recycle the slice"
|
||||||
|
)
|
||||||
|
except Exception as exc: # noqa: BLE001 # pragma: no cover
|
||||||
|
result_queue.put(("error", repr(exc)))
|
||||||
|
return
|
||||||
|
finally:
|
||||||
|
del proxy, source
|
||||||
|
if pool is not None:
|
||||||
|
pool.shutdown()
|
||||||
|
del pool
|
||||||
|
gc.collect()
|
||||||
|
result_queue.put(("ok", None))
|
||||||
|
|
||||||
|
|
||||||
|
class TestCudaVmmTransport(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
if (
|
||||||
|
not torch.cuda.is_available()
|
||||||
|
or torch.version.cuda is None
|
||||||
|
or torch.cuda.device_count() < 3
|
||||||
|
):
|
||||||
|
raise unittest.SkipTest("At least three NVIDIA CUDA GPUs are required")
|
||||||
|
|
||||||
|
def _run_round_trip(self, mode: str):
|
||||||
|
consumer_devices = (1, 2)
|
||||||
|
torch.cuda.set_device(consumer_devices[0])
|
||||||
|
ctx = mp.get_context("spawn")
|
||||||
|
proxy_queue = ctx.Queue()
|
||||||
|
producer_results = ctx.Queue()
|
||||||
|
consumer_done = ctx.Event()
|
||||||
|
producer = ctx.Process(
|
||||||
|
target=_produce_vmm_tensor,
|
||||||
|
args=(proxy_queue, consumer_done, producer_results, mode),
|
||||||
|
)
|
||||||
|
producer.start()
|
||||||
|
proxy = second_proxy = None
|
||||||
|
reconstructed = []
|
||||||
|
producer_result = None
|
||||||
|
try:
|
||||||
|
try:
|
||||||
|
proxy, expected = proxy_queue.get(timeout=60)
|
||||||
|
except queue.Empty:
|
||||||
|
producer_result = producer_results.get(timeout=5)
|
||||||
|
_status, payload = producer_result
|
||||||
|
self.fail(
|
||||||
|
f"CUDA VMM producer failed before sending its proxy: {payload}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if mode == "posix_fallback":
|
||||||
|
self.assertIsNone(proxy.fabric_handle)
|
||||||
|
self.assertIsNotNone(proxy.posix_socket_path)
|
||||||
|
else:
|
||||||
|
self.assertIsNotNone(proxy.fabric_handle)
|
||||||
|
self.assertIsNone(proxy.posix_socket_path)
|
||||||
|
second_proxy = pickle.loads(pickle.dumps(proxy))
|
||||||
|
for tp_rank, (consumer_proxy, device) in enumerate(
|
||||||
|
zip((proxy, second_proxy), consumer_devices)
|
||||||
|
):
|
||||||
|
torch.cuda.set_device(device)
|
||||||
|
with get_parallel().override(
|
||||||
|
attn_tp_size=2,
|
||||||
|
attn_tp_rank=tp_rank,
|
||||||
|
attn_cp_size=1,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
):
|
||||||
|
tensor = consumer_proxy.reconstruct_on_target_device(
|
||||||
|
device, consumer_count=1
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize(device)
|
||||||
|
self.assertEqual(tensor.cpu().tolist(), expected)
|
||||||
|
reconstructed.append(tensor)
|
||||||
|
finally:
|
||||||
|
del reconstructed, second_proxy, proxy
|
||||||
|
_imported_pool_cache_clear()
|
||||||
|
gc.collect()
|
||||||
|
consumer_done.set()
|
||||||
|
producer.join(timeout=60)
|
||||||
|
try:
|
||||||
|
if producer_result is None:
|
||||||
|
producer_result = producer_results.get(timeout=5)
|
||||||
|
status, payload = producer_result
|
||||||
|
self.assertEqual(status, "ok", payload)
|
||||||
|
finally:
|
||||||
|
if producer.is_alive():
|
||||||
|
producer.terminate()
|
||||||
|
producer.join(timeout=10)
|
||||||
|
torch.cuda.set_device(0)
|
||||||
|
self.assertEqual(producer.exitcode, 0)
|
||||||
|
|
||||||
|
def test_posix_fd_fallback_tensor_round_trip_and_recycling(self):
|
||||||
|
self._run_round_trip(mode="posix_fallback")
|
||||||
|
|
||||||
|
def test_auto_prefers_fabric_tensor_round_trip_and_recycling(self):
|
||||||
|
self._run_round_trip(mode="auto")
|
||||||
|
|
||||||
|
def test_reused_chunk_clears_acknowledgements(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 2, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
old = pool.wrap_tensor(torch.ones(1024, dtype=torch.uint8, device="cuda:0"))
|
||||||
|
pool.memory_pool[old.control_offset : old.control_offset + 8].view(
|
||||||
|
torch.int32
|
||||||
|
).fill_(1)
|
||||||
|
torch.cuda.synchronize(0)
|
||||||
|
with pool._lock:
|
||||||
|
pool._recycle_chunks()
|
||||||
|
pool._merge_chunks()
|
||||||
|
|
||||||
|
pool.wrap_tensor(torch.ones(100, dtype=torch.uint8, device="cuda:0"))
|
||||||
|
live = pool.wrap_tensor(torch.ones(256, dtype=torch.uint8, device="cuda:0"))
|
||||||
|
control = pool.memory_pool[
|
||||||
|
live.control_offset : live.control_offset + 8
|
||||||
|
].view(torch.int32)
|
||||||
|
self.assertTrue(torch.equal(control, torch.zeros_like(control)))
|
||||||
|
|
||||||
|
with pool._lock:
|
||||||
|
pool._recycle_chunks()
|
||||||
|
self.assertIn(
|
||||||
|
live.control_offset,
|
||||||
|
[chunk.start for chunk in pool.occupied_chunks],
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_packed_tensors_round_trip_through_one_shared_buffer(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
sources = [
|
||||||
|
torch.arange(24, dtype=torch.float32, device="cuda:0")
|
||||||
|
.reshape(4, 6)
|
||||||
|
.transpose(0, 1),
|
||||||
|
torch.arange(7, dtype=torch.bfloat16),
|
||||||
|
torch.arange(5, dtype=torch.int64, device="cuda:0"),
|
||||||
|
]
|
||||||
|
expected = [source.contiguous().cpu() for source in sources]
|
||||||
|
proxies = reconstructed = None
|
||||||
|
try:
|
||||||
|
stream = MagicMock(wraps=torch.cuda.current_stream(0))
|
||||||
|
with patch("torch.cuda.current_stream", return_value=stream):
|
||||||
|
proxies = pool.wrap_tensors(sources)
|
||||||
|
|
||||||
|
self.assertIsNotNone(proxies)
|
||||||
|
self.assertEqual(stream.synchronize.call_count, 1)
|
||||||
|
self.assertEqual(len(pool.occupied_chunks), 1)
|
||||||
|
self.assertTrue(
|
||||||
|
all(
|
||||||
|
isinstance(proxy, CudaVmmPackedTensorTransportProxy)
|
||||||
|
for proxy in proxies
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(len({proxy.control_offset for proxy in proxies}), 1)
|
||||||
|
|
||||||
|
proxies = pickle.loads(pickle.dumps(proxies))
|
||||||
|
self.assertIs(proxies[0]._packed_owner, proxies[-1]._packed_owner)
|
||||||
|
with get_parallel().override(
|
||||||
|
attn_tp_size=1,
|
||||||
|
attn_tp_rank=0,
|
||||||
|
attn_cp_size=1,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
):
|
||||||
|
reconstructed = [
|
||||||
|
proxy.reconstruct_on_target_device(0, consumer_count=1)
|
||||||
|
for proxy in proxies
|
||||||
|
]
|
||||||
|
torch.cuda.synchronize(0)
|
||||||
|
|
||||||
|
for actual, wanted in zip(reconstructed, expected):
|
||||||
|
self.assertTrue(torch.equal(actual.cpu(), wanted))
|
||||||
|
packed_storage = proxies[
|
||||||
|
0
|
||||||
|
]._packed_owner.reconstruct_tensor.untyped_storage()
|
||||||
|
self.assertTrue(
|
||||||
|
all(
|
||||||
|
tensor.untyped_storage().data_ptr() == packed_storage.data_ptr()
|
||||||
|
for tensor in reconstructed
|
||||||
|
)
|
||||||
|
)
|
||||||
|
with pool._lock:
|
||||||
|
pool._recycle_chunks()
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
finally:
|
||||||
|
del reconstructed, proxies, expected, sources
|
||||||
|
_imported_pool_cache_clear()
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_packed_cancel_is_shared_and_idempotent(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
proxies = pool.wrap_tensors(
|
||||||
|
[
|
||||||
|
torch.ones(8, dtype=torch.float32, device="cuda:0"),
|
||||||
|
torch.ones(8, dtype=torch.float32, device="cuda:0"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(proxies)
|
||||||
|
|
||||||
|
pool.cancel_proxy(proxies[0])
|
||||||
|
pool.cancel_proxy(proxies[1])
|
||||||
|
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
self.assertEqual(
|
||||||
|
sum(chunk.size for chunk in pool.available_chunks),
|
||||||
|
pool.allocation_size,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_packed_reservation_failure_returns_fallback_signal(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
source = torch.empty(
|
||||||
|
pool.allocation_size, dtype=torch.uint8, device="cuda:0"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(pool.wrap_tensors([source]))
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
self.assertEqual(
|
||||||
|
sum(chunk.size for chunk in pool.available_chunks),
|
||||||
|
pool.allocation_size,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_oversized_tensor_falls_back_to_cpu(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
source = torch.empty(
|
||||||
|
pool.allocation_size, dtype=torch.uint8, device="cuda:0"
|
||||||
|
)
|
||||||
|
|
||||||
|
fallback = pool.wrap_tensor(source)
|
||||||
|
|
||||||
|
self.assertTrue(fallback.is_cpu)
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
self.assertEqual(
|
||||||
|
sum(chunk.size for chunk in pool.available_chunks),
|
||||||
|
pool.allocation_size,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_failed_copy_rolls_back_reservation(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
with self.assertRaisesRegex(NotImplementedError, "meta tensor"):
|
||||||
|
pool.wrap_tensor(torch.ones(16, device="meta"))
|
||||||
|
with self.assertRaisesRegex(NotImplementedError, "meta tensor"):
|
||||||
|
pool.wrap_tensors(
|
||||||
|
[
|
||||||
|
torch.ones(16, device="cuda:0"),
|
||||||
|
torch.ones(16, device="meta"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
self.assertEqual(
|
||||||
|
sum(chunk.size for chunk in pool.available_chunks),
|
||||||
|
pool.allocation_size,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_undispatched_proxy_can_be_cancelled_immediately(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
proxy = pool.wrap_tensor(torch.ones(16, device="cuda:0"))
|
||||||
|
|
||||||
|
pool.cancel_proxy(proxy)
|
||||||
|
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
self.assertEqual(
|
||||||
|
sum(chunk.size for chunk in pool.available_chunks),
|
||||||
|
pool.allocation_size,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_failed_cleanup_sync_quarantines_pool(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
stream = MagicMock()
|
||||||
|
stream.synchronize.side_effect = RuntimeError("forced sync failure")
|
||||||
|
try:
|
||||||
|
with (
|
||||||
|
patch("torch.cuda.current_stream", return_value=stream),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "forced sync failure"),
|
||||||
|
):
|
||||||
|
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
|
||||||
|
torch.cuda.synchronize(0)
|
||||||
|
self.assertIsNotNone(pool._pool_error)
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "pool failed"):
|
||||||
|
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
|
||||||
|
finally:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_shutdown_waits_for_active_publisher(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
publisher_entered = threading.Event()
|
||||||
|
allow_publisher_to_finish = threading.Event()
|
||||||
|
shutdown_entered = threading.Event()
|
||||||
|
shutdown_finished = threading.Event()
|
||||||
|
errors = []
|
||||||
|
real_stream = torch.cuda.current_stream(0)
|
||||||
|
stream = MagicMock(wraps=real_stream)
|
||||||
|
|
||||||
|
def synchronize():
|
||||||
|
publisher_entered.set()
|
||||||
|
if not allow_publisher_to_finish.wait(timeout=10):
|
||||||
|
raise TimeoutError("publisher was not released")
|
||||||
|
real_stream.synchronize()
|
||||||
|
|
||||||
|
stream.synchronize.side_effect = synchronize
|
||||||
|
|
||||||
|
def publish():
|
||||||
|
try:
|
||||||
|
with patch("torch.cuda.current_stream", return_value=stream):
|
||||||
|
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
|
||||||
|
except Exception as error: # pragma: no cover
|
||||||
|
errors.append(error)
|
||||||
|
|
||||||
|
def shutdown():
|
||||||
|
shutdown_entered.set()
|
||||||
|
try:
|
||||||
|
pool.shutdown()
|
||||||
|
except Exception as error: # pragma: no cover
|
||||||
|
errors.append(error)
|
||||||
|
finally:
|
||||||
|
shutdown_finished.set()
|
||||||
|
|
||||||
|
publisher = threading.Thread(target=publish)
|
||||||
|
shutdown_thread = threading.Thread(target=shutdown)
|
||||||
|
try:
|
||||||
|
publisher.start()
|
||||||
|
self.assertTrue(publisher_entered.wait(timeout=10))
|
||||||
|
shutdown_thread.start()
|
||||||
|
self.assertTrue(shutdown_entered.wait(timeout=10))
|
||||||
|
self.assertFalse(shutdown_finished.wait(timeout=0.1))
|
||||||
|
finally:
|
||||||
|
allow_publisher_to_finish.set()
|
||||||
|
publisher.join(timeout=10)
|
||||||
|
shutdown_thread.join(timeout=10)
|
||||||
|
if not pool._closed:
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
self.assertFalse(publisher.is_alive())
|
||||||
|
self.assertFalse(shutdown_thread.is_alive())
|
||||||
|
self.assertFalse(errors)
|
||||||
|
|
||||||
|
def test_consumer_copy_failure_releases_slice_without_allowing_retry(self):
|
||||||
|
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
|
||||||
|
try:
|
||||||
|
proxy = pool.wrap_tensor(torch.ones(1, device="cuda:0"))
|
||||||
|
proxy.shape = (2,)
|
||||||
|
with (
|
||||||
|
get_parallel().override(
|
||||||
|
attn_tp_size=1,
|
||||||
|
attn_tp_rank=0,
|
||||||
|
attn_cp_size=1,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
),
|
||||||
|
self.assertRaises(RuntimeError),
|
||||||
|
):
|
||||||
|
proxy.reconstruct_on_target_device(0, consumer_count=1)
|
||||||
|
torch.cuda.synchronize(0)
|
||||||
|
with pool._lock:
|
||||||
|
pool._recycle_chunks()
|
||||||
|
self.assertFalse(pool.occupied_chunks)
|
||||||
|
|
||||||
|
with (
|
||||||
|
get_parallel().override(
|
||||||
|
attn_tp_size=1,
|
||||||
|
attn_tp_rank=0,
|
||||||
|
attn_cp_size=1,
|
||||||
|
attn_cp_rank=0,
|
||||||
|
),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "already released"),
|
||||||
|
):
|
||||||
|
proxy.reconstruct_on_target_device(0, consumer_count=1)
|
||||||
|
finally:
|
||||||
|
_imported_pool_cache_clear()
|
||||||
|
pool.shutdown()
|
||||||
|
|
||||||
|
def test_posix_export_fd_closes_when_allocation_setup_fails(self):
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.utils.cuda_vmm_transport_utils._tensor_from_pointer",
|
||||||
|
side_effect=RuntimeError("forced storage failure"),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.utils.cuda_vmm_transport_utils.os.close",
|
||||||
|
wraps=os.close,
|
||||||
|
) as close_fd,
|
||||||
|
self.assertRaisesRegex(RuntimeError, "forced storage failure"),
|
||||||
|
):
|
||||||
|
_FabricUnavailableCudaVmmMemoryPool(
|
||||||
|
4 << 20, 60, 0, 1, allow_posix_fallback=True
|
||||||
|
)
|
||||||
|
close_fd.assert_called_once()
|
||||||
|
|
||||||
|
def test_stream_setup_failure_releases_pool_and_posix_broker(self):
|
||||||
|
release_allocation = CudaVmmMemoryPool._release_allocation
|
||||||
|
close_broker = _PosixFdBroker.close
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.utils.cuda_vmm_transport_utils.torch.cuda.Stream",
|
||||||
|
side_effect=RuntimeError("forced stream failure"),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
CudaVmmMemoryPool,
|
||||||
|
"_release_allocation",
|
||||||
|
autospec=True,
|
||||||
|
side_effect=release_allocation,
|
||||||
|
) as release_pool,
|
||||||
|
patch.object(
|
||||||
|
_PosixFdBroker,
|
||||||
|
"close",
|
||||||
|
autospec=True,
|
||||||
|
side_effect=close_broker,
|
||||||
|
) as close_fd_broker,
|
||||||
|
self.assertRaisesRegex(RuntimeError, "forced stream failure"),
|
||||||
|
):
|
||||||
|
_FabricUnavailableCudaVmmMemoryPool(
|
||||||
|
4 << 20, 60, 0, 1, allow_posix_fallback=True
|
||||||
|
)
|
||||||
|
|
||||||
|
close_fd_broker.assert_called_once()
|
||||||
|
release_pool.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main(verbosity=2)
|
||||||
@@ -0,0 +1,666 @@
|
|||||||
|
import unittest
|
||||||
|
from contextlib import nullcontext
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, call, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestCudaVmmFeatureTransport(unittest.TestCase):
|
||||||
|
def test_partial_pool_release_can_be_retried(self):
|
||||||
|
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
|
||||||
|
|
||||||
|
pool = object.__new__(vmm.CudaVmmMemoryPool)
|
||||||
|
pool.memory_pool = object()
|
||||||
|
pool.use_fabric = True
|
||||||
|
pool.shareable_handle = b"handle"
|
||||||
|
pool._pool_pointer = 123
|
||||||
|
pool._allocation_handle = 456
|
||||||
|
pool._allocation_mapped = True
|
||||||
|
pool.allocation_size = 4096
|
||||||
|
pool.device_index = 0
|
||||||
|
driver = MagicMock()
|
||||||
|
driver.cuMemUnmap.return_value = "unmap"
|
||||||
|
driver.cuMemAddressFree.return_value = "address_free"
|
||||||
|
driver.cuMemRelease.return_value = "release"
|
||||||
|
failed_once = False
|
||||||
|
|
||||||
|
def check_driver(result, _operation):
|
||||||
|
nonlocal failed_once
|
||||||
|
if result == "address_free" and not failed_once:
|
||||||
|
failed_once = True
|
||||||
|
raise RuntimeError("forced address-free failure")
|
||||||
|
return result
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(vmm, "_get_cuda_driver", return_value=driver),
|
||||||
|
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
|
||||||
|
patch.object(vmm, "check_drv", side_effect=check_driver),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "forced address-free failure"),
|
||||||
|
):
|
||||||
|
pool._release_allocation()
|
||||||
|
|
||||||
|
self.assertFalse(pool._allocation_mapped)
|
||||||
|
self.assertEqual(pool._pool_pointer, 123)
|
||||||
|
self.assertEqual(pool._allocation_handle, 456)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(vmm, "_get_cuda_driver", return_value=driver),
|
||||||
|
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
|
||||||
|
patch.object(vmm, "check_drv", side_effect=lambda result, _: result),
|
||||||
|
):
|
||||||
|
pool._release_allocation()
|
||||||
|
|
||||||
|
self.assertIsNone(pool._pool_pointer)
|
||||||
|
self.assertIsNone(pool._allocation_handle)
|
||||||
|
self.assertEqual(driver.cuMemUnmap.call_count, 1)
|
||||||
|
self.assertEqual(driver.cuMemAddressFree.call_count, 2)
|
||||||
|
self.assertEqual(driver.cuMemRelease.call_count, 1)
|
||||||
|
|
||||||
|
def test_model_class_controls_cuda_vmm_opt_in(self):
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
|
||||||
|
class SupportedModel:
|
||||||
|
supports_cuda_vmm_feature_transport = True
|
||||||
|
|
||||||
|
class UnsupportedModel:
|
||||||
|
pass
|
||||||
|
|
||||||
|
manager = object.__new__(TokenizerManager)
|
||||||
|
manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
|
||||||
|
manager.model_config = object()
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.model_loader.utils.get_model_architecture",
|
||||||
|
return_value=(SupportedModel, "supported"),
|
||||||
|
):
|
||||||
|
manager._validate_cuda_vmm_feature_transport_support()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.model_loader.utils.get_model_architecture",
|
||||||
|
return_value=(UnsupportedModel, "unsupported"),
|
||||||
|
),
|
||||||
|
self.assertRaisesRegex(ValueError, "UnsupportedModel"),
|
||||||
|
):
|
||||||
|
manager._validate_cuda_vmm_feature_transport_support()
|
||||||
|
|
||||||
|
def test_cpu_transport_skips_model_opt_in_lookup(self):
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
|
||||||
|
manager = object.__new__(TokenizerManager)
|
||||||
|
manager.server_args = SimpleNamespace(mm_feature_transport="cpu")
|
||||||
|
manager.model_config = object()
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.model_loader.utils.get_model_architecture"
|
||||||
|
) as get_model_architecture:
|
||||||
|
manager._validate_cuda_vmm_feature_transport_support()
|
||||||
|
|
||||||
|
get_model_architecture.assert_not_called()
|
||||||
|
|
||||||
|
def test_vmm_transport_initializes_pool(self):
|
||||||
|
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
|
||||||
|
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
mm_feature_transport="cuda_vmm",
|
||||||
|
tokenizer_worker_num=2,
|
||||||
|
base_gpu_id=3,
|
||||||
|
enable_dp_attention=False,
|
||||||
|
tp_size=4,
|
||||||
|
nnodes=1,
|
||||||
|
)
|
||||||
|
pool = object()
|
||||||
|
with (
|
||||||
|
patch.object(vmm, "get_mm_feature_pool_size_per_worker", return_value=123),
|
||||||
|
patch.object(vmm, "CudaVmmMemoryPool", return_value=pool) as pool_class,
|
||||||
|
):
|
||||||
|
transport = vmm.CudaVmmFeatureTransport(server_args, SimpleNamespace())
|
||||||
|
|
||||||
|
self.assertIs(transport.pool, pool)
|
||||||
|
pool_class.assert_called_once_with(
|
||||||
|
memory_size=123,
|
||||||
|
recycle_interval=vmm.MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
||||||
|
base_gpu_id=3,
|
||||||
|
consumer_count=4,
|
||||||
|
allow_posix_fallback=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_disabled_transport_is_a_noop(self):
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = CudaVmmFeatureTransport(
|
||||||
|
SimpleNamespace(mm_feature_transport="cpu"), None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(transport.prepare_for_dispatch([None]), [])
|
||||||
|
transport.cancel_for_dispatch([])
|
||||||
|
transport.shutdown()
|
||||||
|
self.assertIsNone(transport.pool)
|
||||||
|
|
||||||
|
def test_vmm_transport_requires_processor(self):
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "multimodal processor"):
|
||||||
|
CudaVmmFeatureTransport(
|
||||||
|
SimpleNamespace(mm_feature_transport="cuda_vmm"), None
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_image_features_are_packed_per_request(self):
|
||||||
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = MagicMock()
|
||||||
|
features = [torch.arange(4), torch.arange(4, 8)]
|
||||||
|
proxies = [object(), object()]
|
||||||
|
transport.pool.wrap_tensors.return_value = proxies
|
||||||
|
items = [
|
||||||
|
MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
|
||||||
|
for feature in features
|
||||||
|
]
|
||||||
|
|
||||||
|
transport.wrap_items(items)
|
||||||
|
|
||||||
|
transport.pool.wrap_tensors.assert_called_once_with(features)
|
||||||
|
transport.pool.wrap_tensor.assert_not_called()
|
||||||
|
self.assertEqual([item.feature for item in items], proxies)
|
||||||
|
|
||||||
|
def test_deferred_features_are_not_packed(self):
|
||||||
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
|
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||||
|
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = MagicMock()
|
||||||
|
features = [torch.arange(4), torch.arange(4, 8)]
|
||||||
|
proxies = [object(), object()]
|
||||||
|
transport.pool.wrap_tensor.side_effect = proxies
|
||||||
|
items = [
|
||||||
|
MultimodalDataItem(
|
||||||
|
modality=Modality.IMAGE,
|
||||||
|
feature=feature,
|
||||||
|
model_specific_data={DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY: True},
|
||||||
|
)
|
||||||
|
for feature in features
|
||||||
|
]
|
||||||
|
|
||||||
|
transport.wrap_items(items)
|
||||||
|
|
||||||
|
transport.pool.wrap_tensors.assert_not_called()
|
||||||
|
self.assertEqual(
|
||||||
|
transport.pool.wrap_tensor.call_args_list,
|
||||||
|
[call(feature) for feature in features],
|
||||||
|
)
|
||||||
|
self.assertEqual([item.feature for item in items], proxies)
|
||||||
|
|
||||||
|
def test_tensor_containers_fail_closed(self):
|
||||||
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = MagicMock()
|
||||||
|
item = MultimodalDataItem(
|
||||||
|
modality=Modality.IMAGE,
|
||||||
|
feature=[torch.arange(4), torch.arange(4, 8)],
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(TypeError, "single tensor"):
|
||||||
|
transport.wrap_items([item])
|
||||||
|
|
||||||
|
transport.pool.wrap_tensor.assert_not_called()
|
||||||
|
transport.pool.wrap_tensors.assert_not_called()
|
||||||
|
|
||||||
|
def test_partial_failure_restores_tensors_and_cancels_packed_chunk_once(self):
|
||||||
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
CudaVmmMemoryPool,
|
||||||
|
CudaVmmPackedTensorTransportProxy,
|
||||||
|
_CudaVmmPackedTransportOwner,
|
||||||
|
)
|
||||||
|
|
||||||
|
owner = object.__new__(_CudaVmmPackedTransportOwner)
|
||||||
|
owner.control_offset = 64
|
||||||
|
owner._producer_cancelled = False
|
||||||
|
proxies = [object.__new__(CudaVmmPackedTensorTransportProxy) for _ in range(2)]
|
||||||
|
for proxy in proxies:
|
||||||
|
proxy._packed_owner = owner
|
||||||
|
|
||||||
|
pool = object.__new__(CudaVmmMemoryPool)
|
||||||
|
pool.wrap_tensors = MagicMock(return_value=proxies)
|
||||||
|
pool.wrap_tensor = MagicMock(side_effect=RuntimeError("copy failed"))
|
||||||
|
pool._cancel_control_offset = MagicMock()
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = pool
|
||||||
|
|
||||||
|
features = [torch.arange(4), torch.arange(4, 8)]
|
||||||
|
embedding = torch.arange(2)
|
||||||
|
items = [
|
||||||
|
MultimodalDataItem(
|
||||||
|
modality=Modality.IMAGE,
|
||||||
|
feature=features[0],
|
||||||
|
precomputed_embeddings=embedding,
|
||||||
|
),
|
||||||
|
MultimodalDataItem(modality=Modality.IMAGE, feature=features[1]),
|
||||||
|
]
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "copy failed"):
|
||||||
|
transport.wrap_items(items)
|
||||||
|
|
||||||
|
for item, feature in zip(items, features, strict=True):
|
||||||
|
self.assertIs(item.feature, feature)
|
||||||
|
self.assertIs(items[0].precomputed_embeddings, embedding)
|
||||||
|
pool._cancel_control_offset.assert_called_once_with(owner.control_offset)
|
||||||
|
|
||||||
|
def test_text_request_uses_base_send_path(self):
|
||||||
|
from sglang.srt.managers import tokenizer_manager
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
|
||||||
|
manager = object.__new__(TokenizerManager)
|
||||||
|
transport = MagicMock()
|
||||||
|
transport.prepare_for_dispatch.return_value = []
|
||||||
|
manager.cuda_vmm_feature_transport = transport
|
||||||
|
manager._dispatch_to_scheduler = MagicMock()
|
||||||
|
state = SimpleNamespace(dispatched=False)
|
||||||
|
manager.rid_to_state = {"test-request": state}
|
||||||
|
tokenized_obj = SimpleNamespace(
|
||||||
|
rid="test-request",
|
||||||
|
mm_inputs=None,
|
||||||
|
time_stats=MagicMock(),
|
||||||
|
wrap_pickle_fields=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj):
|
||||||
|
manager._send_one_request(tokenized_obj)
|
||||||
|
|
||||||
|
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
|
||||||
|
transport.prepare_for_dispatch.assert_called_once_with((None,))
|
||||||
|
transport.cancel_for_dispatch.assert_not_called()
|
||||||
|
self.assertTrue(state.dispatched)
|
||||||
|
|
||||||
|
def test_failed_dispatch_cancels_published_items(self):
|
||||||
|
from sglang.srt.managers import tokenizer_manager
|
||||||
|
from sglang.srt.managers.schedule_batch import (
|
||||||
|
Modality,
|
||||||
|
MultimodalDataItem,
|
||||||
|
MultimodalProcessorOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
|
transport = MagicMock()
|
||||||
|
manager._dispatch_to_scheduler = MagicMock(
|
||||||
|
side_effect=RuntimeError("send failed")
|
||||||
|
)
|
||||||
|
state = SimpleNamespace(dispatched=False)
|
||||||
|
manager.rid_to_state = {"test-request": state}
|
||||||
|
items = [MultimodalDataItem(modality=Modality.IMAGE, feature=torch.arange(2))]
|
||||||
|
tokenized_obj = SimpleNamespace(
|
||||||
|
rid="test-request",
|
||||||
|
mm_inputs=MultimodalProcessorOutput(input_ids=[1], mm_items=items),
|
||||||
|
time_stats=MagicMock(),
|
||||||
|
wrap_pickle_fields=MagicMock(),
|
||||||
|
)
|
||||||
|
transport.prepare_for_dispatch.return_value = items
|
||||||
|
manager.cuda_vmm_feature_transport = transport
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "send failed"),
|
||||||
|
):
|
||||||
|
manager._send_one_request(tokenized_obj)
|
||||||
|
|
||||||
|
transport.prepare_for_dispatch.assert_called_once_with(
|
||||||
|
(tokenized_obj.mm_inputs,)
|
||||||
|
)
|
||||||
|
transport.cancel_for_dispatch.assert_called_once_with(items)
|
||||||
|
self.assertFalse(state.dispatched)
|
||||||
|
|
||||||
|
def test_post_dispatch_failure_does_not_cancel_published_items(self):
|
||||||
|
from sglang.srt.managers import tokenizer_manager
|
||||||
|
from sglang.srt.managers.schedule_batch import (
|
||||||
|
Modality,
|
||||||
|
MultimodalDataItem,
|
||||||
|
MultimodalProcessorOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
manager = object.__new__(tokenizer_manager.TokenizerManager)
|
||||||
|
transport = MagicMock()
|
||||||
|
manager._dispatch_to_scheduler = MagicMock()
|
||||||
|
state = SimpleNamespace(dispatched=False)
|
||||||
|
manager.rid_to_state = {"test-request": state}
|
||||||
|
time_stats = MagicMock()
|
||||||
|
time_stats.set_api_server_dispatch_finish_time.side_effect = RuntimeError(
|
||||||
|
"bookkeeping failed"
|
||||||
|
)
|
||||||
|
items = [MultimodalDataItem(modality=Modality.IMAGE, feature=torch.arange(2))]
|
||||||
|
tokenized_obj = SimpleNamespace(
|
||||||
|
rid="test-request",
|
||||||
|
mm_inputs=MultimodalProcessorOutput(input_ids=[1], mm_items=items),
|
||||||
|
time_stats=time_stats,
|
||||||
|
wrap_pickle_fields=MagicMock(),
|
||||||
|
)
|
||||||
|
transport.prepare_for_dispatch.return_value = items
|
||||||
|
manager.cuda_vmm_feature_transport = transport
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "bookkeeping failed"),
|
||||||
|
):
|
||||||
|
manager._send_one_request(tokenized_obj)
|
||||||
|
|
||||||
|
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
|
||||||
|
transport.cancel_for_dispatch.assert_not_called()
|
||||||
|
self.assertTrue(state.dispatched)
|
||||||
|
|
||||||
|
def test_prepare_batch_cancels_prior_groups_on_failure(self):
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = MagicMock()
|
||||||
|
transport.wrap_items = MagicMock(
|
||||||
|
side_effect=[None, RuntimeError("wrap failed")]
|
||||||
|
)
|
||||||
|
transport.cancel_for_dispatch = MagicMock()
|
||||||
|
item_groups = [[object()], [object()]]
|
||||||
|
mm_inputs_batch = [SimpleNamespace(mm_items=items) for items in item_groups]
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "wrap failed"):
|
||||||
|
transport.prepare_for_dispatch(mm_inputs_batch)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
transport.wrap_items.call_args_list,
|
||||||
|
[call(item_groups[0]), call(item_groups[1])],
|
||||||
|
)
|
||||||
|
transport.cancel_for_dispatch.assert_called_once_with(item_groups[0])
|
||||||
|
|
||||||
|
def test_prepare_batch_returns_flattened_items(self):
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
transport.pool = MagicMock()
|
||||||
|
transport.wrap_items = MagicMock()
|
||||||
|
item_groups = [[object()], [object(), object()]]
|
||||||
|
|
||||||
|
prepared = transport.prepare_for_dispatch(
|
||||||
|
[
|
||||||
|
None,
|
||||||
|
SimpleNamespace(mm_items=[]),
|
||||||
|
*(SimpleNamespace(mm_items=items) for items in item_groups),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(prepared, item_groups[0] + item_groups[1])
|
||||||
|
self.assertEqual(
|
||||||
|
transport.wrap_items.call_args_list,
|
||||||
|
[call(items) for items in item_groups],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_engine_shutdown_is_idempotent(self):
|
||||||
|
from sglang.srt.entrypoints import engine as engine_module
|
||||||
|
from sglang.srt.entrypoints.engine import Engine
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
manager = object.__new__(TokenizerManager)
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
pool = MagicMock()
|
||||||
|
transport.pool = pool
|
||||||
|
manager.cuda_vmm_feature_transport = transport
|
||||||
|
manager._subprocess_watchdog = None
|
||||||
|
engine = object.__new__(Engine)
|
||||||
|
engine.tokenizer_manager = manager
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
engine_module,
|
||||||
|
"kill_process_tree",
|
||||||
|
side_effect=RuntimeError("base failed"),
|
||||||
|
):
|
||||||
|
for _ in range(2):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "base failed"):
|
||||||
|
engine.shutdown()
|
||||||
|
|
||||||
|
self.assertEqual(pool.shutdown.call_count, 2)
|
||||||
|
self.assertIs(transport.pool, pool)
|
||||||
|
|
||||||
|
def test_engine_startup_failure_releases_parent_pool(self):
|
||||||
|
from sglang.srt.entrypoints import engine as engine_module
|
||||||
|
from sglang.srt.entrypoints.engine import Engine
|
||||||
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
manager = object.__new__(TokenizerManager)
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
pool = MagicMock()
|
||||||
|
transport.pool = pool
|
||||||
|
manager.cuda_vmm_feature_transport = transport
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
remote_instance_weight_loader_start_seed_via_transfer_engine=False,
|
||||||
|
reasoning_parser=None,
|
||||||
|
tool_call_parser=None,
|
||||||
|
weight_cache_mode=None,
|
||||||
|
enable_elastic_expert_backup=False,
|
||||||
|
elastic_ep_backend=None,
|
||||||
|
node_rank=0,
|
||||||
|
tokenizer_worker_num=1,
|
||||||
|
check_server_args=MagicMock(),
|
||||||
|
)
|
||||||
|
scheduler_init_result = SimpleNamespace(
|
||||||
|
all_child_pids=[],
|
||||||
|
scheduler_infos=[],
|
||||||
|
wait_for_ready=MagicMock(side_effect=RuntimeError("startup failed")),
|
||||||
|
engine_info_bootstrap_server=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(engine_module, "configure_logger"),
|
||||||
|
patch.object(engine_module, "_set_envs_and_config"),
|
||||||
|
patch.object(engine_module, "load_plugins"),
|
||||||
|
patch.object(
|
||||||
|
Engine,
|
||||||
|
"_launch_scheduler_processes",
|
||||||
|
return_value=(scheduler_init_result, []),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
Engine, "_launch_detokenizer_subprocesses", return_value=([], [])
|
||||||
|
),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "startup failed"),
|
||||||
|
):
|
||||||
|
Engine._launch_subprocesses(
|
||||||
|
server_args=server_args,
|
||||||
|
init_tokenizer_manager_func=MagicMock(return_value=(manager, object())),
|
||||||
|
run_scheduler_process_func=MagicMock(),
|
||||||
|
run_detokenizer_process_func=MagicMock(),
|
||||||
|
port_args=SimpleNamespace(),
|
||||||
|
)
|
||||||
|
|
||||||
|
pool.shutdown.assert_called_once_with()
|
||||||
|
self.assertIs(transport.pool, pool)
|
||||||
|
|
||||||
|
def test_failed_pool_shutdown_remains_retryable(self):
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmFeatureTransport,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = object.__new__(CudaVmmFeatureTransport)
|
||||||
|
pool = MagicMock()
|
||||||
|
pool.shutdown.side_effect = [RuntimeError("shutdown failed"), None]
|
||||||
|
transport.pool = pool
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "shutdown failed"):
|
||||||
|
transport.shutdown()
|
||||||
|
self.assertIs(transport.pool, pool)
|
||||||
|
|
||||||
|
transport.shutdown()
|
||||||
|
self.assertIs(transport.pool, pool)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSchedulerMmTransportBoundary(unittest.TestCase):
|
||||||
|
@staticmethod
|
||||||
|
def _prepare_scheduler(scheduler):
|
||||||
|
scheduler.session_controller = SimpleNamespace(maybe_reap=MagicMock())
|
||||||
|
scheduler._request_dispatcher = MagicMock(return_value=None)
|
||||||
|
scheduler.flush_wrapper = SimpleNamespace(check_pending=MagicMock())
|
||||||
|
scheduler.external_corpus_manager = None
|
||||||
|
|
||||||
|
def test_materializes_inputs_directly_before_base_dispatch(self):
|
||||||
|
from sglang.srt.managers import scheduler as scheduler_module
|
||||||
|
|
||||||
|
scheduler = object.__new__(scheduler_module.Scheduler)
|
||||||
|
scheduler.server_args = SimpleNamespace(
|
||||||
|
mm_feature_transport="cuda_vmm",
|
||||||
|
enable_broadcast_mm_inputs_process=True,
|
||||||
|
)
|
||||||
|
self._prepare_scheduler(scheduler)
|
||||||
|
raw_inputs = object()
|
||||||
|
materialized = object()
|
||||||
|
request = SimpleNamespace(mm_inputs=raw_inputs)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "TokenizedGenerateReqInput", SimpleNamespace
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
scheduler_module.MultimodalInputs,
|
||||||
|
"from_processor_output",
|
||||||
|
return_value=materialized,
|
||||||
|
) as build_inputs,
|
||||||
|
patch.object(
|
||||||
|
scheduler, "_process_and_broadcast_mm_inputs"
|
||||||
|
) as cpu_broadcast,
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "is_health_check_generate_req", return_value=False
|
||||||
|
),
|
||||||
|
):
|
||||||
|
scheduler.process_input_requests([request])
|
||||||
|
|
||||||
|
build_inputs.assert_called_once_with(raw_inputs)
|
||||||
|
self.assertIs(request.mm_inputs, materialized)
|
||||||
|
scheduler._request_dispatcher.assert_called_once_with(request)
|
||||||
|
cpu_broadcast.assert_not_called()
|
||||||
|
|
||||||
|
def test_materializes_batched_inputs_before_dispatch(self):
|
||||||
|
from sglang.srt.managers import scheduler as scheduler_module
|
||||||
|
|
||||||
|
class TokenizedRequest:
|
||||||
|
def __init__(self, mm_inputs):
|
||||||
|
self.mm_inputs = mm_inputs
|
||||||
|
|
||||||
|
class BatchRequest:
|
||||||
|
def __init__(self, batch):
|
||||||
|
self.batch = batch
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
return iter(self.batch)
|
||||||
|
|
||||||
|
scheduler = object.__new__(scheduler_module.Scheduler)
|
||||||
|
scheduler.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
|
||||||
|
self._prepare_scheduler(scheduler)
|
||||||
|
raw_inputs = [object(), object()]
|
||||||
|
materialized = [object(), object()]
|
||||||
|
inner_requests = [TokenizedRequest(value) for value in raw_inputs]
|
||||||
|
request = BatchRequest(inner_requests)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "TokenizedGenerateReqInput", TokenizedRequest
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "TokenizedEmbeddingReqInput", TokenizedRequest
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "BatchTokenizedGenerateReqInput", BatchRequest
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "BatchTokenizedEmbeddingReqInput", BatchRequest
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
scheduler_module.MultimodalInputs,
|
||||||
|
"from_processor_output",
|
||||||
|
side_effect=materialized,
|
||||||
|
) as build_inputs,
|
||||||
|
patch.object(
|
||||||
|
scheduler_module, "is_health_check_generate_req", return_value=False
|
||||||
|
),
|
||||||
|
):
|
||||||
|
scheduler.process_input_requests([request])
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
build_inputs.call_args_list,
|
||||||
|
[call(value) for value in raw_inputs],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
[inner.mm_inputs for inner in inner_requests],
|
||||||
|
materialized,
|
||||||
|
)
|
||||||
|
scheduler._request_dispatcher.assert_called_once_with(request)
|
||||||
|
|
||||||
|
def test_already_materialized_inputs_are_reused(self):
|
||||||
|
from sglang.srt.managers.schedule_batch import MultimodalInputs
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
|
||||||
|
scheduler = object.__new__(Scheduler)
|
||||||
|
mm_inputs = MultimodalInputs(mm_items=[])
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
scheduler, "_process_and_broadcast_mm_inputs"
|
||||||
|
) as process_and_broadcast:
|
||||||
|
self.assertIs(scheduler._get_multimodal_inputs(mm_inputs), mm_inputs)
|
||||||
|
|
||||||
|
process_and_broadcast.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
class TestVmmConsumerCount(unittest.TestCase):
|
||||||
|
def test_proxy_defaults_to_one_consumer(self):
|
||||||
|
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
|
||||||
|
|
||||||
|
proxy = object.__new__(vmm.CudaVmmTensorTransportProxy)
|
||||||
|
proxy.consumer_count = 4
|
||||||
|
self.assertEqual(proxy._resolve_consumer_count(None), 1)
|
||||||
|
self.assertEqual(proxy._resolve_consumer_count(2), 2)
|
||||||
|
|
||||||
|
def test_acknowledgement_ranges_include_cp_rank(self):
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
|
CudaVmmTensorTransportProxy,
|
||||||
|
)
|
||||||
|
|
||||||
|
proxy = object.__new__(CudaVmmTensorTransportProxy)
|
||||||
|
proxy.consumer_count = 4
|
||||||
|
with get_parallel().override(
|
||||||
|
attn_tp_size=2,
|
||||||
|
attn_tp_rank=1,
|
||||||
|
attn_cp_size=2,
|
||||||
|
attn_cp_rank=1,
|
||||||
|
):
|
||||||
|
self.assertEqual(proxy._acknowledgement_range(1), (3, 4))
|
||||||
|
self.assertEqual(proxy._acknowledgement_range(2), (2, 4))
|
||||||
|
self.assertEqual(proxy._acknowledgement_range(4), (0, 4))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -259,6 +259,32 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|||||||
with self.assertRaisesRegex(ValueError, "single node"):
|
with self.assertRaisesRegex(ValueError, "single node"):
|
||||||
server_args._handle_multimodal_feature_transport()
|
server_args._handle_multimodal_feature_transport()
|
||||||
|
|
||||||
|
@patch("sglang.srt.server_args.is_cuda", return_value=False)
|
||||||
|
def test_cuda_vmm_rejects_non_nvidia_platforms(self, _mock_is_cuda):
|
||||||
|
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"):
|
||||||
|
server_args._handle_multimodal_feature_transport()
|
||||||
|
|
||||||
|
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
||||||
|
def test_cuda_vmm_rejects_rust_server(self, _mock_is_cuda):
|
||||||
|
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
|
||||||
|
|
||||||
|
with (
|
||||||
|
envs.SGLANG_RUST_SERVER.override(True),
|
||||||
|
self.assertRaisesRegex(ValueError, "SGLANG_RUST_SERVER"),
|
||||||
|
):
|
||||||
|
server_args._handle_multimodal_feature_transport()
|
||||||
|
|
||||||
|
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
||||||
|
def test_cuda_vmm_rejects_pipeline_parallelism(self, _mock_is_cuda):
|
||||||
|
server_args = ServerArgs(
|
||||||
|
model_path="dummy", mm_feature_transport="cuda_vmm", pp_size=2
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "pipeline parallelism"):
|
||||||
|
server_args._handle_multimodal_feature_transport()
|
||||||
|
|
||||||
|
|
||||||
class TestMambaCacheStochasticRounding(unittest.TestCase):
|
class TestMambaCacheStochasticRounding(unittest.TestCase):
|
||||||
def test_rejects_fp32_ssm_cache(self):
|
def test_rejects_fp32_ssm_cache(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user