mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
fix: support dataclass SGLang ServerArgs (#3581)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user