fix: support dataclass SGLang ServerArgs (#3581)

This commit is contained in:
HEJIAN SANG
2026-09-21 19:35:32 -07:00
committed by GitHub
parent 90ea0b3a8e
commit d3ffe91658
8 changed files with 65 additions and 24 deletions
+6 -6
View File
@@ -4,7 +4,6 @@ import os
import shlex
import sys
import msgspec
from sglang.srt.server_args import ServerArgs
@@ -17,6 +16,7 @@ from miles.utils.lora import (
lora_rollout_enabled,
)
from miles.utils.multi_lora import is_multi_lora_enabled
from miles.utils.workers.argv_utils import _record_field_names
logger = logging.getLogger(__name__)
@@ -183,12 +183,12 @@ def _compute_server_args(
kwargs.update(sglang_overrides)
unused_keys = set(kwargs.keys())
for attr in msgspec.structs.fields(ServerArgs):
if worker_type == "decode" and attr.name == "enable_hierarchical_cache":
for name in _record_field_names(ServerArgs):
if worker_type == "decode" and name == "enable_hierarchical_cache":
continue
if hasattr(args, f"sglang_{attr.name}") and attr.name not in kwargs:
kwargs[attr.name] = getattr(args, f"sglang_{attr.name}")
unused_keys.discard(attr.name)
if hasattr(args, f"sglang_{name}") and name not in kwargs:
kwargs[name] = getattr(args, f"sglang_{name}")
unused_keys.discard(name)
# for compatibility with old args
if len(unused_keys) > 0:
@@ -5,7 +5,6 @@ from collections import defaultdict
from collections.abc import Callable, Sequence
from concurrent.futures import Future, ThreadPoolExecutor
import msgspec
import ray
import torch
from sglang.srt.server_args import ServerArgs
@@ -14,6 +13,7 @@ from miles.backends.training_utils.parallel import get_parallel_state
from miles.backends.training_utils.weight_update.hf_weight_iterator import WeightUpdatePlacement
from miles.backends.training_utils.weight_update.utils import get_data_replica_rank_and_size
from miles.utils import async_utils
from miles.utils.workers.argv_utils import _record_field_names
logger = logging.getLogger(__name__)
@@ -179,7 +179,7 @@ class P2PTransferManager:
def create_server_args_from_dict(data_dict: dict) -> ServerArgs:
valid_fields = {f.name for f in msgspec.structs.fields(ServerArgs)}
valid_fields = set(_record_field_names(ServerArgs))
filtered_data = {k: v for k, v in data_dict.items() if k in valid_fields}
return ServerArgs(**filtered_data)
+1 -1
View File
@@ -107,7 +107,7 @@ def render_cli_argv(
def _record_field_names(record) -> tuple[str, ...]:
"""RouterArgs is a dataclass; ServerArgs is a msgspec Struct since sglang v0.5.20.
"""Accept dataclass and msgspec Struct classes or instances across SGLang versions.
msgspec is imported inside the branch that needs it: this module is on the light
worker entrypoint's import path, whose footprint tests/fast/utils/workers/import_probe.py
@@ -2,7 +2,6 @@ from __future__ import annotations
import argparse
import msgspec
import pytest
@@ -11,6 +10,7 @@ pytest.importorskip("sglang")
from sglang.srt.server_args import ServerArgs
from miles.backends.sglang_utils.arguments import add_sglang_arguments, collect_eval_sglang_overrides
from miles.utils.workers.argv_utils import _record_field_names
def _sglang_flags() -> set[str]:
@@ -57,4 +57,4 @@ class TestAllocatorOwnedServerArgs:
def test_the_skipped_launch_gate_port_names_a_real_server_args_field(self):
"""A renamed upstream field would leave the skip entry stale and quietly re-expose the flag."""
assert "gated_launch_port" in {field.name for field in msgspec.structs.fields(ServerArgs)}
assert "gated_launch_port" in set(_record_field_names(ServerArgs))
@@ -1,9 +1,12 @@
from __future__ import annotations
import dataclasses
from types import SimpleNamespace
import msgspec
import pytest
from miles.backends.sglang_utils import sglang_engine
from miles.backends.sglang_utils.sglang_engine import _compute_server_args
@@ -51,6 +54,21 @@ def compute(args: SimpleNamespace, **overrides: object) -> dict:
return _compute_server_args(args, **kwargs)
@pytest.mark.parametrize(
"record_factory", [dataclasses.make_dataclass, msgspec.defstruct], ids=["dataclass", "msgspec"]
)
def test_server_args_representation_preserves_launch_values(monkeypatch, record_factory):
server_args_type = record_factory(
"ServerArgs",
[("gated_launch_port", int), ("mem_fraction_static", float), ("random_seed", int)],
)
monkeypatch.setattr(sglang_engine, "ServerArgs", server_args_type)
result = compute(make_args(), random_seed=7, sglang_overrides={"random_seed": 99, "unknown_field": True})
assert result == {"gated_launch_port": 30001, "mem_fraction_static": 0.7, "random_seed": 99}
class TestRandomSeed:
def test_the_caller_chosen_seed_reaches_the_engine(self):
"""A seed sglang picks for itself makes a restarted engine replay a different RNG stream."""
@@ -6,7 +6,6 @@ import json
from argparse import Namespace
from typing import Any
import msgspec
import pytest
from tests.fast.backends.sglang_utils.conftest import make_engine_args as _args
from tests.fast.backends.sglang_utils.conftest import tiny_model_path
@@ -17,7 +16,7 @@ from sglang.srt.server_args import ServerArgs
from miles.backends.sglang_utils.server_args_utils import parse_server_args_argv, server_args_to_argv
from miles.backends.sglang_utils.sglang_engine import _compute_server_args
from miles.utils.workers.argv_utils import _actions_by_dest, _render_action_argv, _resolve_action
from miles.utils.workers.argv_utils import _actions_by_dest, _record_field_names, _render_action_argv, _resolve_action
_FIELDS_WITHOUT_A_RENDERABLE_CLI: dict[str, str] = {
"custom_sigquit_handler": "A Python-only callable hook; sglang registers no CLI option for it.",
@@ -80,11 +79,7 @@ def _assert_roundtrips(server_args_dict: dict) -> None:
parsed = parse_server_args_argv(server_args_to_argv(server_args_dict))
device = server_args_dict.get("device") or parsed.device
wanted = ServerArgs(**{**server_args_dict, "device": device})
differing = [
field.name
for field in msgspec.structs.fields(wanted)
if getattr(parsed, field.name) != getattr(wanted, field.name)
]
differing = [name for name in _record_field_names(wanted) if getattr(parsed, name) != getattr(wanted, name)]
assert differing == []
@@ -221,14 +216,14 @@ class TestEveryServerArgsFieldIsRenderable:
[
(
pytest.param(
field.name,
marks=pytest.mark.xfail(reason=_FIELDS_WITHOUT_A_RENDERABLE_CLI[field.name], strict=True),
id=field.name,
name,
marks=pytest.mark.xfail(reason=_FIELDS_WITHOUT_A_RENDERABLE_CLI[name], strict=True),
id=name,
)
if field.name in _FIELDS_WITHOUT_A_RENDERABLE_CLI
else pytest.param(field.name, id=field.name)
if name in _FIELDS_WITHOUT_A_RENDERABLE_CLI
else pytest.param(name, id=name)
)
for field in msgspec.structs.fields(ServerArgs)
for name in _record_field_names(ServerArgs)
],
)
def test_a_field_renders_to_argv_that_parses_back_to_the_same_value(self, field_name: str) -> None:
@@ -1,3 +1,4 @@
import dataclasses
import importlib
import sys
from collections import Counter
@@ -114,6 +115,19 @@ def _query(module, engines: list[_FakeRolloutEngine], pairs: list[tuple[int, int
return module.query_remote_weight_infos(engines, _make_targets(module, pairs))
@pytest.mark.parametrize(
"record_factory", [dataclasses.make_dataclass, msgspec.defstruct], ids=["dataclass", "msgspec"]
)
def test_server_args_from_remote_info_filters_unknown_fields(p2p_transfer_utils, monkeypatch, record_factory):
server_args_type = record_factory("ServerArgs", [("model_path", str)])
monkeypatch.setattr(p2p_transfer_utils, "ServerArgs", server_args_type)
result = p2p_transfer_utils.create_server_args_from_dict({"model_path": "/model", "unknown_field": True})
assert isinstance(result, server_args_type)
assert result.model_path == "/model"
class TestQueryRemoteWeightInfos:
"""Remote-info discovery over the rollout engines' HTTP API."""
@@ -7,6 +7,7 @@ import sys
from pathlib import Path
from typing import Any
import msgspec
import pytest
from pydantic import ValidationError
@@ -16,6 +17,7 @@ from miles.utils.pydantic_utils import FrozenStrictBaseModel
from miles.utils.workers import argv_utils
from miles.utils.workers.argv_utils import (
CONFIG_JSON_FLAG,
_record_field_names,
config_to_argv,
dataclass_to_values,
parse_config_argv,
@@ -23,6 +25,18 @@ from miles.utils.workers.argv_utils import (
)
@pytest.mark.parametrize(
"record_factory", [dataclasses.make_dataclass, msgspec.defstruct], ids=["dataclass", "msgspec"]
)
def test_record_fields_support_classes_and_instances(record_factory):
record_type = record_factory("Args", [("port", int), ("host", str)])
record = record_type(port=30000, host="localhost")
assert _record_field_names(record_type) == ("port", "host")
assert _record_field_names(record) == ("port", "host")
assert dataclass_to_values(record) == {"port": 30000, "host": "localhost"}
class _DemoConfig(FrozenStrictBaseModel):
text: str
count: int