[refactor] Add predicate-keyed registration; migrate the Step3p family (stack 8/15) (#30070)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-07-04 02:21:43 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent 8d8f17e28c
commit 4bf4db09f6
4 changed files with 159 additions and 30 deletions
+83 -2
View File
@@ -60,7 +60,15 @@ class TestModelOverridableWhitelist(CustomTestCase):
self.assertEqual(
model_overridable_fields(ServerArgs),
frozenset({"dtype", "enable_tf32_matmul", "enable_multi_layer_eagle"}),
frozenset(
{
"dtype",
"enable_tf32_matmul",
"enable_multi_layer_eagle",
"swa_full_tokens_ratio",
"disable_hybrid_swa_memory",
}
),
)
def test_non_dataclass_yields_empty_whitelist(self):
@@ -75,6 +83,7 @@ class _IsolatedRegistry(CustomTestCase):
self._patches = [
patch.dict(overrides_module.MODEL_OVERRIDES, clear=True),
patch.dict(overrides_module._MODEL_OVERRIDE_FNS, clear=True),
patch.object(overrides_module, "_PREDICATE_OVERRIDE_FNS", []),
]
for p in self._patches:
p.start()
@@ -131,6 +140,27 @@ class TestModelOverrideRegistry(_IsolatedRegistry):
with self.assertRaises(TypeError):
collect_model_override_declarations("FakeForCausalLM", None, None)
def test_predicate_keyed_provider(self):
from sglang.srt.arg_groups.overrides import register_model_override_predicate
@register_model_override("FakeStep9ForCausalLM")
def _exact(server_args, hf_config):
return {"a": 1}
@register_model_override_predicate(lambda arch: "Step9" in arch)
def _by_predicate(server_args, hf_config):
return {"b": 2}
# matching arch: exact-keyed first, then predicate-keyed
self.assertEqual(
collect_model_override_declarations("FakeStep9ForCausalLM", None, None),
[(_exact.__qualname__, {"a": 1}), (_by_predicate.__qualname__, {"b": 2})],
)
# non-matching arch: predicate does not fire
self.assertEqual(
collect_model_override_declarations("OtherForCausalLM", None, None), []
)
@dataclasses.dataclass
class _FakeAttnGroup(_StaticFlags):
@@ -292,13 +322,14 @@ class TestGoldenModelOverrides(_IsolatedPublish):
"v_head_dim": 16,
}
def _construct(self, arch, model_type, **server_kwargs):
def _construct(self, arch, model_type, config_extra=None, **server_kwargs):
from sglang.srt.server_args import ServerArgs
# Golden resolution must be host-independent: accelerator-less CI
# runners resolve only the base platform, where get_device() raises.
server_kwargs.setdefault("device", "cuda")
config = dict(self._MINI_CONFIG, architectures=[arch], model_type=model_type)
config.update(config_extra or {})
config_dir = tempfile.mkdtemp(prefix="golden_override_")
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
with open(os.path.join(config_dir, "config.json"), "w") as f:
@@ -377,6 +408,56 @@ class TestGoldenModelOverrides(_IsolatedPublish):
[("_mimo_v2_overrides", {"enable_multi_layer_eagle": True})],
)
def test_step3p_hierarchical_cache_golden(self):
# SWA-hybrid arch: the mini config needs layer_types/sliding_window.
config_extra = {
"layer_types": ["sliding_attention", "full_attention"],
"sliding_window": 64,
}
sa = self._construct(
"Step3p5ForCausalLM",
"llama",
config_extra=config_extra,
enable_hierarchical_cache=True,
)
# dual-apply == legacy writes
self.assertEqual(sa.swa_full_tokens_ratio, 1.0)
self.assertTrue(sa.disable_hybrid_swa_memory)
flags = self._publish(sa)
self.assertEqual(flags.swa_full_tokens_ratio, 1.0)
self.assertTrue(flags.disable_hybrid_swa_memory)
def test_step3p_declarations_at_callable_level(self):
from sglang.srt.arg_groups.overrides import _step3p_overrides
self.assertEqual(
_step3p_overrides(
SimpleNamespace(
speculative_algorithm="EAGLE", enable_hierarchical_cache=False
),
None,
),
{"enable_multi_layer_eagle": True},
)
self.assertEqual(
_step3p_overrides(
SimpleNamespace(
speculative_algorithm=None, enable_hierarchical_cache=True
),
None,
),
{"swa_full_tokens_ratio": 1.0, "disable_hybrid_swa_memory": True},
)
self.assertEqual(
_step3p_overrides(
SimpleNamespace(
speculative_algorithm=None, enable_hierarchical_cache=False
),
None,
),
{},
)
class TestDualApplyParity(CustomTestCase):
def test_dual_apply_replays_and_parity_holds(self):