Allow CUDA VMM feature transport with the Rust frontend (#39347)
This commit is contained in:
@@ -851,11 +851,6 @@ def handle_multimodal_feature_transport(server_args: Any):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--mm-feature-transport=cuda_vmm does not support pipeline parallelism."
|
"--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."
|
|
||||||
)
|
|
||||||
pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get()
|
pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get()
|
||||||
handle_kind = "CUDA FABRIC" if cfg.nnodes > 1 else "POSIX FD"
|
handle_kind = "CUDA FABRIC" if cfg.nnodes > 1 else "POSIX FD"
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -686,15 +686,19 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|||||||
handle_multimodal_feature_transport(server_args)
|
handle_multimodal_feature_transport(server_args)
|
||||||
|
|
||||||
@override_platform(is_cuda=True)
|
@override_platform(is_cuda=True)
|
||||||
def test_cuda_vmm_rejects_rust_server(self):
|
def test_cuda_vmm_allows_rust_server(self):
|
||||||
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
|
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
|
||||||
|
|
||||||
with (
|
with (
|
||||||
|
patch.dict(os.environ, {}, clear=False),
|
||||||
envs.SGLANG_RUST_SERVER.override(True),
|
envs.SGLANG_RUST_SERVER.override(True),
|
||||||
self.assertRaisesRegex(ValueError, "SGLANG_RUST_SERVER"),
|
|
||||||
):
|
):
|
||||||
handle_multimodal_feature_transport(server_args)
|
handle_multimodal_feature_transport(server_args)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
|
||||||
|
)
|
||||||
|
|
||||||
@override_platform(is_cuda=True)
|
@override_platform(is_cuda=True)
|
||||||
def test_cuda_vmm_rejects_pipeline_parallelism(self):
|
def test_cuda_vmm_rejects_pipeline_parallelism(self):
|
||||||
server_args = ServerArgs(
|
server_args = ServerArgs(
|
||||||
|
|||||||
Reference in New Issue
Block a user