config: the resolution pipeline moves out of the record (#36789)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-28 10:17:24 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 726665e08e
commit c2928e86d7
30 changed files with 6863 additions and 5639 deletions
@@ -37,6 +37,20 @@ _READ_BEFORE_RESOLUTION = frozenset({"is_embedding"})
# has to be looked at.
_STALE_IN_THE_MODEL_CONFIG = frozenset({"speculative_algorithm"})
# Behind the expert-pack build. `expert_pack_hook.handle_expert_pack` builds a
# model configuration, and it always did -- the walk stopped at the record's
# file and never saw it, so these three read as decided before the first build.
# The call sits behind `load_format != "expert_pack": return`, so it is the
# first build only on an expert-pack launch. Pre-existing; named rather than
# fixed, because fixing it means moving the build or the hook.
_STALE_BEHIND_THE_EXPERT_PACK_BUILD = frozenset(
{
"_speculative_draft_quantization_explicitly_set",
"model_path",
"speculative_draft_model_quantization",
}
)
# The same staleness through the registries: `_handle_model_specific_adjustments`
# builds the model configuration and *then* collects the override declarations,
# both inside one handler body. Named rather than fixed (that means moving the
@@ -98,13 +112,27 @@ def _registry_collection_is_after_the_build():
collection above this handler's own `get_model_config()` call does not move
it above the configuration another handler already cached.
"""
tree = _parsed(_SRT / "server_args.py")
handler = next(
node
for node in ast.walk(tree)
if isinstance(node, ast.FunctionDef)
and node.name == "_handle_model_specific_adjustments"
)
handler = None
for source, wanted in (
(_SRT / "server_args.py", "_handle_model_specific_adjustments"),
*(
(path, "handle_model_specific_adjustments")
for path in sorted((_SRT / "arg_groups").glob("*.py"))
),
):
for node in ast.walk(_parsed(source)):
if isinstance(node, ast.FunctionDef) and node.name == wanted:
if any(
isinstance(child, ast.Call)
and getattr(child.func, "attr", getattr(child.func, "id", None))
== "collect_model_override_declarations"
for child in ast.walk(node)
):
handler = node
break
if handler is not None:
break
assert handler is not None, "the model-specific handler was not found"
build = collect = None
for node in ast.walk(handler):
if not isinstance(node, ast.Call):
@@ -283,6 +311,21 @@ def _hook_declarations(dispatch, source_module):
return out
def _hook_functions():
"""Module-level resolution functions under `arg_groups/`.
A handler that moved out of the record leaves a slot behind that imports
one of these and calls it. Without following that hop the scan stops at
the slot and silently loses everything the handler does.
"""
functions = {}
for path in sorted((_SRT / "arg_groups").glob("*.py")):
for node in _parsed(path).body:
if isinstance(node, ast.FunctionDef):
functions.setdefault(node.name, node)
return functions
def _pipeline():
"""(ordered steps, {step: methods it reaches}) for the resolution dispatch."""
tree = _parsed(_SRT / "server_args.py")
@@ -294,6 +337,28 @@ def _pipeline():
methods = {
node.name: node for node in record.body if isinstance(node, ast.FunctionDef)
}
hooks = _hook_functions()
# Follow exactly one edge: the slot's own `from arg_groups.X import f` /
# `f(self)`. Merging every hook function by bare name would let the walk
# wander into families the slot never calls.
slot_target = {}
for name, node in methods.items():
imported = {
alias.asname or alias.name
for child in ast.walk(node)
if isinstance(child, ast.ImportFrom)
and child.module
and child.module.startswith("sglang.srt.arg_groups")
for alias in child.names
}
called = {
child.func.id
for child in ast.walk(node)
if isinstance(child, ast.Call) and isinstance(child.func, ast.Name)
}
for target in sorted(imported & called & set(hooks)):
slot_target.setdefault(name, target)
methods.update({name: hooks[name] for name in slot_target.values()})
dispatch = methods["_run_resolution_pipeline"]
steps = [
name
@@ -313,14 +378,18 @@ def _pipeline():
return seen
seen.add(name)
for node in ast.walk(methods[name]):
if not isinstance(node, ast.Call):
continue
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
isinstance(node.func, ast.Attribute)
and isinstance(node.func.value, ast.Name)
and node.func.value.id == "self"
and node.func.attr in methods
):
reaches(node.func.attr, seen)
target = slot_target.get(name)
if target is not None:
reaches(target, seen)
return seen
step_lines = {}
@@ -512,6 +581,7 @@ class TestModelConfigReadsResolvedInput(CustomTestCase):
_READ_BEFORE_RESOLUTION
| _STALE_IN_THE_MODEL_CONFIG
| _STALE_FROM_THE_REGISTRIES
| _STALE_BEHIND_THE_EXPERT_PACK_BUILD
)
late = sorted(
field
@@ -520,8 +520,11 @@ class TestProgramsResolveBeforeReadingResolution(CustomTestCase):
declarers = {"_declare", "declare_resolution", "declare_late_resolution"}
fields = set()
field_names = {field.name for field in _dataclasses.fields(_ServerArgs)}
for name in ("server_args.py", "arg_groups/overrides.py"):
tree = ast.parse((srt / name).read_text(encoding="utf-8-sig"))
# The record plus every module under `arg_groups/`: a handler declares
# from whichever of the two it lives in.
sources = [srt / "server_args.py", *sorted((srt / "arg_groups").rglob("*.py"))]
for source in sources:
tree = ast.parse(source.read_text(encoding="utf-8-sig"))
for node in ast.walk(tree):
# Registry data: provider dict keys are field names as
# *data*, invisible to the keyword scan below. Filtered
@@ -72,6 +72,15 @@ _ATTRIBUTE_SPELLED = _BAG_ACCESSORS - {"get_device"}
_OWN = ("server_args.py", "runtime_context.py")
def _pipeline_sources():
"""The record plus every module under `arg_groups/`.
A handler that moved out of the record takes its imports with it, so
seeding the walk from two files would stop covering it.
"""
return [_SRT / "server_args.py", *sorted((_SRT / "arg_groups").rglob("*.py"))]
def _module_of(name):
"""`sglang.srt.a.b` -> the file, if it is one of ours."""
if not name or not name.startswith("sglang.srt."):
@@ -196,9 +205,29 @@ def _functions_in(path):
}
def _locally_shadowed_accessors(path):
"""Accessor names this file imports from somewhere that is not the context.
`get_device` is both the `device` bag accessor and the hardware probe in
`utils.common`. Matching the bare name would report the probe as a bag read,
so a name imported from elsewhere in this file is not the accessor.
"""
shadowed = set()
for node in ast.walk(ast.parse(path.read_text(encoding="utf-8-sig"))):
if isinstance(node, ast.ImportFrom) and node.module:
if node.module.endswith("runtime_context"):
continue
for alias in node.names:
name = alias.asname or alias.name
if name in _BAG_ACCESSORS:
shadowed.add(name)
return shadowed
def _reaches_a_bag(path, entry):
"""Does `entry` in `path` reach a bag accessor, following calls in-module?"""
functions = _functions_in(path)
shadowed = _locally_shadowed_accessors(path)
seen = set()
def walk(name):
@@ -216,7 +245,7 @@ def _reaches_a_bag(path, entry):
continue
if not isinstance(node.func, ast.Name):
continue
if node.func.id in _BAG_ACCESSORS:
if node.func.id in _BAG_ACCESSORS and node.func.id not in shadowed:
return node.lineno
found = walk(node.func.id)
if found is not None:
@@ -241,9 +270,7 @@ class TestResolutionReadsNoBag(CustomTestCase):
def test_the_walk_finds_something_to_walk(self):
"""A collapsed import map would make the pin vacuous."""
imported = _imported_symbols(
[_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"]
)
imported = _imported_symbols(_pipeline_sources())
self.assertGreater(
len(imported),
20,
@@ -282,9 +309,7 @@ class TestResolutionReadsNoBag(CustomTestCase):
)
def test_nothing_the_pipeline_calls_reads_a_bag(self):
imported = _imported_symbols(
[_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"]
)
imported = _imported_symbols(_pipeline_sources())
reachable = {
(path, symbol) for path, symbols in imported.items() for symbol in symbols
} | _registered_entries()
@@ -9,7 +9,7 @@ from types import SimpleNamespace
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 import parallel_hook, pd_disaggregation_hook, serving_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 (
@@ -40,7 +40,9 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
register_cpu_ci(est_time=11, suite="base-c-test-cpu")
# Mock get_device() so all tests run on CPU-only CI runners
_mock_device = patch("sglang.srt.server_args.get_device", return_value="cuda")
_mock_device = patch(
"sglang.srt.arg_groups.serving_hook.get_device", return_value="cuda"
)
_mock_device.start()
@@ -223,7 +225,7 @@ class TestMmEncoderDataParallelLogging(CustomTestCase):
model_path="dummy", mm_enable_dp_encoder=True, tp_size=1
)
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
with self.assertLogs(parallel_hook.logger, level="WARNING") as logs:
server_args._handle_data_parallelism()
self.assertIn("TP=1", logs.output[0])
@@ -234,7 +236,7 @@ class TestMmEncoderDataParallelLogging(CustomTestCase):
model_path="dummy", mm_enable_dp_encoder=True, tp_size=4
)
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(parallel_hook.logger, level="INFO") as logs:
server_args._handle_data_parallelism()
self.assertIn("TP=4", logs.output[0])
@@ -255,7 +257,7 @@ class TestImageProcessorBackend(CustomTestCase):
def test_legacy_flag_maps_to_pil_with_one_warning(self):
server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True)
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_deprecated_args()
self.assertEqual(
@@ -285,7 +287,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
def _set_model_type(server_args, *, is_multimodal):
server_args._model_config = SimpleNamespace(is_multimodal=is_multimodal)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
@@ -295,7 +297,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -307,12 +309,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("base GPU 2", output)
self.assertIn("4 tokenizer worker", output)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_legacy_keep_flag_maps_to_cuda_ipc(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", keep_mm_feature_on_device=True)
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -335,12 +337,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertRaisesRegex(ValueError, "conflicts.*cuda_vmm"):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_explicit_cpu_overrides_legacy_environment(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cpu")
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -361,7 +363,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_text_only_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy")
self._set_model_type(server_args, is_multimodal=False)
@@ -376,7 +378,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_multimodal_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy")
self._set_model_type(server_args, is_multimodal=True)
@@ -391,9 +393,11 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.os.path.exists", return_value=True)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True)
@patch(
"sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
@patch(
"sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport",
return_value=True,
@@ -410,7 +414,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with patch.dict(os.environ, {}, clear=False):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -422,9 +426,11 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("auto-resolved to cuda_vmm", output)
self.assertIn("CUDA FABRIC", output)
@patch("sglang.srt.server_args.os.path.exists", return_value=True)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True)
@patch(
"sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
@patch(
"sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport",
return_value=False,
@@ -439,15 +445,17 @@ class TestMultimodalFeatureTransport(CustomTestCase):
server_args = ServerArgs(model_path="dummy", nnodes=2)
self._set_model_type(server_args, is_multimodal=True)
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
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)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=False)
@patch(
"sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_without_imex_channel(
self, _mock_is_cuda, _mock_is_mnnvl, _mock_path_exists
):
@@ -456,7 +464,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with patch.dict(os.environ, {}, clear=False):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -465,8 +473,10 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("no IMEX channel", "\n".join(logs.output))
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=False)
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch(
"sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=False
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_multinode_non_mnnvl(
self, _mock_is_cuda, _mock_is_mnnvl
):
@@ -482,7 +492,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_language_only_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", language_only=True)
self._set_model_type(server_args, is_multimodal=True)
@@ -496,14 +506,14 @@ class TestMultimodalFeatureTransport(CustomTestCase):
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=False)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=False)
def test_cuda_ipc_rejects_non_nvidia_platforms(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_ipc")
with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_ipc_rejects_multi_node(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy", mm_feature_transport="cuda_ipc", nnodes=2
@@ -512,7 +522,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertRaisesRegex(ValueError, "single node"):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_vmm_is_explicit_and_uses_shared_budget(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
@@ -525,7 +535,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}),
envs.SGLANG_MM_FEATURE_CACHE_MB.override(256),
):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(
@@ -539,14 +549,14 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("2 tokenizer worker", output)
self.assertIn("falls back to inline CPU", output)
@patch("sglang.srt.server_args.is_cuda", return_value=False)
@patch("sglang.srt.arg_groups.serving_hook.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)
@patch("sglang.srt.arg_groups.serving_hook.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")
@@ -556,7 +566,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.serving_hook.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
@@ -577,7 +587,7 @@ class TestMambaCacheStochasticRounding(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "--mamba-ssm-dtype float16"):
server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=False)
@patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=False)
def test_rejects_non_cuda(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
@@ -588,8 +598,8 @@ class TestMambaCacheStochasticRounding(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "NVIDIA CUDA"):
server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.server_args.is_sm100_supported", return_value=False)
@patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=True)
@patch("sglang.srt.arg_groups.mamba_hook.is_sm100_supported", return_value=False)
def test_rejects_triton_without_sm100(self, _mock_sm100, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
@@ -1925,11 +1935,11 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
prefill=PhaseConfig(backend=Backend.TC_PIECEWISE)
)
with (
patch("sglang.srt.server_args.is_hip", return_value=False),
patch("sglang.srt.server_args.is_npu", return_value=False),
patch("sglang.srt.server_args.is_cpu", return_value=False),
patch("sglang.srt.server_args.is_mps", return_value=False),
patch("sglang.srt.server_args.is_xpu", return_value=False),
patch("sglang.srt.arg_groups.cuda_graph_hook.is_hip", return_value=False),
patch("sglang.srt.arg_groups.cuda_graph_hook.is_npu", return_value=False),
patch("sglang.srt.arg_groups.cuda_graph_hook.is_cpu", return_value=False),
patch("sglang.srt.arg_groups.cuda_graph_hook.is_mps", return_value=False),
patch("sglang.srt.arg_groups.cuda_graph_hook.is_xpu", return_value=False),
):
args._disable_tc_piecewise_cudagraph_if_incompatible()
@@ -2661,7 +2671,7 @@ class TestGrpcServerArgs(CustomTestCase):
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:
with self.assertLogs(serving_hook.logger, level="WARNING") as cm:
sa._handle_deprecated_args()
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output))