test: publish resolved config in unit fixtures for the namespace API (#31817)
This commit is contained in:
@@ -31,7 +31,7 @@ _RATCHETS = [
|
|||||||
(
|
(
|
||||||
"set_global_server_args_for_*",
|
"set_global_server_args_for_*",
|
||||||
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
|
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
|
||||||
5,
|
4,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import unittest
|
|||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import sglang.srt.server_args as server_args_module
|
import sglang.srt.server_args as server_args_module
|
||||||
from sglang.srt.arg_groups.arg_utils import A, Arg
|
from sglang.srt.arg_groups.arg_utils import NS, A, Arg
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
Flags,
|
Flags,
|
||||||
ParallelContext,
|
ParallelContext,
|
||||||
@@ -18,6 +18,7 @@ from sglang.srt.runtime_context import (
|
|||||||
_FlagGroupBase,
|
_FlagGroupBase,
|
||||||
get_context,
|
get_context,
|
||||||
get_flags,
|
get_flags,
|
||||||
|
get_memory,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
reset_context,
|
reset_context,
|
||||||
@@ -215,8 +216,10 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
|
|||||||
self.assertIs(get_server_args(), sentinel)
|
self.assertIs(get_server_args(), sentinel)
|
||||||
self.assertIs(get_context().server_args, sentinel)
|
self.assertIs(get_context().server_args, sentinel)
|
||||||
|
|
||||||
def test_tokenizer_alias_is_same_function(self):
|
def test_tokenizer_and_scheduler_setters_are_distinct_role_shims(self):
|
||||||
self.assertIs(
|
# The per-role publish shims are no longer aliases: each records its own
|
||||||
|
# process role via publish(role=...).
|
||||||
|
self.assertIsNot(
|
||||||
server_args_module.set_global_server_args_for_tokenizer,
|
server_args_module.set_global_server_args_for_tokenizer,
|
||||||
server_args_module.set_global_server_args_for_scheduler,
|
server_args_module.set_global_server_args_for_scheduler,
|
||||||
)
|
)
|
||||||
@@ -373,8 +376,10 @@ class TestFlagsTier(_IsolatedServerArgs):
|
|||||||
class _FakeResolvedArgs:
|
class _FakeResolvedArgs:
|
||||||
"""Publishable fixture with a resolvable whitelist (real flat leaves)."""
|
"""Publishable fixture with a resolvable whitelist (real flat leaves)."""
|
||||||
|
|
||||||
page_size: A[int | None, Arg(help="p", resolvable=True)] = None
|
page_size: A[int | None, Arg(help="p", resolvable=True), NS("memory")] = None
|
||||||
sampling_backend: A[str | None, Arg(help="s", resolvable=True)] = None
|
sampling_backend: A[
|
||||||
|
str | None, Arg(help="s", resolvable=True), NS("exec.kernel")
|
||||||
|
] = None
|
||||||
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
@@ -916,12 +921,15 @@ class TestPublishLifecycle(_IsolatedServerArgs):
|
|||||||
get_context().set_server_args(object())
|
get_context().set_server_args(object())
|
||||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||||
|
|
||||||
def test_declare_load_time_override_writes_through(self):
|
def test_declare_load_time_override_writes_the_bag(self):
|
||||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||||
|
|
||||||
args = self._publish(page_size=1)
|
args = self._publish(page_size=1)
|
||||||
declare_load_time_override("model.load_time", {"page_size": 64})
|
declare_load_time_override("model.load_time", {"page_size": 64})
|
||||||
self.assertEqual(args.page_size, 64)
|
# The declaration lands on the config bag; the pristine startup record
|
||||||
|
# (server_args) is untouched.
|
||||||
|
self.assertEqual(get_memory().page_size, 64)
|
||||||
|
self.assertEqual(args.page_size, 1)
|
||||||
|
|
||||||
def test_declare_load_time_override_validates_whitelist(self):
|
def test_declare_load_time_override_validates_whitelist(self):
|
||||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||||
@@ -933,16 +941,14 @@ class TestPublishLifecycle(_IsolatedServerArgs):
|
|||||||
|
|
||||||
def test_declare_load_time_override_records_provenance(self):
|
def test_declare_load_time_override_records_provenance(self):
|
||||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
class _Args(_FakeResolvedArgs):
|
self._publish(page_size=1)
|
||||||
override = ServerArgs.override
|
|
||||||
|
|
||||||
args = _Args(page_size=1)
|
|
||||||
get_context().set_server_args(args)
|
|
||||||
declare_load_time_override("model.load_time", {"page_size": 64})
|
declare_load_time_override("model.load_time", {"page_size": 64})
|
||||||
self.assertEqual(args.page_size, 64)
|
self.assertEqual(get_memory().page_size, 64)
|
||||||
self.assertIn(("model.load_time", {"page_size": 64}), args._resolved_overrides)
|
self.assertIn(
|
||||||
|
("model.load_time", {"page_size": 64}),
|
||||||
|
get_context().overrides_log(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user