Allow CUDA VMM feature transport with the Rust frontend (#39347)

This commit is contained in:
Lianmin Zheng
2026-09-14 17:09:49 -07:00
committed by GitHub
parent 6e755e4114
commit a2b4e8888f
2 changed files with 6 additions and 7 deletions
@@ -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(