config: six more runtime readers ask the bags (#36973)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
b65e677e48
commit
1a3e152f03
@@ -79,20 +79,17 @@ def _run_negotiate_test(rank, test_cases):
|
||||
|
||||
for case in test_cases:
|
||||
# The DP-attention gate is a published config leaf.
|
||||
override = get_context().override_server_args(enable_dp_attention=True)
|
||||
override = get_context().override_server_args(
|
||||
enable_dp_attention=True,
|
||||
prefill_delayer_queue_min_ratio=case.queue_min_ratio,
|
||||
prefill_delayer_max_delay_ms=case.max_delay_ms,
|
||||
prefill_max_requests=case.prefill_max_requests,
|
||||
)
|
||||
override.install()
|
||||
delayer = PrefillDelayer(
|
||||
dp_size=world_size,
|
||||
attn_tp_size=1,
|
||||
cpu_group=cpu_group,
|
||||
server_args=SimpleNamespace(
|
||||
enable_dp_attention=True,
|
||||
disaggregation_mode="null",
|
||||
disable_overlap_schedule=False,
|
||||
prefill_delayer_queue_min_ratio=case.queue_min_ratio,
|
||||
prefill_delayer_max_delay_ms=case.max_delay_ms,
|
||||
prefill_max_requests=case.prefill_max_requests,
|
||||
),
|
||||
max_delay_passes=case.max_delay_passes,
|
||||
token_usage_low_watermark=case.token_usage_low_watermark,
|
||||
)
|
||||
|
||||
@@ -143,23 +143,20 @@ class TestProfileSpsTable(CustomTestCase):
|
||||
|
||||
|
||||
def _build_sps_cost_table_for(testcase, *, sps_table_path):
|
||||
from sglang.srt.runtime_context import get_context, get_server_args
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.srt.speculative.dspark_components.dspark_planner import (
|
||||
build_sps_cost_table,
|
||||
)
|
||||
|
||||
# The table bound reads `max_running_requests` from the published bags, so
|
||||
# the case publishes it; the table path stays on the handed record, which is
|
||||
# what `build_sps_cost_table` takes.
|
||||
# Both the table path and the bound come from the published bags, so the
|
||||
# case publishes them.
|
||||
override = get_context().override_server_args(
|
||||
speculative_dspark_sps_table_path=sps_table_path,
|
||||
max_running_requests=4,
|
||||
)
|
||||
override.install()
|
||||
testcase.addCleanup(override.restore)
|
||||
return build_sps_cost_table(
|
||||
server_args=get_server_args(), verify_num_draft_tokens=5
|
||||
)
|
||||
return build_sps_cost_table(verify_num_draft_tokens=5)
|
||||
|
||||
|
||||
class TestBuildSpsCostTableContract(CustomTestCase):
|
||||
|
||||
@@ -39,17 +39,7 @@ from sglang.srt.observability.metrics_collector import (
|
||||
TokenizerMetricsCollector,
|
||||
resolve_collector_class,
|
||||
)
|
||||
|
||||
|
||||
class _StubArgs:
|
||||
"""Minimal ServerArgs stand-in.
|
||||
|
||||
Avoids triggering the heavy real ServerArgs import chain for unit-level
|
||||
``resolve_collector_class`` cases.
|
||||
"""
|
||||
|
||||
def __init__(self, stat_loggers=None):
|
||||
self.stat_loggers = stat_loggers
|
||||
from sglang.srt.runtime_context import get_context, reset_context
|
||||
|
||||
|
||||
class TestCollectorClassAttrs(unittest.TestCase):
|
||||
@@ -79,30 +69,37 @@ class TestCollectorClassAttrs(unittest.TestCase):
|
||||
|
||||
|
||||
class TestResolveCollectorClass(unittest.TestCase):
|
||||
def test_returns_default_when_server_args_none(self):
|
||||
cls = resolve_collector_class(None, "scheduler", SchedulerMetricsCollector)
|
||||
self.assertIs(cls, SchedulerMetricsCollector)
|
||||
"""The role table is read from the published `observability` bag."""
|
||||
|
||||
def _resolve(self, role, default_cls, **fields):
|
||||
if not fields:
|
||||
return resolve_collector_class(role, default_cls)
|
||||
with get_context().override_server_args(**fields):
|
||||
return resolve_collector_class(role, default_cls)
|
||||
|
||||
def test_returns_default_when_nothing_is_published(self):
|
||||
reset_context()
|
||||
self.assertIs(
|
||||
resolve_collector_class("scheduler", SchedulerMetricsCollector),
|
||||
SchedulerMetricsCollector,
|
||||
)
|
||||
|
||||
def test_returns_default_when_stat_loggers_none(self):
|
||||
cls = resolve_collector_class(
|
||||
_StubArgs(stat_loggers=None), "scheduler", SchedulerMetricsCollector
|
||||
)
|
||||
cls = self._resolve("scheduler", SchedulerMetricsCollector, stat_loggers=None)
|
||||
self.assertIs(cls, SchedulerMetricsCollector)
|
||||
|
||||
def test_returns_default_when_stat_loggers_empty(self):
|
||||
cls = resolve_collector_class(
|
||||
_StubArgs(stat_loggers={}), "scheduler", SchedulerMetricsCollector
|
||||
)
|
||||
cls = self._resolve("scheduler", SchedulerMetricsCollector, stat_loggers={})
|
||||
self.assertIs(cls, SchedulerMetricsCollector)
|
||||
|
||||
def test_returns_default_when_role_missing(self):
|
||||
class MyTokenizer(TokenizerMetricsCollector):
|
||||
pass
|
||||
|
||||
cls = resolve_collector_class(
|
||||
_StubArgs(stat_loggers={"tokenizer": MyTokenizer}),
|
||||
cls = self._resolve(
|
||||
"scheduler",
|
||||
SchedulerMetricsCollector,
|
||||
stat_loggers={"tokenizer": MyTokenizer},
|
||||
)
|
||||
self.assertIs(cls, SchedulerMetricsCollector)
|
||||
|
||||
@@ -110,10 +107,10 @@ class TestResolveCollectorClass(unittest.TestCase):
|
||||
class MyScheduler(SchedulerMetricsCollector):
|
||||
pass
|
||||
|
||||
cls = resolve_collector_class(
|
||||
_StubArgs(stat_loggers={"scheduler": MyScheduler}),
|
||||
cls = self._resolve(
|
||||
"scheduler",
|
||||
SchedulerMetricsCollector,
|
||||
stat_loggers={"scheduler": MyScheduler},
|
||||
)
|
||||
self.assertIs(cls, MyScheduler)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user