|
|
|
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch
|
|
|
|
|
|
|
|
|
|
import sglang.srt.server_args as server_args_module
|
|
|
|
|
from sglang.srt.arg_groups import pd_disaggregation_hook
|
|
|
|
|
from sglang.srt.arg_groups.overrides import resolution_result
|
|
|
|
|
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
|
|
|
|
|
from sglang.srt.entrypoints.sidecar import (
|
|
|
|
|
SGLANG_GRPC_ENDPOINT_ENV,
|
|
|
|
@@ -61,7 +62,7 @@ class TestPrepareServerArgs(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
args.resolve_once()
|
|
|
|
|
|
|
|
|
|
self.assertTrue(args.enable_w4a4_mxfp4_megamoe)
|
|
|
|
|
self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
|
|
|
|
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1")
|
|
|
|
|
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1")
|
|
|
|
|
|
|
|
|
@@ -76,14 +77,14 @@ class TestPrepareServerArgs(CustomTestCase):
|
|
|
|
|
# nothing to be untouched by.
|
|
|
|
|
args.resolve_once()
|
|
|
|
|
|
|
|
|
|
self.assertFalse(args.enable_w4a4_mxfp4_megamoe)
|
|
|
|
|
self.assertFalse(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
|
|
|
|
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
|
|
|
|
|
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0")
|
|
|
|
|
|
|
|
|
|
def test_prefill_decode_interval(self):
|
|
|
|
|
args = ServerArgs(model_path="dummy", prefill_decode_interval=16)
|
|
|
|
|
args.resolve_once()
|
|
|
|
|
self.assertEqual(args.prefill_decode_interval, 16)
|
|
|
|
|
self.assertEqual(resolution_result(args, "prefill_decode_interval"), 16)
|
|
|
|
|
|
|
|
|
|
with self.assertRaisesRegex(
|
|
|
|
|
ValueError, "--prefill-decode-interval must be non-negative"
|
|
|
|
@@ -114,22 +115,24 @@ class TestPrepareServerArgs(CustomTestCase):
|
|
|
|
|
return server_args
|
|
|
|
|
|
|
|
|
|
disabled = _resolved(model_path="dummy")
|
|
|
|
|
self.assertFalse(disabled.enable_return_hidden_states)
|
|
|
|
|
self.assertIsNone(disabled.return_hidden_states_mode)
|
|
|
|
|
self.assertFalse(resolution_result(disabled, "enable_return_hidden_states"))
|
|
|
|
|
self.assertIsNone(resolution_result(disabled, "return_hidden_states_mode"))
|
|
|
|
|
|
|
|
|
|
last = _resolved(
|
|
|
|
|
model_path="dummy",
|
|
|
|
|
return_hidden_states_mode="last",
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(last.enable_return_hidden_states)
|
|
|
|
|
self.assertEqual(last.return_hidden_states_mode, "last")
|
|
|
|
|
self.assertTrue(resolution_result(last, "enable_return_hidden_states"))
|
|
|
|
|
self.assertEqual(resolution_result(last, "return_hidden_states_mode"), "last")
|
|
|
|
|
|
|
|
|
|
legacy_full = _resolved(
|
|
|
|
|
model_path="dummy",
|
|
|
|
|
enable_return_hidden_states=True,
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(legacy_full.enable_return_hidden_states)
|
|
|
|
|
self.assertEqual(legacy_full.return_hidden_states_mode, "full")
|
|
|
|
|
self.assertTrue(resolution_result(legacy_full, "enable_return_hidden_states"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(legacy_full, "return_hidden_states_mode"), "full"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
parsed_last = prepare_server_args(
|
|
|
|
|
[
|
|
|
|
@@ -140,8 +143,10 @@ class TestPrepareServerArgs(CustomTestCase):
|
|
|
|
|
]
|
|
|
|
|
)
|
|
|
|
|
parsed_last.resolve_once()
|
|
|
|
|
self.assertTrue(parsed_last.enable_return_hidden_states)
|
|
|
|
|
self.assertEqual(parsed_last.return_hidden_states_mode, "last")
|
|
|
|
|
self.assertTrue(resolution_result(parsed_last, "enable_return_hidden_states"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(parsed_last, "return_hidden_states_mode"), "last"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# The rejection is resolution's, not the constructor's.
|
|
|
|
|
with self.assertRaisesRegex(
|
|
|
|
@@ -156,13 +161,24 @@ class TestPrepareServerArgs(CustomTestCase):
|
|
|
|
|
def test_draft_quantization_explicitness_survives_asdict_round_trip(self):
|
|
|
|
|
inherited = ServerArgs(model_path="dummy", quantization="modelopt_fp4")
|
|
|
|
|
inherited._handle_missing_default_values()
|
|
|
|
|
self.assertEqual(inherited.speculative_draft_model_quantization, "modelopt_fp4")
|
|
|
|
|
self.assertFalse(inherited._speculative_draft_quantization_explicitly_set)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(inherited, "speculative_draft_model_quantization"),
|
|
|
|
|
"modelopt_fp4",
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(
|
|
|
|
|
resolution_result(
|
|
|
|
|
inherited, "_speculative_draft_quantization_explicitly_set"
|
|
|
|
|
)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
reconstructed = ServerArgs(**dataclasses.asdict(inherited))
|
|
|
|
|
reconstructed._handle_missing_default_values()
|
|
|
|
|
|
|
|
|
|
self.assertFalse(reconstructed._speculative_draft_quantization_explicitly_set)
|
|
|
|
|
self.assertFalse(
|
|
|
|
|
resolution_result(
|
|
|
|
|
reconstructed, "_speculative_draft_quantization_explicitly_set"
|
|
|
|
|
)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_config_nested_dict_args_are_json(self):
|
|
|
|
|
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
|
|
|
|
@@ -219,8 +235,10 @@ class TestImageProcessorBackend(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
server_args._handle_deprecated_args()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.image_processor_backend, "pil")
|
|
|
|
|
self.assertFalse(server_args.disable_fast_image_processor)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "image_processor_backend"), "pil"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(resolution_result(server_args, "disable_fast_image_processor"))
|
|
|
|
|
|
|
|
|
|
def test_legacy_flag_maps_to_pil_with_one_warning(self):
|
|
|
|
|
server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True)
|
|
|
|
@@ -228,8 +246,10 @@ class TestImageProcessorBackend(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
|
|
|
|
server_args._handle_deprecated_args()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.image_processor_backend, "pil")
|
|
|
|
|
self.assertTrue(server_args.disable_fast_image_processor)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "image_processor_backend"), "pil"
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "disable_fast_image_processor"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
sum(
|
|
|
|
|
"--disable-fast-image-processor is deprecated" in x for x in logs.output
|
|
|
|
@@ -266,7 +286,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
output = "\n".join(logs.output)
|
|
|
|
@@ -281,8 +303,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
|
|
|
|
|
self.assertFalse(server_args.keep_mm_feature_on_device)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(
|
|
|
|
|
resolution_result(server_args, "keep_mm_feature_on_device")
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
self.assertIn("deprecated", logs.output[0])
|
|
|
|
@@ -305,7 +331,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
self.assertIn("overrides", logs.output[0])
|
|
|
|
@@ -316,7 +344,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
|
|
|
@@ -329,7 +359,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertNoLogs(server_args_module.logger, level="INFO"):
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
|
|
|
@@ -342,7 +374,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertNoLogs(server_args_module.logger, level="INFO"):
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.os.path.exists", return_value=True)
|
|
|
|
@@ -367,7 +401,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
output = "\n".join(logs.output)
|
|
|
|
@@ -394,7 +430,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu")
|
|
|
|
|
self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output))
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.os.path.exists", return_value=False)
|
|
|
|
@@ -411,7 +447,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.assertIn("no IMEX channel", "\n".join(logs.output))
|
|
|
|
|
|
|
|
|
@@ -427,7 +465,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.is_cuda", return_value=True)
|
|
|
|
@@ -439,7 +479,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cpu")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cpu"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.server_args.is_cuda", return_value=False)
|
|
|
|
@@ -474,7 +516,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
|
|
|
|
|
server_args._handle_multimodal_feature_transport()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
|
|
|
|
|
|
|
|
|
|
output = "\n".join(logs.output)
|
|
|
|
@@ -555,15 +599,22 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
|
|
|
|
|
|
|
|
|
def test_non_pd_defaults_to_round_robin(self):
|
|
|
|
|
server_args = self._load_balance_args(disaggregation_mode="null")
|
|
|
|
|
self.assertEqual(server_args.load_balance_method, "round_robin")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "load_balance_method"), "round_robin"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_pd_prefill_defaults_to_follow_bootstrap_room(self):
|
|
|
|
|
server_args = self._load_balance_args(disaggregation_mode="prefill")
|
|
|
|
|
self.assertEqual(server_args.load_balance_method, "follow_bootstrap_room")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "load_balance_method"),
|
|
|
|
|
"follow_bootstrap_room",
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_pd_decode_defaults_to_round_robin(self):
|
|
|
|
|
server_args = self._load_balance_args(disaggregation_mode="decode")
|
|
|
|
|
self.assertEqual(server_args.load_balance_method, "round_robin")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "load_balance_method"), "round_robin"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_pd_prefill_dcp_warns_about_performance(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -581,7 +632,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
|
|
|
|
disaggregation_transfer_backend="mooncake",
|
|
|
|
|
dcp_size=4,
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(server_args.disable_radix_cache)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
|
|
|
|
|
|
|
|
|
|
def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -601,7 +652,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
|
|
|
|
disaggregation_transfer_backend="fake",
|
|
|
|
|
dcp_size=4,
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(server_args.disable_radix_cache)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
|
|
|
|
|
|
|
|
|
|
def test_pd_decode_dcp_rejects_radix_cache(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -665,8 +716,11 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
|
|
|
|
disaggregation_transfer_backend="mooncake_tcp",
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.assertFalse(server_args.disable_radix_cache)
|
|
|
|
|
self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake")
|
|
|
|
|
self.assertFalse(resolution_result(server_args, "disable_radix_cache"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "disaggregation_transfer_backend"),
|
|
|
|
|
"mooncake",
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestSkipTokenizerInit(unittest.TestCase):
|
|
|
|
@@ -681,8 +735,8 @@ class TestSkipTokenizerInit(unittest.TestCase):
|
|
|
|
|
server_args._handle_tokenizer_batching()
|
|
|
|
|
|
|
|
|
|
# Tokenizer fanout preserved; detokenizer coerced to 1 (no decode work).
|
|
|
|
|
self.assertEqual(server_args.tokenizer_worker_num, 4)
|
|
|
|
|
self.assertEqual(server_args.detokenizer_worker_num, 1)
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "tokenizer_worker_num"), 4)
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "detokenizer_worker_num"), 1)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
|
|
|
|
@@ -870,7 +924,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
from sglang.srt.arg_groups.overrides import resolved_view
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
|
|
|
|
self.assertEqual(args.page_size, 1) # the field stays pristine
|
|
|
|
|
self.assertEqual(resolved_view(args).page_size, 128)
|
|
|
|
|
|
|
|
|
|
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
|
|
|
|
@@ -883,7 +937,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
from sglang.srt.arg_groups.overrides import resolved_view
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
|
|
|
|
self.assertEqual(args.page_size, 1) # the field stays pristine
|
|
|
|
|
self.assertEqual(resolved_view(args).page_size, 128)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -918,12 +972,12 @@ class TestContextParallelServerArgs(CustomTestCase):
|
|
|
|
|
def test_canonical_prefill_cp_requires_strategy(self):
|
|
|
|
|
args = self.parser.parse_args(["--model", "dummy", "--enable-prefill-cp"])
|
|
|
|
|
|
|
|
|
|
self.assertTrue(args.enable_prefill_cp)
|
|
|
|
|
self.assertIsNone(args.cp_strategy)
|
|
|
|
|
self.assertTrue(resolution_result(args, "enable_prefill_cp"))
|
|
|
|
|
self.assertIsNone(resolution_result(args, "cp_strategy"))
|
|
|
|
|
|
|
|
|
|
server_args = self._new_cp_args(
|
|
|
|
|
enable_prefill_cp=args.enable_prefill_cp,
|
|
|
|
|
cp_strategy=args.cp_strategy,
|
|
|
|
|
enable_prefill_cp=resolution_result(args, "enable_prefill_cp"),
|
|
|
|
|
cp_strategy=resolution_result(args, "cp_strategy"),
|
|
|
|
|
)
|
|
|
|
|
with self.assertRaisesRegex(ValueError, "--cp-strategy"):
|
|
|
|
|
server_args._handle_context_parallelism()
|
|
|
|
@@ -940,16 +994,18 @@ class TestContextParallelServerArgs(CustomTestCase):
|
|
|
|
|
)
|
|
|
|
|
server_args = self._new_cp_args(
|
|
|
|
|
enable_dsa_prefill_context_parallel=(
|
|
|
|
|
args.enable_dsa_prefill_context_parallel
|
|
|
|
|
resolution_result(args, "enable_dsa_prefill_context_parallel")
|
|
|
|
|
),
|
|
|
|
|
dsa_prefill_cp_mode=args.dsa_prefill_cp_mode,
|
|
|
|
|
dsa_prefill_cp_mode=resolution_result(args, "dsa_prefill_cp_mode"),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
server_args._handle_legacy_cp_arguments()
|
|
|
|
|
|
|
|
|
|
self.assertTrue(server_args.enable_prefill_cp)
|
|
|
|
|
self.assertEqual(server_args.cp_strategy, "interleave")
|
|
|
|
|
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "cp_strategy"), "interleave")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self):
|
|
|
|
|
server_args = self._new_cp_args(
|
|
|
|
@@ -961,10 +1017,18 @@ class TestContextParallelServerArgs(CustomTestCase):
|
|
|
|
|
server_args._handle_legacy_cp_arguments()
|
|
|
|
|
server_args._handle_context_parallelism()
|
|
|
|
|
|
|
|
|
|
self.assertTrue(server_args.enable_dsa_prefill_context_parallel)
|
|
|
|
|
self.assertFalse(server_args.enable_prefill_context_parallel)
|
|
|
|
|
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
|
|
|
|
|
self.assertEqual(server_args.prefill_cp_mode, "round-robin-split")
|
|
|
|
|
self.assertTrue(
|
|
|
|
|
resolution_result(server_args, "enable_dsa_prefill_context_parallel")
|
|
|
|
|
)
|
|
|
|
|
self.assertFalse(
|
|
|
|
|
resolution_result(server_args, "enable_prefill_context_parallel")
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "prefill_cp_mode"), "round-robin-split"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_context_parallel_handler_initializes_cp_strategy(self):
|
|
|
|
|
server_args = self._new_cp_args(
|
|
|
|
@@ -1049,15 +1113,25 @@ class TestContextParallelServerArgs(CustomTestCase):
|
|
|
|
|
server_args._handle_legacy_cp_arguments()
|
|
|
|
|
server_args._handle_context_parallelism()
|
|
|
|
|
|
|
|
|
|
self.assertTrue(server_args.enable_prefill_cp)
|
|
|
|
|
self.assertEqual(server_args.cp_strategy, strategy)
|
|
|
|
|
self.assertEqual(server_args.dsa_prefill_cp_mode, mode)
|
|
|
|
|
self.assertEqual(server_args.prefill_cp_mode, mode)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
server_args.enable_dsa_prefill_context_parallel, expect_dsa
|
|
|
|
|
resolution_result(server_args, "cp_strategy"), strategy
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
server_args.enable_prefill_context_parallel, expect_generic
|
|
|
|
|
resolution_result(server_args, "dsa_prefill_cp_mode"), mode
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "prefill_cp_mode"), mode
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(
|
|
|
|
|
server_args, "enable_dsa_prefill_context_parallel"
|
|
|
|
|
),
|
|
|
|
|
expect_dsa,
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "enable_prefill_context_parallel"),
|
|
|
|
|
expect_generic,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -1308,7 +1382,7 @@ class TestSSLArgs(unittest.TestCase):
|
|
|
|
|
ssl_certfile="cert.pem",
|
|
|
|
|
enable_ssl_refresh=True,
|
|
|
|
|
)
|
|
|
|
|
self.assertTrue(server_args.enable_ssl_refresh)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enable_ssl_refresh"))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestHiCacheArgs(unittest.TestCase):
|
|
|
|
@@ -1328,10 +1402,17 @@ class TestHiCacheArgs(unittest.TestCase):
|
|
|
|
|
expected_mem_layout: str,
|
|
|
|
|
expected_decode_backend: str | None = None,
|
|
|
|
|
):
|
|
|
|
|
self.assertEqual(args.hicache_io_backend, expected_io_backend)
|
|
|
|
|
self.assertEqual(args.hicache_mem_layout, expected_mem_layout)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "hicache_io_backend"), expected_io_backend
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "hicache_mem_layout"), expected_mem_layout
|
|
|
|
|
)
|
|
|
|
|
if expected_decode_backend is not None:
|
|
|
|
|
self.assertEqual(args.decode_attention_backend, expected_decode_backend)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "decode_attention_backend"),
|
|
|
|
|
expected_decode_backend,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_hicache_io_backend_and_mem_layout_compatibility(self):
|
|
|
|
|
cases = [
|
|
|
|
@@ -1409,9 +1490,9 @@ class TestHiCacheArgs(unittest.TestCase):
|
|
|
|
|
)
|
|
|
|
|
args._handle_hicache()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.hicache_io_backend, "kernel")
|
|
|
|
|
self.assertEqual(args.hicache_mem_layout, "page_first")
|
|
|
|
|
self.assertIsNone(args.decode_attention_backend)
|
|
|
|
|
self.assertEqual(resolution_result(args, "hicache_io_backend"), "kernel")
|
|
|
|
|
self.assertEqual(resolution_result(args, "hicache_mem_layout"), "page_first")
|
|
|
|
|
self.assertIsNone(resolution_result(args, "decode_attention_backend"))
|
|
|
|
|
|
|
|
|
|
def test_decode_offload_rejects_host_pool_retraction(self):
|
|
|
|
|
args = self._make_args(
|
|
|
|
@@ -1494,11 +1575,19 @@ class TestDecoupledSpecArgs(CustomTestCase):
|
|
|
|
|
"/tmp/tr",
|
|
|
|
|
]
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(server_args.decoupled_spec_role, "verifier")
|
|
|
|
|
self.assertEqual(server_args.decoupled_spec_bind_endpoint, "ipc:///tmp/v")
|
|
|
|
|
self.assertEqual(server_args.decoupled_spec_connect_endpoints, ["ipc:///tmp/d"])
|
|
|
|
|
self.assertEqual(server_args.decoupled_spec_rank, 0)
|
|
|
|
|
self.assertEqual(server_args.spec_trace_dir, "/tmp/tr")
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "decoupled_spec_role"), "verifier"
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "decoupled_spec_bind_endpoint"),
|
|
|
|
|
"ipc:///tmp/v",
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(server_args, "decoupled_spec_connect_endpoints"),
|
|
|
|
|
["ipc:///tmp/d"],
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "decoupled_spec_rank"), 0)
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "spec_trace_dir"), "/tmp/tr")
|
|
|
|
|
|
|
|
|
|
def test_decoupled_spec_role_rejects_invalid_choice(self):
|
|
|
|
|
with self.assertRaises(SystemExit):
|
|
|
|
@@ -1533,10 +1622,10 @@ class TestAdaptiveSpecArgs(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
handle_speculative_decoding(args)
|
|
|
|
|
|
|
|
|
|
self.assertTrue(args.speculative_adaptive)
|
|
|
|
|
self.assertEqual(args.speculative_eagle_topk, 1)
|
|
|
|
|
self.assertEqual(args.speculative_num_steps, 3)
|
|
|
|
|
self.assertEqual(args.speculative_num_draft_tokens, 4)
|
|
|
|
|
self.assertTrue(resolution_result(args, "speculative_adaptive"))
|
|
|
|
|
self.assertEqual(resolution_result(args, "speculative_eagle_topk"), 1)
|
|
|
|
|
self.assertEqual(resolution_result(args, "speculative_num_steps"), 3)
|
|
|
|
|
self.assertEqual(resolution_result(args, "speculative_num_draft_tokens"), 4)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestWaterfillArgs(CustomTestCase):
|
|
|
|
@@ -1552,10 +1641,9 @@ class TestWaterfillArgs(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
from sglang.srt.arg_groups.overrides import resolved_view
|
|
|
|
|
|
|
|
|
|
# dual-apply retired: the fields stay pristine, the declarations win
|
|
|
|
|
self.assertTrue(server_args.disable_shared_experts_fusion)
|
|
|
|
|
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
|
|
|
|
|
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
|
|
|
|
|
|
|
|
|
def test_waterfill_overrides_moe_a2a_backend_to_deepep(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -1570,7 +1658,7 @@ class TestWaterfillArgs(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.moe_a2a_backend, "none") # pristine
|
|
|
|
|
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "deepep")
|
|
|
|
|
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
|
|
|
|
|
|
|
|
|
def test_waterfill_keeps_megamoe_backend(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -1586,7 +1674,7 @@ class TestWaterfillArgs(CustomTestCase):
|
|
|
|
|
|
|
|
|
|
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "megamoe")
|
|
|
|
|
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
|
|
|
|
|
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
|
|
|
|
|
|
|
|
|
def test_waterfill_supports_deepep_low_latency_mode(self):
|
|
|
|
|
server_args = ServerArgs(
|
|
|
|
@@ -1598,9 +1686,9 @@ class TestWaterfillArgs(CustomTestCase):
|
|
|
|
|
# dummy-model path short-circuits __post_init__; invoke the handler directly.
|
|
|
|
|
server_args._handle_a2a_moe()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(server_args.deepep_mode, "low_latency")
|
|
|
|
|
self.assertFalse(server_args.disable_cuda_graph)
|
|
|
|
|
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
|
|
|
|
self.assertEqual(resolution_result(server_args, "deepep_mode"), "low_latency")
|
|
|
|
|
self.assertFalse(resolution_result(server_args, "disable_cuda_graph"))
|
|
|
|
|
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestPrefillOnlyDisableKvCache(unittest.TestCase):
|
|
|
|
@@ -1635,7 +1723,7 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase):
|
|
|
|
|
|
|
|
|
|
def test_valid_minimal_config_constructs(self):
|
|
|
|
|
sa = self._validate_prefill_only_args()
|
|
|
|
|
self.assertTrue(sa.prefill_only_disable_kv_cache)
|
|
|
|
|
self.assertTrue(resolution_result(sa, "prefill_only_disable_kv_cache"))
|
|
|
|
|
|
|
|
|
|
def test_rejects_when_not_embedding(self):
|
|
|
|
|
with self.assertRaisesRegex(ValueError, "requires --is-embedding"):
|
|
|
|
@@ -1725,15 +1813,27 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
|
|
|
|
def test_cuda_graph_prefill_role_defaults_disable_decode_graph(self):
|
|
|
|
|
args = self._handled_args(disaggregation_mode="prefill")
|
|
|
|
|
|
|
|
|
|
self.assertFalse(args.disable_cuda_graph)
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
|
|
|
|
self.assertFalse(resolution_result(args, "disable_cuda_graph"))
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").decode.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.BREAKABLE,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_cuda_graph_decode_role_defaults_disable_prefill_graph(self):
|
|
|
|
|
args = self._handled_args(disaggregation_mode="decode")
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
|
|
|
|
self.assertNotEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
self.assertNotEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").decode.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_cuda_graph_global_disable_still_disables_both_phases_for_all_roles(self):
|
|
|
|
|
for disaggregation_mode in ("prefill", "decode", "null"):
|
|
|
|
@@ -1744,10 +1844,12 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
args.cuda_graph_config.decode.backend, Backend.DISABLED
|
|
|
|
|
resolution_result(args, "cuda_graph_config").decode.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
args.cuda_graph_config.prefill.backend, Backend.DISABLED
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self):
|
|
|
|
@@ -1756,7 +1858,9 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
|
|
|
|
|
cuda_graph_backend_decode=Backend.FULL,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").decode.backend, Backend.FULL
|
|
|
|
|
)
|
|
|
|
|
self.assertIn((Phase.DECODE, "backend"), args._cuda_graph_config_locked)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -1782,12 +1886,18 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
|
|
|
|
|
def test_enable_lora_keeps_breakable_prefill_graph(self):
|
|
|
|
|
args = self._handled_args(enable_lora=True)
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.BREAKABLE,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_lora_paths_keep_breakable_prefill_graph(self):
|
|
|
|
|
args = self._handled_args(lora_paths=["dummy/lora-adapter"])
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.BREAKABLE,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_lora_still_disables_tc_piecewise_prefill_graph(self):
|
|
|
|
|
# Pin the tc_piecewise LoRA rule itself, with the hardware rule
|
|
|
|
@@ -1811,7 +1921,10 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
|
|
|
|
|
):
|
|
|
|
|
args._disable_tc_piecewise_cudagraph_if_incompatible()
|
|
|
|
|
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
|
|
|
@@ -1840,7 +1953,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
|
|
|
|
is_multimodal=True,
|
|
|
|
|
allowlisted=False,
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.DISABLED,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_allowlisted_multimodal_arch_keeps_prefill_breakable(self):
|
|
|
|
|
args = self._handled_args(
|
|
|
|
@@ -1848,7 +1964,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
|
|
|
|
|
is_multimodal=True,
|
|
|
|
|
allowlisted=True,
|
|
|
|
|
)
|
|
|
|
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
|
|
|
|
|
self.assertEqual(
|
|
|
|
|
resolution_result(args, "cuda_graph_config").prefill.backend,
|
|
|
|
|
Backend.BREAKABLE,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def test_allowlist_membership(self):
|
|
|
|
|
from sglang.srt.configs.model_config import (
|
|
|
|
@@ -2046,20 +2165,20 @@ class TestGrpcServerArgs(CustomTestCase):
|
|
|
|
|
def test_http_only_high_port_does_not_derive_grpc_port(self):
|
|
|
|
|
sa = self._args(port=56000)
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertIsNone(sa.grpc_port)
|
|
|
|
|
self.assertIsNone(resolution_result(sa, "grpc_port"))
|
|
|
|
|
|
|
|
|
|
def test_grpc_port_enables_native_and_env_knobs(self):
|
|
|
|
|
sa = self._args(grpc_port=50051)
|
|
|
|
|
with envs.SGLANG_GRPC_WORKER_THREADS.override(8):
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertEqual(sa.grpc_port, 50051)
|
|
|
|
|
self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
|
|
|
|
|
self.assertEqual(sa.grpc_worker_threads, 8)
|
|
|
|
|
|
|
|
|
|
def test_env_grpc_port_enables_native(self):
|
|
|
|
|
sa = self._args(port=30000)
|
|
|
|
|
with envs.SGLANG_GRPC_PORT.override(45000):
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertEqual(sa.grpc_port, 45000)
|
|
|
|
|
self.assertEqual(resolution_result(sa, "grpc_port"), 45000)
|
|
|
|
|
|
|
|
|
|
@staticmethod
|
|
|
|
|
def _sidecar_parser():
|
|
|
|
@@ -2196,20 +2315,20 @@ class TestGrpcServerArgs(CustomTestCase):
|
|
|
|
|
def test_legacy_smg_derives_grpc_port_from_http_port(self):
|
|
|
|
|
sa = self._args(port=30000, smg_grpc_mode=True)
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertEqual(sa.grpc_port, 40000)
|
|
|
|
|
self.assertEqual(resolution_result(sa, "grpc_port"), 40000)
|
|
|
|
|
|
|
|
|
|
def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self):
|
|
|
|
|
sa = self._args(grpc_mode=True)
|
|
|
|
|
with self.assertLogs(server_args_module.logger, level="WARNING") as cm:
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertTrue(sa.smg_grpc_mode)
|
|
|
|
|
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
|
|
|
|
|
self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output))
|
|
|
|
|
|
|
|
|
|
def test_legacy_smg_takes_precedence_over_grpc_port(self):
|
|
|
|
|
sa = self._args(grpc_port=50051, smg_grpc_mode=True)
|
|
|
|
|
sa._handle_deprecated_args()
|
|
|
|
|
self.assertTrue(sa.smg_grpc_mode)
|
|
|
|
|
self.assertEqual(sa.grpc_port, 50051)
|
|
|
|
|
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
|
|
|
|
|
self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
|
|
|
|
|
|
|
|
|
|
def test_native_grpc_rejects_multi_tokenizer(self):
|
|
|
|
|
sa = self._args(grpc_port=40000, tokenizer_worker_num=2)
|
|
|
|
@@ -2254,7 +2373,7 @@ class TestGrpcServerArgs(CustomTestCase):
|
|
|
|
|
tokenizer_manager=MagicMock(),
|
|
|
|
|
template_manager=MagicMock(),
|
|
|
|
|
scheduler_info={},
|
|
|
|
|
grpc_port=server_args.grpc_port,
|
|
|
|
|
grpc_port=resolution_result(server_args, "grpc_port"),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.assertEqual(handle, "handle")
|
|
|
|
|