Migrate unit tests and preserve Verifiers contract coverage

Materialize the final GU operation with unit-test migrations and regression coverage. Retire the AllConfig parallel-property test case alongside the explicitly deleted op48-14; preserve TrainerConfig product, no-duplicate-field, and immutability assertions.
This commit is contained in:
Tom
2026-10-01 15:43:36 +08:00
parent eb9e76737b
commit 6501f53c71
413 changed files with 142331 additions and 2842 deletions
+81
View File
@@ -0,0 +1,81 @@
from dataclasses import dataclass
from pathlib import Path
import pytest
from tests.ci import ci_utils
from miles.utils.audit_utils.config_snapshot.converter import ConfigSnapshotConverter
from miles.utils.audit_utils.config_snapshot.models import (
ConfigSnapshotContext,
ConfigSnapshotPoint,
ConfigSnapshotRecord,
)
from miles.utils.audit_utils.config_snapshot.runner import ConfigSnapshotTestRunner
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
from miles.utils.test_utils.snapshot import SNAPSHOT_RECORD_DIR_ENV_VAR, SNAPSHOT_UPDATE_ENV_VAR, dump_snapshot
@dataclass(frozen=True)
class _SnapshotFileCase:
test_file: Path
record_root: Path
golden: Path
record: ConfigSnapshotRecord
def write_golden(self, *, value: str) -> None:
record = self.record.model_copy(update={"config": {"args": {"value": value}}})
self.golden.parent.mkdir(parents=True, exist_ok=True)
self.golden.write_text(dump_snapshot(ConfigSnapshotConverter.convert([record])))
@pytest.fixture
def snapshot_file_case(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> _SnapshotFileCase:
relative = Path("tests/e2e/test_snapshot_probe.py")
script = tmp_path / relative
script.parent.mkdir(parents=True)
script.write_text(
"import json\n"
"import os\n"
"import time\n"
"from pathlib import Path\n"
"records = Path(os.environ['MILES_SNAPSHOT_RECORD_DIR'])\n"
"records.mkdir(parents=True, exist_ok=True)\n"
"(records / 'record.json').write_text(Path(__file__).with_suffix('.json').read_text())\n"
"metric_dir = Path(os.environ['MILES_CI_GATE_RECORD_DIR'])\n"
"Path(__file__).with_suffix('.capture').write_text(str(metric_dir))\n"
"(metric_dir / 'probe.jsonl').write_text(json.dumps({'metric': 'train/grad_norm', 'series': [[0, 1.5]]}) + '\\n')\n"
"time.sleep(float(os.environ.get('FILE_RUN_TEST_SLEEP', '0')))\n"
"raise SystemExit(int(os.environ.get('FILE_RUN_TEST_EXIT', '0')))\n"
)
record = ConfigSnapshotRecord(
context=ConfigSnapshotContext(
name="tests/e2e/test_snapshot_probe/run-0000",
deploy_component="all",
deploy_instance_id="default",
source=SimpleProcessIdentity(component="main"),
run_uuid="test-run",
capture_id="test-capture",
),
point=ConfigSnapshotPoint(stage="process_config", index=0),
config={"args": {"value": "actual"}},
)
script.with_suffix(".json").write_text(record.model_dump_json())
record_root = tmp_path / "record-root"
monkeypatch.chdir(tmp_path)
monkeypatch.setattr(ci_utils, "__file__", str(tmp_path / "tests/ci/ci_utils.py"))
monkeypatch.setenv(SNAPSHOT_RECORD_DIR_ENV_VAR, str(record_root))
monkeypatch.setenv("CI", "false")
monkeypatch.setenv(ci_utils.CI_GATE_RECORD_DIR_ENV, "")
for name in (
SNAPSHOT_UPDATE_ENV_VAR,
"FILE_RUN_TEST_SLEEP",
"FILE_RUN_TEST_EXIT",
ci_utils.CI_GATE_RECORD_DIR_ENV,
):
monkeypatch.delenv(name, raising=False)
return _SnapshotFileCase(
test_file=relative,
record_root=record_root,
golden=ConfigSnapshotTestRunner.golden_path(test=str(relative), repo_root=tmp_path),
record=record,
)
+8 -3
View File
@@ -175,7 +175,8 @@ def test_main_fails_closed_on_an_unregistered_file(monkeypatch, tmp_path, capsys
assert "::error::" in capsys.readouterr().err
def test_target_workflow_keeps_orchestration_trusted_and_checks_out_exact_head():
def test_target_workflow_keeps_orchestration_trusted_and_checks_out_exact_head() -> None:
"""Resolved file jobs preserve trust boundaries and use the snapshot-aware CUDA entrypoint."""
root = Path(__file__).parents[3]
workflow = (root / ".github/workflows/run-ci-file.yml").read_text()
gpu_workflow = (root / ".github/workflows/_run-ci.yml").read_text()
@@ -223,9 +224,13 @@ def test_target_workflow_keeps_orchestration_trusted_and_checks_out_exact_head()
assert "if: always()" in report_job
assert "tests.ci.run_suite" not in workflow
assert "pytest '${{ inputs.test_file }}' -v -x" in workflow
assert "python3 '${{ inputs.test_file }}'" in workflow
assert "python3 -m tests.ci.run_file" in workflow
assert "--test-file '${{ inputs.test_file }}'" in workflow
assert "--timeout-seconds '${{ needs.resolve-file-run.outputs.timeout_seconds }}'" in workflow
assert '"$(( ${{ needs.resolve-file-run.outputs.timeout_seconds }} + 120 ))s"' in workflow
assert "python3 '${{ inputs.test_file }}'" not in workflow
assert workflow.count("timeout --signal=TERM --kill-after=30s") == 2
assert workflow.count("'${{ needs.resolve-file-run.outputs.timeout_seconds }}s'") == 2
assert workflow.count("'${{ needs.resolve-file-run.outputs.timeout_seconds }}s'") == 1
for reusable in (gpu_workflow, cpu_workflow):
assert "plan_already_resolved:" in reusable
assert "if: ${{ !inputs.plan_already_resolved }}" in reusable
+82
View File
@@ -0,0 +1,82 @@
import json
import logging
import os
from pathlib import Path
import pytest
from tests.ci.ci_register import register_cpu_ci
from tests.ci.ci_utils import CI_GATE_RECORD_DIR_ENV
from tests.ci.run_file import app
from tests.ci.test.conftest import _SnapshotFileCase
from typer.testing import CliRunner
register_cpu_ci(est_time=10, suite="stage-a-cpu", labels=[])
@pytest.mark.parametrize("existing_directory", [False, True])
def test_child_captures_metrics_without_a_database_store(
snapshot_file_case: _SnapshotFileCase,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
existing_directory: bool,
) -> None:
"""Single-file runs retain suite-equivalent metric capture without database access."""
snapshot_file_case.write_golden(value="actual")
supplied_directory = tmp_path / "metrics"
if existing_directory:
monkeypatch.setenv(CI_GATE_RECORD_DIR_ENV, str(supplied_directory))
monkeypatch.setenv("NEON_DATABASE_URL", "postgresql://invalid.invalid/forbidden")
result = CliRunner().invoke(app, ["--test-file", str(snapshot_file_case.test_file), "--timeout-seconds", "10"])
assert result.exit_code == 0, result.output
record_directory = Path(snapshot_file_case.test_file.with_suffix(".capture").read_text())
base_directory = Path(os.environ[CI_GATE_RECORD_DIR_ENV])
assert record_directory.is_relative_to(base_directory)
if existing_directory:
assert base_directory == supplied_directory
assert json.loads((record_directory / "probe.jsonl").read_text()) == {
"metric": "train/grad_norm",
"series": [[0, 1.5]],
}
assert record_directory.with_suffix(".merged.jsonl").is_file()
@pytest.mark.parametrize("golden_value", ["actual", "different"])
def test_successful_child_completes_attempt_and_checks_the_golden(
snapshot_file_case: _SnapshotFileCase, golden_value: str
) -> None:
"""A real successful child completes its attempt even when the snapshot comparison fails."""
snapshot_file_case.write_golden(value=golden_value)
before = snapshot_file_case.golden.read_bytes()
result = CliRunner().invoke(app, ["--test-file", str(snapshot_file_case.test_file), "--timeout-seconds", "10"])
assert result.exit_code == (0 if golden_value == "actual" else -1), result.output
[attempt] = snapshot_file_case.record_root.glob("*/*/attempt.json")
assert json.loads(attempt.read_text()) == {"test": str(snapshot_file_case.test_file), "completed": True}
assert (attempt.parent / "records/record.json").is_file()
assert snapshot_file_case.golden.read_bytes() == before
@pytest.mark.parametrize("failure", ["exit", "timeout"])
def test_failed_or_timed_out_child_keeps_an_incomplete_attempt(
snapshot_file_case: _SnapshotFileCase,
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
failure: str,
) -> None:
"""Real child failure and timeout retain raw records without claiming successful completion."""
snapshot_file_case.write_golden(value="actual")
caplog.set_level(logging.INFO)
monkeypatch.setenv(
"FILE_RUN_TEST_EXIT" if failure == "exit" else "FILE_RUN_TEST_SLEEP", "7" if failure == "exit" else "60"
)
result = CliRunner().invoke(app, ["--test-file", str(snapshot_file_case.test_file), "--timeout-seconds", "2"])
assert result.exit_code == -1, result.output
assert ("returned exit code 7" if failure == "exit" else "after 2 seconds") in caplog.text
[attempt] = snapshot_file_case.record_root.glob("*/*/attempt.json")
assert json.loads(attempt.read_text()) == {"test": str(snapshot_file_case.test_file), "completed": False}
assert (attempt.parent / "records/record.json").is_file()
+9
View File
@@ -14,6 +14,15 @@ from tests.fast.fixtures.rollout_fixtures import rollout_env
_ = rollout_env, generation_env
@pytest.fixture(autouse=True)
def isolate_config_snapshot_generated_values(monkeypatch: pytest.MonkeyPatch) -> None:
"""Each test owns an isolated process-local provenance registry."""
from miles.utils.audit_utils.config_snapshot import generated_values
monkeypatch.setattr(generated_values, "_generated_values", [])
monkeypatch.delenv(generated_values.GENERATED_VALUES_ENV_VAR, raising=False)
@pytest.fixture(autouse=True)
def no_env_reporting(monkeypatch):
"""Constructing a worker configures its logger, which in a real process starts a thread that
@@ -9,6 +9,7 @@ from pathlib import Path
import pytest
import torch
from tests.ci.ci_register import register_cuda_ci
from tests.e2e.long.verifiers_contract import verify_sdk_contract
from miles.utils.external_utils import command_utils
@@ -39,6 +40,13 @@ def prepare():
f"{VERIFIERS_VENV}/bin/python -m pip install "
f"-r {command_utils.repo_base_dir}/examples/experimental/verifiers/requirements.txt"
)
verify_sdk_contract(
execute=U.exec_command_cpu,
venv=VERIFIERS_VENV,
repo_root=Path(command_utils.repo_base_dir),
report_path=RUN_DIR / "sdk-contract.xml",
pythonpath=f"{VERIFIERS_SITE_PACKAGES}:{command_utils.repo_base_dir}:{MEGATRON_PATH}:{os.environ.get('PYTHONPATH', '')}",
)
# prime pins no upper bound on prime-sandboxes, and 0.3.0 dropped the
# CommandRequest the pinned prime imports, so the tool env has to cap it.
U.exec_command_cpu("uv tool install 'prime==0.6.19' --with 'prime-sandboxes<0.3'")
+46
View File
@@ -0,0 +1,46 @@
import shlex
from collections.abc import Callable
from pathlib import Path
from xml.etree import ElementTree
def verify_sdk_contract(
*,
execute: Callable[[str], str | None],
venv: Path,
repo_root: Path,
report_path: Path,
pythonpath: str,
) -> None:
test_file = repo_root / "tests/fast/examples/experimental/verifiers/test_verifiers_rollout.py"
execute(
shlex.join(
[
"env",
"CUDA_VISIBLE_DEVICES=",
f"VIRTUAL_ENV={venv}",
f"PYTHONPATH={pythonpath}",
"uv",
"run",
"--no-project",
"--active",
"python",
"-m",
"pytest",
f"{test_file}::test_canonical_tokenizer_selects_tool_renderer_for_ambiguous_local_checkpoint",
f"{test_file}::test_verifiers_episode_owns_group_reward_computation",
"-v",
f"--junitxml={report_path}",
]
)
)
_assert_contract_report(report_path)
def _assert_contract_report(report_path: Path) -> None:
cases = ElementTree.parse(report_path).getroot().findall(".//testcase")
assert len(cases) == 3, f"Expected three Verifiers SDK contract cases, found {len(cases)}"
for case in cases:
assert all(
case.find(status) is None for status in ("skipped", "failure", "error")
), f"Verifiers SDK contract did not pass: {case.attrib}"
+18 -11
View File
@@ -9,6 +9,8 @@ register_cuda_ci(
)
from types import SimpleNamespace
import pytest
import torch
from tools.convert_hf_to_mxfp8 import quantize_mxfp8 as tool_quantize_mxfp8
@@ -84,12 +86,21 @@ def _te_mxfp8_reference(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tenso
)
def _quantizer_args(
*, extra_high_precision_layers_megatron: tuple[str, ...] | None = None, **backend_fields: object
) -> SimpleNamespace:
return SimpleNamespace(
backend=SimpleNamespace(**backend_fields),
extra_high_precision_layers_megatron=extra_high_precision_layers_megatron,
)
def test_mxfp8_quantize_params_respects_extra_high_precision_layers_megatron():
weight = torch.randn((4, MXFP8_GROUP_SIZE), dtype=torch.bfloat16)
converted_named_params = [
("model.layers.0.mlp.experts.0.down_proj.weight", weight),
]
args = type("Args", (), {"extra_high_precision_layers_megatron": ("linear_fc2",)})()
args = _quantizer_args(extra_high_precision_layers_megatron=("linear_fc2",))
out = quantize_params_mxfp8(
args=args,
@@ -107,16 +118,12 @@ def test_mxfp8_quantize_params_respects_first_last_layers_bf16(layer_idx):
converted_named_params = [
("model.layers.0.mlp.experts.0.down_proj.weight", weight),
]
args = type(
"Args",
(),
{
"first_last_layers_bf16": True,
"num_layers": 4,
"num_layers_at_start_in_bf16": 1,
"num_layers_at_end_in_bf16": 1,
},
)()
args = _quantizer_args(
first_last_layers_bf16=True,
num_layers=4,
num_layers_at_start_in_bf16=1,
num_layers_at_end_in_bf16=1,
)
out = quantize_params_mxfp8(
args=args,
+19 -14
View File
@@ -11,7 +11,7 @@ register_cuda_ci(
import json
import os
import sys
from types import ModuleType
from types import ModuleType, SimpleNamespace
import pytest
import safetensors
@@ -148,11 +148,20 @@ def test_nvfp4_global_scale_exact_without_host_scalar_read(device, e4m3_max, sha
torch.testing.assert_close(decoded.cpu(), torch.div(1.0, expected_encode), rtol=0, atol=0, equal_nan=True)
def _quantizer_args(
*, extra_high_precision_layers_megatron: tuple[str, ...] | None = None, **backend_fields: object
) -> SimpleNamespace:
return SimpleNamespace(
backend=SimpleNamespace(**backend_fields),
extra_high_precision_layers_megatron=extra_high_precision_layers_megatron,
)
def test_nvfp4_quantize_params_requires_complete_gated_pair():
weight = torch.randn((4, NVFP4_GROUP_SIZE), dtype=torch.float32)
with pytest.raises(ValueError, match="requires gate/up tensors to be quantized together"):
quantize_params_nvfp4(
args=None,
args=_quantizer_args(),
megatron_name="decoder.layers.0.mlp.experts.linear_fc1.weight0",
converted_named_params=[
("model.layers.0.mlp.experts.0.gate_proj.weight", weight),
@@ -167,7 +176,7 @@ def test_nvfp4_quantize_params_respects_extra_high_precision_layers_megatron():
("model.layers.0.mlp.experts.0.gate_proj.weight", weight),
("model.layers.0.mlp.experts.0.up_proj.weight", weight),
]
args = type("Args", (), {"extra_high_precision_layers_megatron": ("linear_fc1",)})()
args = _quantizer_args(extra_high_precision_layers_megatron=("linear_fc1",))
out = quantize_params_nvfp4(
args=args,
@@ -186,16 +195,12 @@ def test_nvfp4_quantize_params_respects_first_last_layers_bf16(layer_idx):
("model.layers.0.mlp.experts.0.gate_proj.weight", weight),
("model.layers.0.mlp.experts.0.up_proj.weight", weight),
]
args = type(
"Args",
(),
{
"first_last_layers_bf16": True,
"num_layers": 4,
"num_layers_at_start_in_bf16": 1,
"num_layers_at_end_in_bf16": 1,
},
)()
args = _quantizer_args(
first_last_layers_bf16=True,
num_layers=4,
num_layers_at_start_in_bf16=1,
num_layers_at_end_in_bf16=1,
)
out = quantize_params_nvfp4(
args=args,
@@ -222,7 +227,7 @@ def test_nvfp4_quantize_params_omits_static_input_scale(monkeypatch):
)
out = quantize_params_nvfp4(
args=None,
args=_quantizer_args(),
megatron_name="decoder.layers.0.mlp.experts.linear_fc1.weight0",
converted_named_params=[
("model.layers.0.mlp.experts.0.gate_proj.weight", weight),
@@ -26,7 +26,7 @@ import argparse
import os
import sys
from pathlib import Path
from types import ModuleType
from types import ModuleType, SimpleNamespace
from typing import Any
from unittest.mock import MagicMock, patch
@@ -109,6 +109,25 @@ from miles.utils.debug_utils.run_megatron.worker.main import ( # noqa: E402
_MODULE = "miles.utils.debug_utils.run_megatron.worker.main"
def _parsed_standalone_args(*, trainers: list[SimpleNamespace], **overrides: Any) -> SimpleNamespace:
defaults = {
"train_backend": "megatron",
"debug_train_only": True,
"debug_rollout_only": False,
"offload_train": False,
"colocate": False,
"starts_inference_engines": False,
"hf_checkpoint": "/model",
"actor_num_nodes": 1,
"actor_num_gpus_per_node": 1,
}
return SimpleNamespace(**(defaults | overrides), raw_megatron=SimpleNamespace(trainers=trainers))
def _trainer_config(*, world_size: int) -> SimpleNamespace:
return SimpleNamespace(backend=SimpleNamespace(world_size=world_size), offload_train=False)
class TestParseArgs:
def test_script_options_are_translated_for_the_shared_parser(self) -> None:
"""Use the shared parser with standalone topology and checkpoint options."""
@@ -128,19 +147,30 @@ class TestParseArgs:
"miles",
]
captured: list[str] = []
parsed = object()
actor, critic = SimpleNamespace(role="actor"), SimpleNamespace(role="critic")
parsed = _parsed_standalone_args(
actor_num_nodes=2,
actor_num_gpus_per_node=4,
advantage_estimator="ppo",
critic_load="/checkpoint",
trainers=[actor, critic],
)
trainer_config = _trainer_config(world_size=8)
def parse_shared() -> object:
def parse_shared() -> SimpleNamespace:
captured.extend(sys.argv[1:])
return parsed
with patch.object(sys, "argv", argv), patch.dict(
os.environ, {"WORLD_SIZE": "8", "LOCAL_WORLD_SIZE": "4"}
), patch(f"{_MODULE}.parse_args", side_effect=parse_shared):
), patch(f"{_MODULE}.parse_args", side_effect=parse_shared), patch(
f"{_MODULE}.compute_trainer_config", return_value=trainer_config
) as compute:
args, script_args = _parse_args()
assert sys.argv is argv
assert args is parsed
assert args is trainer_config
compute.assert_called_once_with(parsed, critic)
assert script_args.ref_load == Path("/checkpoint")
for flag, value in [
("--hf-checkpoint", "/model"),
@@ -169,11 +199,15 @@ class TestParseArgs:
]
captured: list[str] = []
def parse_shared() -> object:
def parse_shared() -> SimpleNamespace:
captured.extend(sys.argv[1:])
return object()
return _parsed_standalone_args(trainers=[SimpleNamespace(role="actor")])
with patch.object(sys, "argv", argv), patch(f"{_MODULE}.parse_args", side_effect=parse_shared):
with patch.object(sys, "argv", argv), patch.dict(
os.environ, {"WORLD_SIZE": "1", "LOCAL_WORLD_SIZE": "1"}
), patch(f"{_MODULE}.parse_args", side_effect=parse_shared), patch(
f"{_MODULE}.compute_trainer_config", return_value=_trainer_config(world_size=1)
):
_parse_args()
assert captured.count("--load") == 1
@@ -249,10 +249,9 @@ def _memory_args(chunk_size: int, vocab_size: int) -> Namespace:
qkv_format="thd",
rollout_temperature=1.0,
true_on_policy_mode=True,
bf16=True,
fp16=False,
train_backend="megatron",
backend=Namespace(bf16=True, fp16=False, vocab_size=vocab_size),
log_probs_chunk_size=chunk_size,
vocab_size=vocab_size,
allgather_cp=False,
debug_unified_grad_fused_logprob=False,
)
@@ -71,7 +71,8 @@ def _capture_scenario(scenario: _Scenario) -> dict[str, Any]:
arguments.extend(scenario.arguments)
with _environment(arguments=arguments, legacy=scenario.legacy):
_, parser = parse_args_and_get_parser()
parsed = {"minimal": vars(parser.parse_args(arguments))}
minimal = vars(parser.parse_args(arguments))
parsed = {"minimal": minimal}
variants = {
"lora_disabled": ["--no-sglang-lora-use-virtual-experts"],
"sglang_alias": ["--sglang-tp-size", "2"],
@@ -79,10 +80,17 @@ def _capture_scenario(scenario: _Scenario) -> dict[str, Any]:
"eval_false": ["--no-eval-sglang-enable-metrics"],
}
for name, extra in variants.items():
parsed[name] = vars(parser.parse_args(arguments + extra))
parsed[name] = _parsed_difference(base=minimal, parsed=vars(parser.parse_args(arguments + extra)))
return {"argv": arguments, "legacy": scenario.legacy, "schema": snapshot_parser(parser), "parsed": parsed}
def _parsed_difference(*, base: dict[str, Any], parsed: dict[str, Any]) -> dict[str, Any]:
return {
"changed": {name: value for name, value in parsed.items() if name not in base or base[name] != value},
"removed": sorted(base.keys() - parsed.keys()),
}
class _HookConfig(BaseConfig):
snapshot_hook: A[int, Arg()] = 23
@@ -3,8 +3,4 @@ from typing import Any
def snapshot_parser(parser: argparse.ArgumentParser) -> dict[str, Any]:
destinations = dict.fromkeys([action.dest for action in parser._actions] + list(parser._defaults))
return {
"parser": parser,
"effective_defaults": {dest: parser.get_default(dest) for dest in destinations},
}
return {"parser": parser}
@@ -1,3 +1,4 @@
import json
from pathlib import Path
import pytest
@@ -15,6 +16,22 @@ def result_files(tmp_path: Path) -> Path:
"ref/latest_checkpointed_iteration.txt": "3\n",
"data.jsonl": '{"prompt": "1+1", "label": "2"}\n',
"data2.jsonl": '{"prompt": "2+2", "label": "4"}\n',
"hf/config.json": json.dumps(
{
"architectures": ["LlamaForCausalLM"],
"model_type": "llama",
"hidden_size": 128,
"intermediate_size": 512,
"num_attention_heads": 2,
"num_key_value_heads": 2,
"num_hidden_layers": 1,
"vocab_size": 1024,
"max_position_embeddings": 4096,
"rms_norm_eps": 1e-05,
"rope_theta": 10000.0,
"tie_word_embeddings": True,
}
),
}
for name, content in files.items():
path = tmp_path / name
@@ -44,6 +44,12 @@ def result_scenarios() -> dict[str, ResultScenario]:
error=NotImplementedError,
message="not supported on the megatron backend",
)
scenarios["fsdp_reject_lora"] = ResultScenario(
backend="fsdp",
arguments=("--lora-rank", "8", "--target-modules", "q_proj,v_proj"),
error=AssertionError,
message="LoRA injection is not implemented for FSDP",
)
scenarios["fsdp_true_on_policy"] = ResultScenario(backend="fsdp", arguments=("--true-on-policy-mode",))
scenarios["fsdp_prefill_logprobs"] = ResultScenario(
backend="fsdp", arguments=("--true-on-policy-mode", "--recompute-logprobs-via-prefill")
@@ -58,6 +64,9 @@ def result_scenarios() -> dict[str, ResultScenario]:
return scenarios
_HF_CHECKPOINT = ("--hf-checkpoint", "$FIXTURES/hf", "--ffn-hidden-size", "512")
def _shared_variants() -> dict[str, tuple[str, ...]]:
return {
"learning_rate": ("--lr", "0.000003"),
@@ -84,7 +93,14 @@ def _shared_variants() -> dict[str, tuple[str, ...]]:
"load_rollout": ("--load-debug-rollout-data", "$FIXTURES/rollout.pt"),
"save": ("--save", "$FIXTURES/save", "--save-interval", "2"),
"dump_details": ("--dump-details", "$FIXTURES/dump"),
"explicit_event_directory": ("--save", "$FIXTURES/save", "--save-debug-event-data", "$FIXTURES/events"),
"explicit_event_directory": (
"--save",
"$FIXTURES/save",
"--save-interval",
"2",
"--save-debug-event-data",
"$FIXTURES/events",
),
"load_fallback": ("--load", "$FIXTURES/missing", "--ref-load", "$FIXTURES/ref", "--ref-ckpt-step", "7"),
"load_existing": ("--load", "$FIXTURES/checkpoint"),
"custom_yaml": ("--custom-config-path", "$FIXTURES/custom.yaml", "--lr", "0.000003"),
@@ -145,8 +161,15 @@ def _shared_variants() -> dict[str, tuple[str, ...]]:
"2",
),
"sglang_tp_derived": ("--sglang-tp-size", "8", "--rollout-num-gpus-per-engine", "2"),
"rollout_ft": ("--use-fault-tolerance",),
"ft_explicit_api": ("--use-fault-tolerance", "--api-server-port", "23456", "--no-mini-ft-controller-enable"),
"rollout_ft": ("--use-fault-tolerance", "--update-weight-transfer-mode", "p2p"),
"ft_explicit_api": (
"--use-fault-tolerance",
"--update-weight-transfer-mode",
"p2p",
"--api-server-port",
"23456",
"--no-mini-ft-controller-enable",
),
"opd_external": ("--use-opd", "--opd-type", "sglang"),
"rollout_logprobs": ("--use-rollout-logprobs",),
"ci": ("--ci-test", "--no-enable-sample-ownership-checker", "--save-debug-event-data", "$FIXTURES/events"),
@@ -158,7 +181,6 @@ def _shared_variants() -> dict[str, tuple[str, ...]]:
),
"ray_rpc": ("--worker-comm-backend", "rpc"),
"external_rollout": ("--rollout-external-engine-addrs", "127.0.0.1:30000"),
"lora_targets": ("--lora-rank", "8", "--target-modules", "q_proj,v_proj", "--exclude-modules", "v_proj"),
}
@@ -181,7 +203,26 @@ def _megatron_variants() -> dict[str, tuple[str, ...]]:
"2",
),
"ppo": ("--advantage-estimator", "ppo"),
"ppo_save": ("--advantage-estimator", "ppo", "--save", "$FIXTURES/save", "--critic-lr", "0.000001"),
"ppo_save": (
"--advantage-estimator",
"ppo",
"--save",
"$FIXTURES/save",
"--save-interval",
"2",
"--critic-lr",
"0.000001",
),
"lora_targets": (
*_HF_CHECKPOINT,
"--lora-rank",
"8",
"--target-modules",
"attn",
"--exclude-modules",
"v_proj",
),
"lora_default_targets": (*_HF_CHECKPOINT, "--lora-rank", "8"),
"offload_disk": (
"--offload-train",
"--offload-train-target",
@@ -285,7 +326,6 @@ def _rejected_variants() -> dict[str, tuple[tuple[str, ...], type[Exception], st
AssertionError,
"not compatible with --colocate",
),
"lora_without_targets": (("--lora-rank", "8"), AssertionError, "'--target-modules' is required"),
"mini_ft_without_api": (
("--mini-ft-controller-enable", "--api-server-port", "0"),
ValueError,
@@ -49,4 +49,5 @@ def fsdp_debug_actor() -> actor_module.FSDPTrainRayActor:
actor = object.__new__(actor_module.FSDPTrainRayActor)
actor._heartbeat = SimpleHeartbeat()
actor._init_once = InitOnce(type(actor).__name__)
actor._config_snapshot_train_recorded = True
return actor
+6 -2
View File
@@ -4,6 +4,7 @@ from contextlib import nullcontext
from types import ModuleType
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args
from miles.backends.fsdp_utils import actor as actor_module
from miles.backends.megatron_utils.ft.types import TrainStepOutcome, TrainStepOutput
@@ -55,7 +56,8 @@ class TestFSDPInit:
actor._rank = 0
actor._heartbeat = SimpleHeartbeat()
actor._init_once = InitOnce(type(actor).__name__)
args = Namespace(
args = make_trainer_args(
train_backend="fsdp",
debug_deterministic_collective=False,
distributed_backend="nccl",
distributed_timeout_minutes=1,
@@ -91,7 +93,9 @@ class TestFSDPTrainParallelConfigWiring:
fsdp_debug_actor: actor_module.FSDPTrainRayActor,
) -> None:
"""FSDP init records its live DP layout without precomputed scheduling and train hands it to the loader."""
args = Namespace(dumper_enable=False, seed=0, offload_train=False, debug_rollout_only=True)
args = make_trainer_args(
train_backend="fsdp", dumper_enable=False, seed=0, offload_train=False, debug_rollout_only=True
)
received: list[TrainParallelConfig] = []
def load_rollout_data(
@@ -8,6 +8,7 @@ from typing import Any
from unittest.mock import Mock
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args
from miles.utils.init_once import InitOnce
@@ -94,7 +95,7 @@ def _args(tmp_path: Path, **overrides) -> Namespace:
global_batch_size=1,
)
defaults.update(overrides)
return Namespace(**defaults)
return make_trainer_args(**defaults)
def _write_checkpoint(directory: Path, *, iteration: int) -> str:
@@ -133,10 +134,10 @@ def _watch_load(actor_module, monkeypatch, *, args: Namespace, iteration: int) -
seen: dict[str, Any] = {}
def fake_load_checkpoint(*_args: Any, **_kwargs: Any) -> _CheckpointLoadResult:
seen["args_during_load"] = vars(args).copy()
seen["args_during_load"] = vars(args.backend).copy()
return _CheckpointLoadResult(
iteration=iteration,
restored_trained_iteration=not args.finetune or iteration > 0,
restored_trained_iteration=not args.backend.finetune or iteration > 0,
native_optimizer_restored=False,
)
@@ -250,8 +251,18 @@ class TestTheCheckpointAReloadRollsBackTo:
_actor(actor_module, role="actor", args=args).load_state()
assert args.load == str(tmp_path / "pretrain")
assert (args.finetune, args.no_load_optim, args.no_load_rng, args.ckpt_step) == (True, True, True, 3)
assert args.backend.load == str(tmp_path / "pretrain")
assert (
args.backend.finetune,
args.backend.no_load_optim,
args.backend.no_load_rng,
args.backend.ckpt_step,
) == (
True,
True,
True,
3,
)
def test_a_critic_reload_reads_the_critic_directory(self, actor_module, tmp_path, monkeypatch):
"""A critic's own arguments carry its checkpoint dirs, so reading them reads the critic's."""
@@ -10,6 +10,8 @@ import pytest
from tests.ci.ci_register import register_cpu_ci
from miles.utils.args.runtime import TrainerConfig
register_cpu_ci(est_time=1, suite="stage-a-cpu", labels=[])
@@ -21,14 +23,14 @@ def apply_bridge_runtime_config() -> Callable:
function = next(
node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "_apply_bridge_runtime_config"
)
namespace = {"argparse": argparse}
namespace = {"argparse": argparse, "TrainerConfig": TrainerConfig}
exec(compile(ast.Module(body=[function], type_ignores=[]), str(path), "exec"), namespace)
return namespace["_apply_bridge_runtime_config"]
@pytest.fixture
def runtime_args() -> argparse.Namespace:
return argparse.Namespace(
backend = argparse.Namespace(
tensor_model_parallel_size=1,
pipeline_model_parallel_size=1,
expert_model_parallel_size=1,
@@ -52,7 +54,12 @@ def runtime_args() -> argparse.Namespace:
fp8_recipe=None,
attention_backend="auto",
moe_token_dispatcher_type="alltoall",
decoder_first_pipeline_num_layers=None,
decoder_last_pipeline_num_layers=None,
moe_router_bias_update_rate=None,
moe_aux_loss_coeff=None,
)
return argparse.Namespace(backend=backend, dsa_attention_backend=None)
@pytest.mark.parametrize(
@@ -62,20 +69,16 @@ def runtime_args() -> argparse.Namespace:
(True, True, True),
(False, False, False),
(False, True, True),
(None, False, False),
(None, True, True),
],
)
def test_bridge_mtp_detachment(
apply_bridge_runtime_config: Callable,
runtime_args: argparse.Namespace,
enabled: bool | None,
enabled: bool,
initial_detach: bool,
expected_detach: bool,
) -> None:
# A missing flag covers callers that only register Megatron's arguments.
if enabled is not None:
runtime_args.enable_mtp_training = enabled
runtime_args.enable_mtp_training = enabled
provider = SimpleNamespace(mtp_num_layers=1, mtp_detach_heads=initial_detach)
apply_bridge_runtime_config(provider, runtime_args)
@@ -1,8 +1,8 @@
from argparse import Namespace
from pathlib import Path
from types import SimpleNamespace
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args
from miles.backends.megatron_utils import checkpoint
@@ -37,8 +37,9 @@ class TestCheckpointTrainingProvenance:
(load_path / "metadata.json").write_text("{}")
else:
(load_path / "latest_checkpointed_iteration.txt").write_text(source)
args = Namespace(load=str(load_path), finetune=finetune, ckpt_step=ckpt_step, lora_rank=0)
monkeypatch.setattr(checkpoint, "get_args", lambda: args)
args = make_trainer_args(
load=str(load_path), finetune=finetune, ckpt_step=ckpt_step, lora_rank=0, lora_adapter_path=None
)
monkeypatch.setattr(checkpoint, "_load_checkpoint_megatron", lambda **_kwargs: (0, 123))
checkpointing_context = (
None
@@ -52,17 +53,20 @@ class TestCheckpointTrainingProvenance:
opt_param_scheduler=None,
checkpointing_context=checkpointing_context,
skip_load_to_model_and_opt=False,
args=args,
)
assert result == (0, expected, False)
def test_loading_hf_weights_preserves_non_finetune_rollout_numbering(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
@pytest.mark.parametrize("finetune", [False, True])
def test_loading_hf_weights_never_claims_restored_training_progress(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch, finetune: bool
) -> None:
"""HF loading preserves the existing non-finetune rollout numbering."""
"""HF weights initialize a model without restoring any completed rollout."""
(tmp_path / "config.json").write_text("{}")
args = Namespace(load=str(tmp_path), finetune=False, ckpt_step=None, lora_rank=0)
monkeypatch.setattr(checkpoint, "get_args", lambda: args)
args = make_trainer_args(
load=str(tmp_path), finetune=finetune, ckpt_step=None, lora_rank=0, lora_adapter_path=None
)
monkeypatch.setattr(checkpoint, "_load_checkpoint_hf", lambda **_kwargs: (0, 0))
result = checkpoint.load_checkpoint(
@@ -71,10 +75,12 @@ class TestCheckpointTrainingProvenance:
opt_param_scheduler=None,
checkpointing_context=None,
skip_load_to_model_and_opt=False,
args=args,
)
assert result == (0, True, False)
assert result == (0, False, False)
@pytest.mark.parametrize("source", ["native", "hf"])
@pytest.mark.parametrize(
("adapter_result", "expected"),
[
@@ -88,12 +94,16 @@ class TestCheckpointTrainingProvenance:
self,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
source: str,
adapter_result: tuple[bool, int | None, bool],
expected: bool,
) -> None:
"""Adapter weights alone cannot be mistaken for a saved iteration-zero training state."""
(tmp_path / "latest_checkpointed_iteration.txt").write_text("release")
args = Namespace(
if source == "native":
(tmp_path / "latest_checkpointed_iteration.txt").write_text("release")
else:
(tmp_path / "config.json").write_text("{}")
args = make_trainer_args(
load=str(tmp_path),
finetune=True,
ckpt_step=None,
@@ -101,8 +111,8 @@ class TestCheckpointTrainingProvenance:
lora_adapter_path=str(tmp_path / "adapter"),
no_load_optim=False,
)
monkeypatch.setattr(checkpoint, "get_args", lambda: args)
monkeypatch.setattr(checkpoint, "_load_checkpoint_megatron", lambda **_kwargs: (0, 123))
monkeypatch.setattr(checkpoint, "_load_checkpoint_hf", lambda **_kwargs: (0, 0))
monkeypatch.setattr(checkpoint, "load_lora_adapter", lambda *_args, **_kwargs: adapter_result)
result = checkpoint.load_checkpoint(
@@ -111,6 +121,7 @@ class TestCheckpointTrainingProvenance:
opt_param_scheduler=None,
checkpointing_context=None,
skip_load_to_model_and_opt=False,
args=args,
)
assert result == (0, expected, adapter_result[2])
@@ -9,6 +9,7 @@ import torch
import torch.distributed as dist
from megatron.core.dist_checkpointing.mapping import ShardedTensor
from megatron.core.dist_checkpointing.tensor_aware_state_dict import MCoreTensorAwareStateDict
from tests.fast.fixtures.args_fixtures import make_trainer_args
from torch.utils._pytree import tree_flatten_with_path, tree_unflatten
from miles.backends.megatron_utils.ft import checkpoint_transfer, in_memory_checkpoint
@@ -228,6 +229,7 @@ class TestSendCkptRecords:
with caplog.at_level(logging.INFO, logger=_CKPT_TRANSFER_LOGGER):
checkpoint_transfer.send_ckpt(
args=make_trainer_args(),
indep_dp=GroupInfo(rank=0, size=2, group=None),
model=[],
optimizer=object(),
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
import pytest
import torch
from tests.fast.fixtures.args_fixtures import make_trainer_args
from miles.backends.megatron_utils.ft import indep_dp
from miles.utils.ft_utils.indep_dp import IndepDPInfo
@@ -129,7 +130,7 @@ class TestAllreduceGradsAndLossesAcrossReplicas:
@staticmethod
def _run(pg, util) -> tuple[bool, dict[str, float]]:
args = SimpleNamespace(calculate_per_token_loss=False)
args = make_trainer_args(calculate_per_token_loss=False)
with patch.object(indep_dp.GeneralPGUtil, "create", return_value=util):
return indep_dp.allreduce_grads_and_losses_across_replicas(
args, [_make_model_chunk()], _make_parallel_state(pg), losses_reduced=[]
@@ -157,7 +158,7 @@ class TestAllreduceGradsAndLossesAcrossReplicas:
pg = SimpleNamespace(errored=lambda: None)
util = FakeCrossCellPGUtil()
calls = []
args = SimpleNamespace(calculate_per_token_loss=False)
args = make_trainer_args(calculate_per_token_loss=False)
with patch.object(indep_dp.GeneralPGUtil, "create", return_value=util):
consensus, _ = indep_dp.allreduce_grads_and_losses_across_replicas(
@@ -175,7 +176,7 @@ class TestAllreduceGradsAndLossesAcrossReplicas:
"""A failed metadata collective cannot leave a successful optimizer step."""
pg = SimpleNamespace(errored=lambda: None)
util = FakeCrossCellPGUtil()
args = SimpleNamespace(calculate_per_token_loss=False)
args = make_trainer_args(calculate_per_token_loss=False)
with patch.object(indep_dp.GeneralPGUtil, "create", return_value=util):
consensus, _ = indep_dp.allreduce_grads_and_losses_across_replicas(
@@ -62,17 +62,19 @@ class TestSetRandomSeedFromArgs:
_patch_state(monkeypatch, _parallel_state())
monkeypatch.setattr(initialize.tensor_parallel, "model_parallel_cuda_manual_seed", lambda *args: None)
args = SimpleNamespace(
rank=0,
seed=1729,
data_parallel_random_init=True,
te_rng_tracker=False,
inference_rng_tracker=True,
backend=SimpleNamespace(
rank=0,
seed=1729,
data_parallel_random_init=True,
te_rng_tracker=False,
inference_rng_tracker=True,
)
)
with caplog.at_level(logging.INFO, logger=initialize.__name__):
initialize.set_random_seed_from_args(args)
first_values = (random.random(), np.random.random(), torch.rand(1).item())
args.rank = 1
args.backend.rank = 1
initialize.set_random_seed_from_args(args)
second_values = (random.random(), np.random.random(), torch.rand(1).item())
@@ -5,10 +5,10 @@ save_checkpoint_with_lora / load_checkpoint — the latter using mocks to avoid
GPU / distributed requirements.
"""
from argparse import Namespace
from unittest.mock import MagicMock, patch
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args
from miles.backends.megatron_utils.checkpoint import _is_megatron_checkpoint, save_checkpoint_with_lora
@@ -69,28 +69,26 @@ class TestIsMegatronCheckpoint:
class TestSaveCheckpointWithLoRA:
@patch("miles.backends.megatron_utils.checkpoint.get_args")
@patch("miles.backends.megatron_utils.checkpoint.save_lora_checkpoint")
@patch("miles.backends.megatron_utils.checkpoint.is_lora_model", return_value=True)
def test_lora_model_saves_adapter(self, mock_is_lora, mock_save_lora, mock_get_args, tmp_path):
mock_get_args.return_value = Namespace(save=str(tmp_path))
def test_lora_model_saves_adapter(self, mock_is_lora, mock_save_lora, tmp_path):
args = make_trainer_args(save=str(tmp_path))
model = [MagicMock()]
publisher = MagicMock()
save_checkpoint_with_lora(42, model, MagicMock(), MagicMock(), publisher=publisher)
save_checkpoint_with_lora(42, model, MagicMock(), MagicMock(), args=args, publisher=publisher)
mock_save_lora.assert_called_once()
call_args = mock_save_lora.call_args
assert call_args.kwargs["publisher"] is publisher
assert "adapter" in call_args[1].get("save_dir", call_args[0][2] if len(call_args[0]) > 2 else "")
@patch("miles.backends.megatron_utils.checkpoint.get_args")
@patch("miles.backends.megatron_utils.checkpoint.save_checkpoint")
@patch("miles.backends.megatron_utils.checkpoint.is_lora_model", return_value=False)
def test_non_lora_model_saves_regular(self, mock_is_lora, mock_save_ckpt, mock_get_args, tmp_path):
mock_get_args.return_value = Namespace(save=str(tmp_path))
def test_non_lora_model_saves_regular(self, mock_is_lora, mock_save_ckpt, tmp_path):
args = make_trainer_args(save=str(tmp_path))
model = [MagicMock()]
save_checkpoint_with_lora(42, model, MagicMock(), MagicMock())
save_checkpoint_with_lora(42, model, MagicMock(), MagicMock(), args=args)
mock_save_ckpt.assert_called_once()
@@ -4,10 +4,10 @@ Validates that setup_model_and_optimizer, save, and save_hf_model correctly
route to LoRA-specific code paths depending on configuration — without GPU.
"""
from argparse import Namespace
from unittest.mock import MagicMock, patch
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args, make_trainer_config
# ---------------------------------------------------------------------------
# _ensure_model_list
@@ -39,19 +39,19 @@ class TestShouldDisableForwardPreHook:
def test_both_true(self):
from miles.backends.megatron_utils.model import should_disable_forward_pre_hook
args = Namespace(use_distributed_optimizer=True, overlap_param_gather=True)
args = make_trainer_args(use_distributed_optimizer=True, overlap_param_gather=True)
assert should_disable_forward_pre_hook(args) is True
def test_optimizer_false(self):
from miles.backends.megatron_utils.model import should_disable_forward_pre_hook
args = Namespace(use_distributed_optimizer=False, overlap_param_gather=True)
args = make_trainer_args(use_distributed_optimizer=False, overlap_param_gather=True)
assert should_disable_forward_pre_hook(args) is False
def test_overlap_false(self):
from miles.backends.megatron_utils.model import should_disable_forward_pre_hook
args = Namespace(use_distributed_optimizer=True, overlap_param_gather=False)
args = make_trainer_args(use_distributed_optimizer=True, overlap_param_gather=False)
assert should_disable_forward_pre_hook(args) is False
@@ -66,7 +66,7 @@ class TestSetupModelAndOptimizerLoraBranch:
"""Verify that LoRA-enabled actor + bridge mode routes to _setup_lora_model_via_bridge."""
def _make_args(self, lora_rank=32, role="actor", mode="bridge"):
return Namespace(
return make_trainer_config(
lora_rank=lora_rank,
lora_adapter_path=None,
custom_model_provider_path=None,
@@ -179,16 +179,15 @@ class TestSaveLoRaBranch:
@patch(f"{_MODEL_MODULE}.enable_forward_pre_hook")
@patch(f"{_MODEL_MODULE}.disable_forward_pre_hook")
@patch(f"{_MODEL_MODULE}.should_disable_forward_pre_hook", return_value=False)
@patch(f"{_MODEL_MODULE}.get_args")
@patch(f"{_MODEL_MODULE}.save_checkpoint_with_lora")
@patch(f"{_MODEL_MODULE}.is_lora_model", return_value=True)
def test_lora_model_calls_lora_save(
self, mock_is_lora, mock_save_lora, mock_get_args, mock_should, mock_disable, mock_enable, mock_save_hashes
self, mock_is_lora, mock_save_lora, mock_should, mock_disable, mock_enable, mock_save_hashes
):
from miles.backends.megatron_utils.model import save
model = [MagicMock()]
save(42, model, MagicMock(), MagicMock())
save(make_trainer_args(), 42, model, MagicMock(), MagicMock())
mock_save_lora.assert_called_once()
@@ -196,15 +195,14 @@ class TestSaveLoRaBranch:
@patch(f"{_MODEL_MODULE}.enable_forward_pre_hook")
@patch(f"{_MODEL_MODULE}.disable_forward_pre_hook")
@patch(f"{_MODEL_MODULE}.should_disable_forward_pre_hook", return_value=False)
@patch(f"{_MODEL_MODULE}.get_args")
@patch(f"{_MODEL_MODULE}.save_checkpoint")
@patch(f"{_MODEL_MODULE}.is_lora_model", return_value=False)
def test_non_lora_model_calls_regular_save(
self, mock_is_lora, mock_save_ckpt, mock_get_args, mock_should, mock_disable, mock_enable, mock_save_hashes
self, mock_is_lora, mock_save_ckpt, mock_should, mock_disable, mock_enable, mock_save_hashes
):
from miles.backends.megatron_utils.model import save
model = [MagicMock()]
save(42, model, MagicMock(), MagicMock())
save(make_trainer_args(), 42, model, MagicMock(), MagicMock())
mock_save_ckpt.assert_called_once()
@@ -5,6 +5,7 @@ from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.args_fixtures import make_trainer_args
import miles.backends.megatron_utils.lora.utils as lora_utils
from miles.backends.megatron_utils.lora.utils import (
@@ -37,10 +38,6 @@ class TestIsLoraEnabled:
args = Namespace(lora_rank=0, lora_adapter_path=None)
assert is_lora_enabled(args) is False
def test_disabled_missing_attrs(self):
args = Namespace()
assert is_lora_enabled(args) is False
# ---------------------------------------------------------------------------
# is_lora_weight_name / _is_adapter_param_name
@@ -146,7 +143,7 @@ class TestSaveLoraCheckpointTrainingState:
adapter = torch.nn.Parameter(torch.ones(2))
model = [SimpleNamespace(named_parameters=lambda: [("layers.0.self_attention.lora_A.weight", adapter)])]
args = Namespace(megatron_to_hf_mode="bridge", no_save_optim=no_save_optim)
args = make_trainer_args(megatron_to_hf_mode="bridge", no_save_optim=no_save_optim)
optimizer = SimpleNamespace(state_dict=lambda: {"step": 7})
save_lora_checkpoint(
model,
@@ -33,7 +33,7 @@ def _write_yaml(data: dict, tmp_path) -> str:
def _model_args(args: Namespace, *, model_id: str) -> Namespace:
return compute_trainer_args(args, resolve_megatron_config(args).get(model_id))
return compute_trainer_args(args, resolve_megatron_config(args, base_args={}).get(model_id))
def _make_args(megatron_config: str | None = None, **overrides) -> Namespace:
@@ -74,6 +74,8 @@ def _make_args(megatron_config: str | None = None, **overrides) -> Namespace:
critic_save=None,
critic_lr=None,
critic_lr_warmup_iters=None,
critic_num_nodes=1,
critic_num_gpus_per_node=8,
fp16=False,
seq_length=4096,
vocab_size=None,
@@ -89,7 +91,7 @@ def _make_args(megatron_config: str | None = None, **overrides) -> Namespace:
class TestResolveMegatronConfig:
def test_a_run_without_the_flag_synthesizes_a_plain_actor_trainer(self):
"""Legacy single policy runs must keep working, with no model id anywhere downstream."""
config = resolve_megatron_config(_make_args())
config = resolve_megatron_config(_make_args(), base_args={})
assert [(t.trainer_id, t.role, t.model_id, t.overrides) for t in config.trainers] == [
("actor", "actor", None, {})
@@ -101,13 +103,13 @@ class TestResolveMegatronConfig:
"""Configs written against the first name of the field must keep resolving."""
path = _write_yaml({"megatron": [{"model_id": "a"}, {"model_id": "b"}]}, tmp_path)
assert resolve_megatron_config(_make_args(path)).model_ids == ["a", "b"]
assert resolve_megatron_config(_make_args(path), base_args={}).model_ids == ["a", "b"]
def test_the_yaml_model_ids_become_the_trainer_model_ids(self, tmp_path):
"""The `model_id` field is the source of truth for trainer_model_id and spec names."""
path = _write_yaml({"trainers": [{"model_id": "a", "overrides": {"lr": 1e-5}}, {"model_id": "b"}]}, tmp_path)
config = resolve_megatron_config(_make_args(path))
config = resolve_megatron_config(_make_args(path), base_args={})
assert config.model_ids == ["a", "b"]
assert config.leader_model_id == "a"
@@ -118,17 +120,17 @@ class TestResolveMegatronConfig:
path = _write_yaml({"trainers": [{"model_id": "eval"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="shared eval rollouts"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_the_first_model_is_the_leader_policy(self, tmp_path):
"""The leader owns the global checkpoint index, so its identity must be positional and stable."""
path = _write_yaml({"trainers": [{"model_id": "second"}, {"model_id": "first"}]}, tmp_path)
assert resolve_megatron_config(_make_args(path)).leader_model_id == "second"
assert resolve_megatron_config(_make_args(path), base_args={}).leader_model_id == "second"
def test_an_inline_base64_payload_is_accepted(self, tmp_path):
"""Launchers that cannot ship a file still need to pass the config."""
config = resolve_megatron_config(_make_args(encode_megatron_config("solo")))
config = resolve_megatron_config(_make_args(encode_megatron_config("solo")), base_args={})
assert config.model_ids == ["solo"]
@@ -137,13 +139,13 @@ class TestResolveMegatronConfig:
path = _write_yaml({"trainers": [{"model_id": "a"}, {"model_id": "a"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="trainer ids must be unique"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_a_trainer_id_defaults_to_the_model_id_and_the_role(self, tmp_path):
"""The trainer id addresses a pool, so its default must stay the name every deployment already uses."""
path = _write_yaml({"trainers": [{"model_id": "a"}, {"model_id": "b"}]}, tmp_path)
config = resolve_megatron_config(_make_args(path))
config = resolve_megatron_config(_make_args(path), base_args={})
assert [trainer.trainer_id for trainer in config.trainers] == ["a-actor", "b-actor"]
assert [trainer.role for trainer in config.trainers] == ["actor", "actor"]
@@ -152,21 +154,21 @@ class TestResolveMegatronConfig:
"""A deployment that already named its pools must be able to keep those names."""
path = _write_yaml({"trainers": [{"model_id": "a", "trainer_id": "legacy-actor"}]}, tmp_path)
assert resolve_megatron_config(_make_args(path)).trainers[0].trainer_id == "legacy-actor"
assert resolve_megatron_config(_make_args(path), base_args={}).trainers[0].trainer_id == "legacy-actor"
def test_an_explicit_trainer_id_colliding_with_a_derived_one_is_refused(self, tmp_path):
"""Uniqueness has to hold across both spellings, or two trainers would share one engine pool."""
path = _write_yaml({"trainers": [{"model_id": "a"}, {"model_id": "b", "trainer_id": "a-actor"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="trainer ids must be unique"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_a_trainer_id_that_is_not_a_dns_label_is_refused(self, tmp_path):
"""A trainer id is embedded in Kubernetes pool names, which must be lowercase DNS labels."""
path = _write_yaml({"trainers": [{"model_id": "a", "trainer_id": "Legacy_Actor"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="trainer ids"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_a_trainer_id_too_long_for_its_controller_pool_is_refused(self, tmp_path):
"""A trainer id must leave room for the controller suffix in its Kubernetes pool name."""
@@ -175,13 +177,13 @@ class TestResolveMegatronConfig:
)
with pytest.raises(pydantic.ValidationError, match="longer than"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_several_entries_of_one_model_id_are_not_a_multi_policy_run(self, tmp_path):
"""An actor and a critic of one policy share its id, and one policy is not several policies."""
path = _write_yaml({"trainers": [{"model_id": "a"}]}, tmp_path)
config = resolve_megatron_config(_make_args(path, use_critic=True))
config = resolve_megatron_config(_make_args(path, use_critic=True), base_args={})
assert [trainer.trainer_id for trainer in config.trainers] == ["a-actor", "a-critic"]
assert config.model_ids == ["a"]
@@ -193,34 +195,34 @@ class TestResolveMegatronConfig:
path = _write_yaml({"trainers": [{"model_id": "a", "override": {"lr": 1e-5}}]}, tmp_path)
with pytest.raises(Exception, match="override"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_getting_an_unknown_model_id_fails_loudly(self, tmp_path):
"""Callers routing by model id must not silently fall back to another policy."""
path = _write_yaml({"trainers": [{"model_id": "a"}]}, tmp_path)
with pytest.raises(KeyError, match="Unknown trainer model id"):
resolve_megatron_config(_make_args(path)).get("b")
resolve_megatron_config(_make_args(path), base_args={}).get("b")
def test_a_config_declaring_no_trainer_is_refused(self, tmp_path):
"""An empty list would resolve to a run with nothing to train, and fail much later and less clearly."""
path = _write_yaml({"trainers": []}, tmp_path)
with pytest.raises(AssertionError, match="must declare at least one trainer"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_getting_a_model_id_answers_its_first_trainer(self, tmp_path):
"""Callers ask by model id and expect the actor: the critic of that policy is addressed by role."""
path = _write_yaml({"trainers": [{"model_id": "a"}]}, tmp_path)
config = resolve_megatron_config(_make_args(path, use_critic=True))
config = resolve_megatron_config(_make_args(path, use_critic=True), base_args={})
assert [trainer.role for trainer in config.trainers] == ["actor", "critic"]
assert config.get("a").role == "actor"
def test_a_run_without_the_flag_has_no_leader_model_id(self):
"""A single policy run has no leader to index the trainers by, and must answer None rather than invent one."""
assert resolve_megatron_config(_make_args()).leader_model_id is None
assert resolve_megatron_config(_make_args(), base_args={}).leader_model_id is None
class TestDerivedPerPolicyArgs:
@@ -229,7 +231,7 @@ class TestDerivedPerPolicyArgs:
path = _write_yaml({"trainers": [{"model_id": "../evil"}, {"model_id": "b"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="not usable as Kubernetes pool names"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
@pytest.mark.parametrize("model_id", ["policy_a", "PolicyA", "-policy", "policy-", "policy.a"])
def test_a_model_id_that_is_not_a_dns_label_is_refused(self, tmp_path, model_id):
@@ -237,14 +239,14 @@ class TestDerivedPerPolicyArgs:
path = _write_yaml({"trainers": [{"model_id": model_id}, {"model_id": "b"}]}, tmp_path)
with pytest.raises(pydantic.ValidationError, match="not usable as Kubernetes pool names"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
@pytest.mark.parametrize("model_id", ["default", "policy-a", "a1", "a-b-c"])
def test_lowercase_dns_labels_are_accepted(self, tmp_path, model_id):
"""The ids the docs and examples use must survive validation."""
path = _write_yaml({"trainers": [{"model_id": model_id}, {"model_id": "other"}]}, tmp_path)
assert resolve_megatron_config(_make_args(path)).model_ids == [model_id, "other"]
assert resolve_megatron_config(_make_args(path), base_args={}).model_ids == [model_id, "other"]
class TestOverrideCoercion:
@@ -254,7 +256,7 @@ class TestOverrideCoercion:
{"trainers": [{"model_id": "a", "overrides": {"lr": "5e-7", "global_batch_size": "128"}}]}, tmp_path
)
overrides = resolve_megatron_config(_make_args(path)).get("a").overrides
overrides = resolve_megatron_config(_make_args(path), base_args={}).get("a").overrides
assert overrides == {"lr": 5e-7, "global_batch_size": 128}
@@ -265,21 +267,21 @@ class TestOverrideCoercion:
)
with pytest.raises(AssertionError, match="not a boolean"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_an_override_without_a_value_is_refused(self, tmp_path):
"""A key written with an empty YAML value reads as None, which no argument can be set to."""
path = _write_yaml({"trainers": [{"model_id": "a", "overrides": {"eps_clip_high": None}}]}, tmp_path)
with pytest.raises(AssertionError, match="no value"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_an_argument_outside_the_per_policy_whitelist_is_refused(self, tmp_path):
"""Rhythm arguments are read from the base command line, so accepting them here would do nothing."""
path = _write_yaml({"trainers": [{"model_id": "a", "overrides": {"num_rollout": 3}}]}, tmp_path)
with pytest.raises(AssertionError, match="num_rollout"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
class TestModelDefinitionOverrides:
@@ -506,7 +508,7 @@ class TestComputeTrainerArgs:
path = _write_yaml({"trainers": [{"model_id": "a", "overrides": {"no_such_flag": 3}}]}, tmp_path)
with pytest.raises(AssertionError, match="no_such_flag"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
def test_a_whitelisted_argument_this_run_does_not_declare_is_refused(self, tmp_path):
"""A whitelist entry is not a promise that every backend's parser declares it."""
@@ -556,7 +558,10 @@ class TestTrainerCheckpointDirs:
path = _write_yaml({"trainers": [{"model_id": "a", "trainer_id": "a-second"}, {"model_id": "b"}]}, tmp_path)
args = _make_args(path, save="/ckpt/run")
saves = [compute_trainer_args(args, trainer).save for trainer in resolve_megatron_config(args).trainers]
saves = [
compute_trainer_args(args, trainer).save
for trainer in resolve_megatron_config(args, base_args={}).trainers
]
assert saves == ["/ckpt/run/trainers/a-second", "/ckpt/run/trainers/b-actor"]
@@ -592,7 +597,7 @@ class TestTrainerCheckpointDirs:
critic_save="/ckpt/critic",
critic_load=str(old),
)
[_, critic] = resolve_megatron_config(args).trainers
[_, critic] = resolve_megatron_config(args, base_args={}).trainers
model = compute_trainer_args(args, critic)
@@ -618,7 +623,7 @@ class TestTrainerCheckpointDirs:
)
with pytest.raises(AssertionError, match="sets 'load', which it may not override"):
resolve_megatron_config(_make_args(path))
resolve_megatron_config(_make_args(path), base_args={})
class TestPerPolicyCheckpointResolution:
@@ -697,7 +702,7 @@ class TestMultiPolicyIds:
def test_a_run_without_a_megatron_config_carries_no_trainer_model_id(self):
"""The unnamed legacy actor keeps using the unnamespaced single-policy key."""
args = _make_args()
[trainer] = resolve_megatron_config(args).trainers
[trainer] = resolve_megatron_config(args, base_args={}).trainers
assert compute_trainer_args(args, trainer).trainer_model_id is None
@@ -715,16 +720,9 @@ class TestMultiPolicyIds:
class TestSynthesizedCriticTrainer:
def test_arguments_that_do_not_carry_use_critic_yet_still_resolve(self):
"""use_critic is derived while the arguments are validated, and the config is resolved before that."""
args = _make_args()
del args.use_critic
assert [trainer.role for trainer in resolve_megatron_config(args).trainers] == ["actor"]
def test_a_run_without_the_flag_synthesizes_the_critic_beside_the_actor(self):
"""The critic used to be assembled in specs and in the worker; the config is now the only source."""
config = resolve_megatron_config(_make_args(use_critic=True))
config = resolve_megatron_config(_make_args(use_critic=True), base_args={})
assert [(t.trainer_id, t.role, t.model_id) for t in config.trainers] == [
("actor", "actor", None),
@@ -736,7 +734,7 @@ class TestSynthesizedCriticTrainer:
path = _write_yaml({"trainers": [{"model_id": "a"}, {"model_id": "b"}]}, tmp_path)
with pytest.raises(AssertionError, match="does not support --use-critic"):
resolve_megatron_config(_make_args(path, use_critic=True))
resolve_megatron_config(_make_args(path, use_critic=True), base_args={})
def test_a_config_that_declares_its_own_critic_is_refused(self, tmp_path):
"""The critic checkpoint, learning rate and neutralized knobs only reach the synthesized critic, so a
@@ -752,13 +750,13 @@ class TestSynthesizedCriticTrainer:
)
with pytest.raises(AssertionError, match="declares a critic for"):
resolve_megatron_config(_make_args(path, use_critic=True))
resolve_megatron_config(_make_args(path, use_critic=True), base_args={})
def test_the_critic_of_a_named_policy_inherits_its_id_and_its_overlay(self, tmp_path):
"""The critic trains the same policy, so it must be addressed by that policy and see its settings."""
path = _write_yaml({"trainers": [{"model_id": "alpha", "overrides": {"eps_clip": 0.3}}]}, tmp_path)
[_, critic] = resolve_megatron_config(_make_args(path, use_critic=True)).trainers
[_, critic] = resolve_megatron_config(_make_args(path, use_critic=True), base_args={}).trainers
assert (critic.trainer_id, critic.model_id, critic.role) == ("alpha-critic", "alpha", "critic")
assert critic.overrides["eps_clip"] == 0.3
@@ -775,7 +773,7 @@ class TestSynthesizedCriticTrainer:
critic_lr_warmup_iters=3,
)
critic_args = compute_trainer_args(args, resolve_megatron_config(args).trainers[1])
critic_args = compute_trainer_args(args, resolve_megatron_config(args, base_args={}).trainers[1])
assert (critic_args.kl_coef, critic_args.use_opd, critic_args.disable_param_buffers_cpu_backup) == (
0,
@@ -793,7 +791,7 @@ class TestSynthesizedCriticTrainer:
"""The two trainers share one command line, so a leaked critic override would retrain the actor."""
args = _make_args(use_critic=True, save="/ckpt/run", load="/ckpt/run", critic_load="/ckpt/critic")
actor_args = compute_trainer_args(args, resolve_megatron_config(args).trainers[0])
actor_args = compute_trainer_args(args, resolve_megatron_config(args, base_args={}).trainers[0])
assert (actor_args.load, actor_args.save, actor_args.lr, actor_args.kl_coef) == (
"/ckpt/run",
@@ -806,7 +804,7 @@ class TestSynthesizedCriticTrainer:
"""kl_coef and load are not per-policy yaml arguments, yet the critic must still be able to set them."""
args = _make_args(use_critic=True, critic_load="/ckpt/critic")
overrides = set(resolve_megatron_config(args).trainers[1].overrides)
overrides = set(resolve_megatron_config(args, base_args={}).trainers[1].overrides)
assert {"kl_coef", "load"} <= overrides
assert not {"kl_coef", "load"} & set(PER_POLICY_ARGS)
@@ -816,13 +814,13 @@ class TestSynthesizedCriticTrainer:
path = _write_yaml({"trainers": [{"model_id": "alpha", "overrides": {"lr": 5e-7}}]}, tmp_path)
args = _make_args(path, use_critic=True)
critic_args = compute_trainer_args(args, resolve_megatron_config(args).trainers[1])
critic_args = compute_trainer_args(args, resolve_megatron_config(args, base_args={}).trainers[1])
assert (critic_args.lr, critic_args.lr_warmup_iters) == (None, None)
def test_the_critic_overlay_names_exactly_the_fields_the_worker_used_to_swap(self):
"""A new critic_* argument that nobody wires in here would be read from the command line and ignored."""
overrides = resolve_megatron_config(_make_args(use_critic=True)).trainers[1].overrides
overrides = resolve_megatron_config(_make_args(use_critic=True), base_args={}).trainers[1].overrides
assert set(overrides) == {
"kl_coef",
@@ -838,7 +836,11 @@ class TestSynthesizedCriticTrainer:
"""The overlay order is what neutralizes the critic, so a policy override of the same field must not win."""
path = _write_yaml({"trainers": [{"model_id": "alpha", "overrides": {"lr": 5e-7, "eps_clip": 0.3}}]}, tmp_path)
overrides = resolve_megatron_config(_make_args(path, use_critic=True, critic_lr=2e-6)).trainers[1].overrides
overrides = (
resolve_megatron_config(_make_args(path, use_critic=True, critic_lr=2e-6), base_args={})
.trainers[1]
.overrides
)
assert (overrides["lr"], overrides["eps_clip"]) == (2e-6, 0.3)
@@ -7,6 +7,7 @@ from unittest.mock import Mock
import pytest
import torch
from tests.fast.fixtures.args_fixtures import make_trainer_args, with_backend_values
from tests.fast.utils.test_utils.fault_injector.fakes import _arm_marker_hook
from miles.backends.training_utils.data import DataIterator
@@ -110,7 +111,7 @@ def make_train_one_step_args(**overrides: Any) -> Namespace:
enable_witness=False,
save_local_weight_checksum=False,
)
return Namespace(**{**defaults, **overrides})
return make_trainer_args(**{**defaults, **overrides})
@pytest.fixture
@@ -134,7 +135,6 @@ def train_one_step_env(monkeypatch) -> TrainOneStepEnv:
forward_backward_engine=FakeForwardBackwardEngine(),
)
monkeypatch.setattr(model_module, "get_args", lambda: env.args)
monkeypatch.setattr(model_module, "get_parallel_state", lambda: env.parallel_state)
monkeypatch.setattr("miles.backends.training_utils.parallel.get_parallel_state", lambda: env.parallel_state)
monkeypatch.setattr(model_module, "get_forward_backward_func", lambda: env.forward_backward_engine)
@@ -264,7 +264,7 @@ def test_forward_only_omits_sampling_mask_for_callbacks_that_do_not_replay_sampl
callback_kwargs.update(kwargs)
return {}
args = Namespace(
args = make_trainer_args(
data_pad_size_multiplier=1,
qkv_format="thd",
allgather_cp=False,
@@ -352,7 +352,7 @@ class TestTrainOneStepModelCompanion:
identity_fields["lineage_source_sample_indices"] = [7]
identity_fields["lineage_output_indices"] = [0]
identity_fields["lineage_output_counts"] = [1]
train_one_step_env.args.check_for_nan_in_loss_and_grad = False
train_one_step_env.args = with_backend_values(train_one_step_env.args, check_for_nan_in_loss_and_grad=False)
def forward_backward_engine(**kwargs: Any) -> list[dict[str, Any]]:
kwargs["data_iterator"][0].offset = 1
@@ -391,7 +391,7 @@ class TestTrainOneStepModelCompanion:
identity_fields["lineage_source_sample_indices"] = [7]
identity_fields["lineage_output_indices"] = [1]
identity_fields["lineage_output_counts"] = [2]
train_one_step_env.args.check_for_nan_in_loss_and_grad = False
train_one_step_env.args = with_backend_values(train_one_step_env.args, check_for_nan_in_loss_and_grad=False)
train_one_step_env.forward_backward_engine = FakeForwardBackwardEngine()
def forward_backward_engine(**kwargs: Any) -> list[dict[str, Any]]:
@@ -432,7 +432,7 @@ class TestTrainOneStepModelCompanion:
identity_fields["lineage_output_indices"] = [0]
identity_fields["lineage_output_counts"] = [1]
train_one_step_env.parallel_state.indep_dp.size = 2
train_one_step_env.args.check_for_nan_in_loss_and_grad = False
train_one_step_env.args = with_backend_values(train_one_step_env.args, check_for_nan_in_loss_and_grad=False)
def forward_backward_engine(**kwargs: Any) -> list[dict[str, Any]]:
kwargs["data_iterator"][0].offset = 1
@@ -475,8 +475,8 @@ def test_ft_discard_stays_invalid_with_finite_gradient_norm(
env = train_one_step_env
env.args.enable_sample_ownership_checker = False
env.args.check_for_nan_in_loss_and_grad = False
env.args.calculate_per_token_loss = False
env.args = with_backend_values(env.args, check_for_nan_in_loss_and_grad=False)
env.args = with_backend_values(env.args, calculate_per_token_loss=False)
env.parallel_state.indep_dp.size = 2
def reject_optimizer_step() -> None:
@@ -1,12 +1,12 @@
import sys
import types
from argparse import Namespace
from contextlib import ExitStack
from pathlib import Path
from typing import TYPE_CHECKING
from unittest.mock import MagicMock, patch
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args, make_trainer_config
if TYPE_CHECKING:
from miles.backends.megatron_utils.model import LoadCheckpointOutput
@@ -168,7 +168,7 @@ def _patch_initialize_side_effects(stack: ExitStack) -> None:
def test_initialize_does_not_step_scheduler_restored_from_checkpoint():
from miles.backends.megatron_utils.model import LoadCheckpointOutput, initialize_model_and_optimizer
args = Namespace(use_checkpoint_opt_param_scheduler=True, global_batch_size=8, finetune=False)
args = make_trainer_config(use_checkpoint_opt_param_scheduler=True, global_batch_size=8, finetune=False)
model = [_FakeModelChunk()]
optimizer = object()
opt_param_scheduler = MagicMock()
@@ -198,7 +198,7 @@ def test_initialize_does_not_step_scheduler_restored_from_checkpoint():
def test_initialize_steps_scheduler_when_checkpoint_did_not_restore_it():
from miles.backends.megatron_utils.model import LoadCheckpointOutput, initialize_model_and_optimizer
args = Namespace(use_checkpoint_opt_param_scheduler=False, global_batch_size=8, finetune=False)
args = make_trainer_config(use_checkpoint_opt_param_scheduler=False, global_batch_size=8, finetune=False)
model = [_FakeModelChunk()]
optimizer = object()
opt_param_scheduler = MagicMock()
@@ -250,7 +250,7 @@ def _load_model_state_with(
)
_patch_initialize_side_effects(stack)
return load_model_state(
Namespace(
make_trainer_args(
use_checkpoint_opt_param_scheduler=True,
global_batch_size=8,
finetune=finetune,
@@ -0,0 +1,95 @@
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from safetensors.torch import save_file
from tests.fast.fixtures.args_fixtures import replace_config_values
from tests.fast.fixtures.inkling_provider_fixtures import (
inkling_adapter_model,
inkling_provider_env,
inkling_reload_env,
)
from miles.utils.args.custom_function import CustomFunctionConfig
from miles_plugins.models.inkling.lora import InklingLoRAAdapter
_ = inkling_adapter_model, inkling_provider_env, inkling_reload_env
class TestNativeInklingSetup:
def test_typed_provider_injects_trainable_adapters_before_model_wrapping(
self, inkling_provider_env: SimpleNamespace
) -> None:
"""A typed Inkling provider still selects real native LoRA injection."""
env = inkling_provider_env
[model], optimizer, scheduler = env.module.setup_model_and_optimizer(env.args, role="actor")
assert isinstance(model.lora_lm_head_adapter, InklingLoRAAdapter)
assert model.lora_lm_head_adapter.head_A.requires_grad
assert model.lora_lm_head_adapter.head_B.requires_grad
assert not model.output_layer.weight.requires_grad
assert optimizer is None and scheduler is None
@pytest.mark.parametrize("provider", [None, CustomFunctionConfig(path="models.other.provider")])
def test_other_providers_still_reject_unsupported_native_lora(
self, inkling_provider_env: SimpleNamespace, provider: CustomFunctionConfig | None
) -> None:
"""The typed-provider fix does not accept unsupported native LoRA models."""
env = inkling_provider_env
args = replace_config_values(env.args, custom_model_provider_path=provider)
with pytest.raises(AssertionError, match="Native LoRA injection is only implemented"):
env.module.setup_model_and_optimizer(args, role="actor")
@pytest.mark.parametrize("inkling", [True, False])
def test_only_inkling_disables_muon_qkv_splitting(
self, inkling_provider_env: SimpleNamespace, inkling: bool
) -> None:
"""Typed Inkling providers retain the fused-qkvr Muon safeguard."""
env = inkling_provider_env
provider = env.args.custom_model_provider_path
if not inkling:
provider = CustomFunctionConfig(path="models.other.provider")
args = replace_config_values(
env.args, custom_model_provider_path=provider, lora_rank=0, debug_disable_optimizer=False
)
_, optimizer, _ = env.module.setup_model_and_optimizer(args, role="actor")
assert optimizer.config.muon_split_qkv is (not inkling)
class TestNativeInklingReload:
@pytest.mark.parametrize("native_optimizer_restored", [False, True])
def test_disk_adapter_reloads_only_when_native_optimizer_was_not_restored(
self, inkling_reload_env: SimpleNamespace, tmp_path: Path, native_optimizer_restored: bool
) -> None:
"""Typed Inkling reloads adapter tensors and masters without overwriting native optimizer restores."""
env = inkling_reload_env
env.native_optimizer_restored = native_optimizer_restored
adapter = env.model.lora_lm_head_adapter
initial = adapter.head_B.detach().clone()
save_file(
{
"language_model.lm_head.lora_A.weight": torch.full_like(adapter.head_A, 3),
"language_model.lm_head.lora_B.weight": torch.full_like(adapter.head_B, 7),
},
str(tmp_path / "adapter_model.safetensors"),
)
args = replace_config_values(env.args, lora_adapter_path=str(tmp_path))
env.module.load_model_state(
args,
model=[env.model],
optimizer=env.optimizer,
opt_param_scheduler=None,
role="actor",
checkpointing_context=None,
)
if native_optimizer_restored:
torch.testing.assert_close(adapter.head_B, initial)
assert env.optimizer.masters == {}
else:
torch.testing.assert_close(adapter.head_B, torch.full_like(adapter.head_B, 7))
torch.testing.assert_close(env.optimizer.masters["lora_lm_head_adapter.head_B"], adapter.head_B)
@@ -10,6 +10,7 @@ from unittest.mock import Mock, call
import pytest
import ray
import torch
from tests.fast.fixtures.args_fixtures import make_trainer_args
from tests.fast.train_parallel_config_utils import make_train_parallel_config
from miles.backends.megatron_utils.ft.types import TrainStepOutcome, TrainStepOutput
@@ -76,7 +77,10 @@ def actor_module():
def _worker(actor_module, role, *, asleep=True):
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(offload_train=True, debug_rollout_only=False, enable_sample_ownership_checker=False)
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
offload_train=True, debug_rollout_only=False, enable_sample_ownership_checker=False
)
worker.role = role
worker._asleep = asleep
worker._heartbeat = Mock()
@@ -199,7 +203,7 @@ class TestTrainParallelConfigWiring:
"""Both MultiLoRA entry points send the post-loss cell topology to data loading, not a cached one."""
lora_actor_module = importlib.import_module("miles.backends.megatron_utils.lora.actor")
worker = object.__new__(lora_actor_module.MultiLoRATrainRayActor)
worker.args = Namespace()
worker.args = make_trainer_args()
worker.model = object()
worker._heartbeat = Mock()
trivial = GroupInfo(rank=0, size=1, group=None)
@@ -249,7 +253,8 @@ def test_compute_log_prob_replays_sampling_support_only_for_actor_scores(
):
"""Policy log-prob forwards preserve the model's logits precision."""
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(use_sampling_support_replay=True)
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(use_sampling_support_replay=True)
worker.model = [object()]
forward_only = Mock(return_value={"log_probs": []})
monkeypatch.setattr(actor_module, "forward_only", forward_only)
@@ -264,7 +269,8 @@ def test_compute_log_prob_replays_sampling_support_only_for_actor_scores(
def test_save_model_does_not_manage_lifecycle(actor_module, monkeypatch):
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
async_save=False,
custom_megatron_post_save_hook_path=None,
debug_rollout_only=False,
@@ -288,7 +294,12 @@ def test_save_model_does_not_manage_lifecycle(actor_module, monkeypatch):
worker.save_model(6)
save.assert_called_once_with(
6, worker.model, worker.optimizer, worker.opt_param_scheduler, snapshot_publisher=worker.snapshot_publisher
worker.args,
6,
worker.model,
worker.optimizer,
worker.opt_param_scheduler,
snapshot_publisher=worker.snapshot_publisher,
)
worker.wake_up.assert_not_called()
worker.sleep.assert_not_called()
@@ -298,7 +309,8 @@ def test_save_model_does_not_manage_lifecycle(actor_module, monkeypatch):
def test_force_sync_save_overlaps_hf_export_with_async_checkpoint(actor_module, monkeypatch):
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
async_save=True,
custom_megatron_post_save_hook_path=None,
debug_rollout_only=False,
@@ -332,7 +344,9 @@ def test_update_weights_only_uses_temporary_process_groups_when_asleep(actor_mod
from miles.ray.rollout.inference_controller import UpdatableEngines
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
debug_rollout_only=False,
debug_skip_weight_update=True,
debug_train_only=False,
@@ -365,7 +379,8 @@ def test_update_weights_only_uses_temporary_process_groups_when_asleep(actor_mod
def _lifecycle_worker(actor_module, monkeypatch, asleep):
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
offload_train=True,
rematerialize_param_from_master_weight=False,
clear_quantized_weight_workspaces_on_offload=False,
@@ -452,11 +467,12 @@ def _actor_train_args(**overrides):
skip_actor_forward_only=False,
enable_sample_ownership_checker=False,
)
return Namespace(**(defaults | overrides))
return make_trainer_args(**(defaults | overrides))
def _actor_reuse_worker(actor_module, **args_overrides):
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker._config_snapshot_train_recorded = True
worker.args = _actor_train_args(use_critic=False, **args_overrides)
worker.model = [object()]
worker.optimizer = object()
@@ -527,7 +543,7 @@ def test_actor_logprob_forward_is_explicit_single_step_opt_in(
assert worker._compute_log_prob.call_count == int(not skip_actor_forward_only and not use_rollout_logprobs)
actor_module.compute_advantages_and_returns.assert_called_once_with(worker.args, rollout_data)
train_call = actor_module.train.call_args
assert train_call.args[6] is rollout_data["num_rollouts"]
assert train_call.args[7] is rollout_data["num_rollouts"]
assert train_call.kwargs == {
"witness_info": None,
"attempt": 0,
@@ -680,7 +696,8 @@ def _patch_shared_train_helpers(actor_module: Any, monkeypatch: pytest.MonkeyPat
def _critic_worker(actor_module: Any) -> Any:
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(global_batch_size=1, loss_type=None)
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(global_batch_size=1, loss_type="value_loss")
worker.role = "critic"
worker.model = object()
worker.optimizer = object()
@@ -691,7 +708,8 @@ def _critic_worker(actor_module: Any) -> Any:
def _actor_worker(actor_module: Any) -> Any:
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
colocate=True,
compute_advantages_and_returns=True,
get_mismatch_metrics=False,
@@ -830,7 +848,8 @@ class _RecordingWeightUpdater:
def _weight_update_worker(actor_module: Any, monkeypatch: pytest.MonkeyPatch) -> Any:
worker = object.__new__(actor_module.MegatronTrainRayActor)
worker.args = Namespace(
worker._config_snapshot_train_recorded = True
worker.args = make_trainer_args(
ci_test=False,
debug_rollout_only=False,
debug_skip_weight_update=False,
@@ -4,10 +4,11 @@ from tests.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=60, suite="stage-a-cpu", labels=[])
from argparse import Namespace
from types import SimpleNamespace
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args
from tests.fast.fixtures.sglang_config_fixtures import with_parser_defaults_and_sglang_config
from miles.backends.megatron_utils.update_weight.hf_weight_iterator import MegatronHfWeightIteratorBase
from miles.backends.training_utils.weight_update.hf_weight_iterator import WeightUpdatePlacement
@@ -23,7 +24,12 @@ class _Iterator(MegatronHfWeightIteratorBase):
def _make(speculative, mtp_num_layers):
model = [SimpleNamespace(config=SimpleNamespace(mtp_num_layers=mtp_num_layers))]
args = Namespace(sglang_speculative_algorithm=speculative, q_lora_rank=None)
args = make_trainer_args(
**with_parser_defaults_and_sglang_config(
dict(sglang_speculative_algorithm=speculative, fp16=False, rollout_num_gpus=1)
),
q_lora_rank=None
)
return _Iterator(
args, model, placement=WeightUpdatePlacement(gather_pp=True), model_name="qwen", quantization_config=None
)
@@ -0,0 +1,46 @@
from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.inkling_provider_fixtures import inkling_tower_env
from miles.utils.args.custom_function import CustomFunctionConfig
_ = inkling_tower_env
class TestInklingMultimodalTowers:
@pytest.mark.parametrize(
("provider", "materialize", "expected_names"),
[
(
CustomFunctionConfig(path="miles_plugins.models.inkling.model.inkling_mm_model_provider"),
True,
["audio.weight", "visual.weight"],
),
(
CustomFunctionConfig(path="miles_plugins.models.inkling.model.inkling_mm_model_provider"),
False,
[],
),
(CustomFunctionConfig(path="miles_plugins.models.inkling.model.inkling_model_provider"), True, []),
(None, True, []),
],
ids=["multimodal", "non-materializing-rank", "text", "no-provider"],
)
def test_only_materializing_multimodal_provider_emits_frozen_tower_weights(
self,
inkling_tower_env: SimpleNamespace,
provider: CustomFunctionConfig | None,
materialize: bool,
expected_names: list[str],
) -> None:
"""A typed multimodal provider sends frozen towers while text and non-materializing paths remain empty."""
env = inkling_tower_env
args = SimpleNamespace(custom_model_provider_path=provider, hf_checkpoint=str(env.checkpoint))
units = list(env.module._iter_mm_tower_units(args, materialize=materialize))
assert [unit[0][0] for unit in units] == expected_names
for [(name, tensor)] in units:
torch.testing.assert_close(tensor, env.tensors[name])
@@ -0,0 +1,45 @@
from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.args_fixtures import replace_config_values
from tests.fast.fixtures.inkling_provider_fixtures import inkling_adapter_model, inkling_provider_env
from miles.utils.args.custom_function import CustomFunctionConfig
_ = inkling_adapter_model, inkling_provider_env
class TestNativeInklingExport:
def test_typed_provider_exports_actual_adapter_tensors(
self, inkling_provider_env: SimpleNamespace, inkling_adapter_model: torch.nn.Module
) -> None:
"""Typed Inkling providers export HF-named tensors through the native adapter exporter."""
from miles.backends.megatron_utils.update_weight.hf_weight_iterator_direct import HfWeightIteratorDirect
iterator = object.__new__(HfWeightIteratorDirect)
iterator.args = inkling_provider_env.args
iterator.model = [inkling_adapter_model]
iterator.model_name = "inkling"
adapter = inkling_adapter_model.lora_lm_head_adapter
named = dict(iterator._export_pp_local_lora(adapter=None))
assert set(named) == {"language_model.lm_head.lora_A.weight", "language_model.lm_head.lora_B.weight"}
torch.testing.assert_close(named["language_model.lm_head.lora_A.weight"], adapter.head_A.to(torch.bfloat16))
torch.testing.assert_close(named["language_model.lm_head.lora_B.weight"], adapter.head_B.to(torch.bfloat16))
@pytest.mark.parametrize("provider", [None, CustomFunctionConfig(path="models.other.provider")])
def test_unsupported_providers_still_reject_raw_adapter_export(
self, inkling_provider_env: SimpleNamespace, provider: CustomFunctionConfig | None
) -> None:
"""Absent and unrelated typed providers cannot select the Inkling raw exporter."""
from miles.backends.megatron_utils.update_weight.hf_weight_iterator_direct import HfWeightIteratorDirect
iterator = object.__new__(HfWeightIteratorDirect)
iterator.args = replace_config_values(inkling_provider_env.args, custom_model_provider_path=provider)
iterator.model = []
iterator.model_name = "other"
with pytest.raises(NotImplementedError, match="Raw LoRA export is not implemented"):
iterator._export_pp_local_lora(adapter=None)
+3 -1
View File
@@ -7,6 +7,8 @@ from argparse import Namespace
from pathlib import Path
from typing import Any
from tests.fast.fixtures.sglang_config_fixtures import with_parser_defaults_and_sglang_config
_TINY_MODEL_CONFIG: dict[str, Any] = {
"architectures": ["LlamaForCausalLM"],
"model_type": "llama",
@@ -58,4 +60,4 @@ def make_engine_args(**overrides: Any) -> Namespace:
colocate=False,
)
defaults.update(overrides)
return Namespace(**defaults)
return Namespace(**with_parser_defaults_and_sglang_config(defaults))
@@ -1,6 +1,8 @@
from __future__ import annotations
import argparse
import sys
from unittest.mock import patch
import pytest
@@ -10,6 +12,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.arguments import parse_args
from miles.utils.workers.argv_utils import _record_field_names
@@ -25,7 +28,9 @@ def _parse_sglang_args(argv: list[str]) -> argparse.Namespace:
class TestSglangModelRoutersDefault:
def test_parsing_without_model_routers_sets_none(self):
"""Parsing without multi-policy routers exposes a safe None default."""
args = _parse_sglang_args([])
argv = ["pytest", "--train-backend", "fsdp", "--rollout-batch-size", "1", "--num-rollout", "1"]
with patch.object(sys, "argv", argv):
args = parse_args()
assert args.sglang_model_routers is None
@@ -5,6 +5,7 @@ from types import SimpleNamespace
import msgspec
import pytest
from tests.fast.fixtures.sglang_config_fixtures import with_parser_defaults_and_sglang_config
from miles.backends.sglang_utils import sglang_engine
from miles.backends.sglang_utils.sglang_api_client import WorkerType
@@ -28,11 +29,12 @@ def make_args(**overrides: object) -> SimpleNamespace:
lora_adapter_path=None,
debug_rollout_only=False,
debug_skip_weight_update=False,
multi_lora=False,
multi_lora_n_adapters=1,
lora_adapter_targets=[f"model.layers.*.self_attn.{projection}_proj" for projection in ("q", "k", "v")],
)
defaults.update(overrides)
return SimpleNamespace(**defaults)
return SimpleNamespace(**with_parser_defaults_and_sglang_config(defaults))
def compute(args: SimpleNamespace, **overrides: object) -> dict:
@@ -11,6 +11,7 @@ from miles.backends.sglang_utils.router_args_utils import (
parse_router_args_argv,
router_args_to_argv,
)
from miles.utils.args.configs.router import RouterConfig
def _make_router_cli_parser() -> argparse.ArgumentParser:
@@ -35,7 +36,7 @@ def _make_miles_args(**overrides: object) -> Namespace:
values = vars(_make_prefixed_router_cli_parser().parse_args([]))
values.update(sglang_router_request_timeout_secs=600, sglang_router_policy=None)
values.update(overrides)
return Namespace(**values)
return Namespace(**values, **RouterConfig.from_args(Namespace(**values)))
class TestRouterArgsToArgv:
@@ -5,14 +5,17 @@ register_cpu_ci(est_time=20, suite="stage-a-cpu", labels=[])
import argparse
from miles.backends.sglang_utils.arguments import add_sglang_arguments, validate_args
from miles.utils.args.configs.router import RouterConfig
from miles.utils.http_utils import router_worker_base_urls
def _args(argv):
parser = add_sglang_arguments(argparse.ArgumentParser())
RouterConfig.add_arguments(parser)
args = parser.parse_args(argv)
args.rollout_num_gpus_per_engine = 4
args.true_on_policy_mode = False
args.recompute_logprobs_via_prefill = False
args.use_session_server = False
# Registered by RouterArgs.add_cli_args, not by add_sglang_arguments.
args.router_assignment_mode = "random"
@@ -3,13 +3,16 @@ from __future__ import annotations
from argparse import Namespace
import pytest
from tests.fast.fixtures.sglang_config_fixtures import resolve_sglang_config, resolve_sglang_config_and_scaling
from miles.backends.sglang_utils.sglang_api_client import WorkerType
from miles.backends.sglang_utils.sglang_config import (
ModelConfig,
ServerGroupConfig,
ServerGroupScalingConfig,
SglangConfig,
_compute_megatron_num_gpus,
_compute_rollout_offset,
resolve_sglang_config,
)
from miles.utils.external_utils.command_utils.common import encode_pseudo_file
@@ -33,7 +36,7 @@ def _make_args(**overrides) -> Namespace:
critic_num_nodes=0,
critic_num_gpus_per_node=0,
use_critic=False,
critic_train_only=False,
multi_lora=False,
)
defaults.update(overrides)
defaults.setdefault("starts_inference_engines", not defaults["debug_train_only"] or defaults["eval_num_gpus"] > 0)
@@ -43,9 +46,14 @@ def _make_args(**overrides) -> Namespace:
def _resolve_yaml(tmp_path, yaml_text: str, **args_overrides):
config, _ = _resolve_yaml_and_scaling(tmp_path, yaml_text, **args_overrides)
return config
def _resolve_yaml_and_scaling(tmp_path, yaml_text: str, **args_overrides):
cfg_path = tmp_path / "config.yaml"
cfg_path.write_text(yaml_text)
return resolve_sglang_config(_make_args(sglang_config=str(cfg_path), **args_overrides))
return resolve_sglang_config_and_scaling(_make_args(sglang_config=str(cfg_path), **args_overrides))
class TestNumGpusPerEnginePrecedence:
@@ -140,32 +148,18 @@ class TestResolvedServerGroupValidation:
def test_a_resolved_group_with_zero_gpus_is_rejected(self):
"""A group reserving no GPUs cannot host an engine and must fail construction."""
with pytest.raises(ValueError, match="greater than 0"):
ServerGroupConfig(
worker_type="regular",
num_gpus=0,
num_gpus_per_engine=2,
gpu_offset=0,
engine_offset=0,
needs_offload=False,
)
ServerGroupScalingConfig(num_gpus=0, gpu_offset=0, engine_offset=0)
def test_a_resolved_group_with_non_positive_gpus_per_engine_is_rejected(self):
"""A non-positive engine width would make the engine count division meaningless."""
with pytest.raises(ValueError, match="greater than 0"):
ServerGroupConfig(
worker_type="regular",
num_gpus=8,
num_gpus_per_engine=-1,
gpu_offset=0,
engine_offset=0,
needs_offload=False,
)
ServerGroupConfig(worker_type="regular", num_gpus_per_engine=-1, needs_offload=False)
class TestNumServerCells:
def test_non_placeholder_engine_cells_are_counted(self, tmp_path):
"""The server cell count includes every engine except placeholder reservations."""
cfg = _resolve_yaml(
cfg, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n"
" - name: actor\n"
@@ -182,7 +176,7 @@ class TestNumServerCells:
rollout_num_gpus=16,
)
assert cfg.models[0].num_server_cells == 4
assert cfg.models[0].num_server_cells(scaling) == 4
class TestOverridesResolution:
@@ -222,12 +216,12 @@ class TestOverridesResolution:
class TestYamlShapeValidation:
def test_engine_groups_is_accepted_as_an_alias_for_server_groups(self, tmp_path):
"""The documented engine_groups spelling keeps parsing."""
cfg = _resolve_yaml(
_, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n - name: actor\n engine_groups:\n - worker_type: regular\n num_gpus: 8\n",
rollout_num_gpus=8,
)
assert cfg.models[0].server_groups[0].num_gpus == 8
assert scaling.group(model_name="actor", group_index=0).num_gpus == 8
def test_a_yaml_without_the_sglang_key_is_rejected(self, tmp_path):
"""A config missing the top-level sglang key fails loudly."""
@@ -283,13 +277,13 @@ class TestPrefillNumServersPath:
@pytest.mark.parametrize("multi_lora", [False, True])
def test_prefill_num_servers_counts_engines_not_gpus(self, multi_lora):
"""prefill_num_servers is a server count, so its GPU span scales with the engine width."""
cfg = resolve_sglang_config(
cfg, scaling = resolve_sglang_config_and_scaling(
_make_args(
rollout_num_gpus=16, prefill_num_servers=3, rollout_num_gpus_per_engine=2, multi_lora=multi_lora
)
)
groups = cfg.models[0].server_groups
assert [(group.worker_type, group.num_gpus) for group in groups] == [
groups = zip(cfg.models[0].server_groups, scaling.groups["default"], strict=True)
assert [(group.worker_type, group_scaling.num_gpus) for group, group_scaling in groups] == [
(WorkerType.PREFILL, 6),
(WorkerType.DECODE, 10),
]
@@ -305,7 +299,7 @@ class TestPrefillNumServersPath:
class TestGpuOffset:
def test_gpu_offsets_accumulate_across_groups_and_models_including_placeholders(self, tmp_path):
"""Each group's gpu_offset equals the num_gpus sum of all preceding groups, counting placeholders."""
cfg = _resolve_yaml(
_, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n"
" - name: actor\n"
@@ -321,14 +315,14 @@ class TestGpuOffset:
" num_gpus: 8\n",
rollout_num_gpus=16,
)
assert [group.gpu_offset for group in cfg.models[0].server_groups] == [0, 4]
assert cfg.models[1].server_groups[0].gpu_offset == 8
assert [group.gpu_offset for group in scaling.groups["actor"]] == [0, 4]
assert scaling.groups["ref"][0].gpu_offset == 8
class TestEngineOffset:
def test_engine_offsets_count_the_workers_of_every_preceding_group(self, tmp_path):
"""Groups of different engine widths contribute different worker counts, so a gpu offset alone cannot number them."""
cfg = _resolve_yaml(
_, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n"
" - name: actor\n"
@@ -348,12 +342,12 @@ class TestEngineOffset:
rollout_num_gpus=28,
num_gpus_per_node=4,
)
assert [group.engine_offset for group in cfg.models[0].server_groups] == [0, 4]
assert cfg.models[1].server_groups[0].engine_offset == 5
assert [group.engine_offset for group in scaling.groups["actor"]] == [0, 4]
assert scaling.groups["ref"][0].engine_offset == 5
def test_a_group_whose_engine_spans_nodes_contributes_one_worker_per_node(self, tmp_path):
"""A cross-node engine is launched by one actor per node, so it consumes that many numbers, not one."""
cfg = _resolve_yaml(
_, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n"
" - name: actor\n"
@@ -367,19 +361,19 @@ class TestEngineOffset:
rollout_num_gpus=20,
num_gpus_per_node=4,
)
assert [group.engine_offset for group in cfg.models[0].server_groups] == [0, 4]
assert [group.engine_offset for group in scaling.groups["actor"]] == [0, 4]
def test_prefill_and_decode_groups_are_numbered_in_that_order(self):
"""The legacy --prefill-num-servers layout has no YAML to carry offsets, so the cursor must number it too."""
cfg = resolve_sglang_config(
_, scaling = resolve_sglang_config_and_scaling(
_make_args(rollout_num_gpus=16, prefill_num_servers=3, rollout_num_gpus_per_engine=2)
)
assert [group.engine_offset for group in cfg.models[0].server_groups] == [0, 3]
assert [group.engine_offset for group in scaling.groups["default"]] == [0, 3]
def test_the_generated_eval_model_is_numbered_after_every_rollout_engine(self):
"""The eval fleet is appended without YAML, and reusing the rollout numbers would clone their RNG streams."""
cfg = resolve_sglang_config(
_, scaling = resolve_sglang_config_and_scaling(
_make_args(
rollout_num_gpus=8,
rollout_num_gpus_per_engine=2,
@@ -388,8 +382,8 @@ class TestEngineOffset:
)
)
assert cfg.models[0].server_groups[0].engine_offset == 0
assert cfg.models[1].server_groups[0].engine_offset == 4
assert scaling.groups["default"][0].engine_offset == 0
assert scaling.groups["eval"][0].engine_offset == 4
class TestNeedsOffload:
@@ -583,17 +577,11 @@ class TestRolloutOffset:
class TestMegatronNumGpus:
def test_compute_megatron_num_gpus_for_critic_train_only(self):
"""With only the critic training, the megatron span is the critic's own gpus, not the actor's."""
args = _make_args(
critic_train_only=True,
debug_rollout_only=False,
actor_num_nodes=1,
actor_num_gpus_per_node=8,
critic_num_nodes=1,
critic_num_gpus_per_node=4,
)
assert _compute_megatron_num_gpus(args) == 4
def test_a_namespace_carrying_the_removed_critic_train_only_field_is_rejected(self):
"""critic_train_only no longer resizes the megatron span, so a namespace still carrying it fails loudly."""
args = _make_args(critic_train_only=True, critic_num_nodes=1, critic_num_gpus_per_node=4)
with pytest.raises(AssertionError, match="critic_train_only is not supported"):
_compute_megatron_num_gpus(args)
class TestHostPortOverrideRejection:
@@ -716,3 +704,23 @@ class TestSglangConfigFileArg:
cfg = resolve_sglang_config(_make_args(sglang_config=encode_pseudo_file(self._YAML)))
assert [model.name for model in cfg.models] == ["actor"]
class TestCommonValue:
def test_groups_that_agree_share_their_override(self):
"""Every engine group overriding a field to the same value yields that value."""
group = ServerGroupConfig(
worker_type=WorkerType.REGULAR, num_gpus_per_engine=1, needs_offload=False, overrides={"x": 2}
)
model = ModelConfig(name="actor", model_path=None, server_groups=[group, group], update_weights=True)
assert SglangConfig(models=[model], base_args={"x": 1}).common_value("x") == 2
def test_a_config_without_engine_groups_reads_the_cli_value(self):
"""External engines leave no local engine group, so the CLI value is the common value."""
assert SglangConfig(models=[], base_args={"x": 1}).common_value("x") == 1
def test_a_field_nobody_declares_is_rejected(self):
"""Reading an unknown field fails instead of returning a default."""
with pytest.raises(AttributeError, match="No common field 'y'"):
SglangConfig(models=[], base_args={"x": 1}).common_value("y")
+1 -1
View File
@@ -23,7 +23,7 @@ class _RecordingModelCls:
def _actor(attn_implementation, model_cls):
actor = object.__new__(FSDPTrainRayActor)
actor.args = SimpleNamespace(attn_implementation=attn_implementation)
actor.args = SimpleNamespace(backend=SimpleNamespace(attn_implementation=attn_implementation))
actor._get_model_cls = lambda: model_cls
return actor
+3 -3
View File
@@ -286,7 +286,7 @@ def test_model_type_verified_accepts_recorded_models(model_type, caplog):
from miles.backends.fsdp_utils.adaptations.class_patches import check_model_type_verified
check_model_type_verified(SimpleNamespace(model_type=model_type), SimpleNamespace(rank=0))
check_model_type_verified(SimpleNamespace(model_type=model_type), SimpleNamespace(backend=SimpleNamespace(rank=0)))
assert not caplog.records
@@ -303,7 +303,7 @@ def test_model_type_verified_warns_once_on_rank_zero(caplog):
from miles.backends.fsdp_utils.adaptations.class_patches import _MODEL_PATCH_HOOKS
hook = next(h for h in _MODEL_PATCH_HOOKS if h.name == "model_type_verified")
hook.apply(SimpleNamespace(model_type="qwen3_5_moe"), SimpleNamespace(rank=0))
hook.apply(SimpleNamespace(model_type="qwen3_5_moe"), SimpleNamespace(backend=SimpleNamespace(rank=0)))
assert len(caplog.records) == 1
assert "model_type='qwen3_5_moe' has no recorded FSDP validation" in caplog.text
@@ -315,7 +315,7 @@ def test_model_type_verified_is_silent_on_nonzero_rank(caplog):
from miles.backends.fsdp_utils.adaptations.class_patches import _MODEL_PATCH_HOOKS
hook = next(h for h in _MODEL_PATCH_HOOKS if h.name == "model_type_verified")
hook.apply(SimpleNamespace(model_type="qwen3_5_moe"), SimpleNamespace(rank=1))
hook.apply(SimpleNamespace(model_type="qwen3_5_moe"), SimpleNamespace(backend=SimpleNamespace(rank=1)))
assert not caplog.records
+26 -29
View File
@@ -6,6 +6,7 @@ from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.sglang_config_fixtures import make_sglang_config
from miles.backends.fsdp_utils.adaptations.precision import (
apply_fp32_master,
@@ -17,9 +18,19 @@ from miles.backends.training_utils.data import _rollout_logprob_dtype
from miles.true_on_policy.contracts import QWEN3_DENSE_TRUE_ON_POLICY_V1
def _args(
*, fp16: bool, keep_fp32_master: bool, true_on_policy_mode: bool = False, contract: str | None = None
) -> SimpleNamespace:
return SimpleNamespace(
backend=SimpleNamespace(fp16=fp16, keep_fp32_master=keep_fp32_master),
true_on_policy_mode=true_on_policy_mode,
sglang=make_sglang_config(true_on_policy_contract=contract),
)
def test_resolve_precision_policy_uses_independent_fp32_master_switch_and_dtypes():
dense = SimpleNamespace(model_type="qwen3")
bf16_args = SimpleNamespace(fp16=False, keep_fp32_master=True)
bf16_args = _args(fp16=False, keep_fp32_master=True)
p = resolve_precision_policy(dense, bf16_args)
assert p.keep_fp32_master
@@ -27,13 +38,13 @@ def test_resolve_precision_policy_uses_independent_fp32_master_switch_and_dtypes
disabled = resolve_precision_policy(
SimpleNamespace(model_type="glm4_moe_lite"),
SimpleNamespace(fp16=True, keep_fp32_master=False),
_args(fp16=True, keep_fp32_master=False),
)
assert not disabled.keep_fp32_master
assert disabled.param_dtype == torch.float16 and disabled.reduce_dtype == torch.float32
assert disabled == resolve_precision_policy(
dense,
SimpleNamespace(fp16=True, keep_fp32_master=False),
_args(fp16=True, keep_fp32_master=False),
)
@@ -48,20 +59,17 @@ def test_fp32_master_cli_defaults_enabled_and_can_be_disabled(monkeypatch):
def test_fsdp_args_expose_effective_compute_precision(monkeypatch):
for cli_args, expected_dtype in (([], torch.bfloat16), (["--fp16"], torch.float16)):
monkeypatch.setattr(sys, "argv", ["miles", *cli_args])
args = load_fsdp_args()
args.true_on_policy_mode = True
backend = load_fsdp_args()
args = SimpleNamespace(backend=backend, true_on_policy_mode=True)
assert args.bf16 == (not args.fp16)
assert backend.bf16 == (not backend.fp16)
assert resolve_precision_policy(None, args).param_dtype is expected_dtype
assert _rollout_logprob_dtype(args) is expected_dtype
def test_qwen3_formal_true_on_policy_resolves_fp32_params_with_bf16_autocast():
args = SimpleNamespace(
fp16=False,
keep_fp32_master=True,
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
args = _args(
fp16=False, keep_fp32_master=True, true_on_policy_mode=True, contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name
)
policy = resolve_precision_policy(SimpleNamespace(model_type="qwen3"), args)
@@ -84,12 +92,7 @@ def test_qwen3_formal_true_on_policy_resolves_fp32_params_with_bf16_autocast():
def test_qwen3_formal_precision_does_not_leak_to_other_modes(model_type, true_on_policy_mode, contract):
policy = resolve_precision_policy(
SimpleNamespace(model_type=model_type),
SimpleNamespace(
fp16=False,
keep_fp32_master=True,
true_on_policy_mode=true_on_policy_mode,
sglang_true_on_policy_contract=contract,
),
_args(fp16=False, keep_fp32_master=True, true_on_policy_mode=true_on_policy_mode, contract=contract),
)
assert policy.param_dtype is torch.bfloat16
@@ -101,11 +104,8 @@ def test_qwen3_formal_true_on_policy_rejects_fp16():
with pytest.raises(ValueError, match="requires bf16 training"):
resolve_precision_policy(
SimpleNamespace(model_type="qwen3"),
SimpleNamespace(
fp16=True,
keep_fp32_master=True,
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
_args(
fp16=True, keep_fp32_master=True, true_on_policy_mode=True, contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name
),
)
@@ -114,11 +114,11 @@ def test_qwen3_formal_true_on_policy_rejects_disabled_fp32_master():
with pytest.raises(ValueError, match="requires fp32 master weights"):
resolve_precision_policy(
SimpleNamespace(model_type="qwen3"),
SimpleNamespace(
_args(
fp16=False,
keep_fp32_master=False,
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
),
)
@@ -134,11 +134,8 @@ def test_precision_forward_context_uses_policy_autocast(monkeypatch):
monkeypatch.setattr(torch, "autocast", fake_autocast)
policy = resolve_precision_policy(
SimpleNamespace(model_type="qwen3"),
SimpleNamespace(
fp16=False,
keep_fp32_master=True,
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
_args(
fp16=False, keep_fp32_master=True, true_on_policy_mode=True, contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name
),
)
@@ -3,6 +3,7 @@ from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.sglang_config_fixtures import make_sglang_config
from transformers.models.qwen3 import modeling_qwen3
from transformers.models.qwen3.configuration_qwen3 import Qwen3Config
@@ -75,17 +76,17 @@ def test_qwen3_instance_patch_registry_is_contract_gated(monkeypatch):
model,
SimpleNamespace(model_type=model_type),
SimpleNamespace(
backend=SimpleNamespace(fp16=fp16),
true_on_policy_mode=true_on_policy_mode,
sglang_true_on_policy_contract=contract,
fp16=fp16,
sglang=make_sglang_config(true_on_policy_contract=contract),
),
)
assert calls == []
args = SimpleNamespace(
backend=SimpleNamespace(fp16=False),
true_on_policy_mode=True,
sglang_true_on_policy_contract=formal_contract,
fp16=False,
sglang=make_sglang_config(true_on_policy_contract=formal_contract),
)
config = SimpleNamespace(model_type="qwen3")
assert hook.applies_to(config, args)
@@ -141,10 +142,9 @@ def test_qwen3_formal_sync_preserves_post_update_fp32_values():
policy = resolve_precision_policy(
model.config,
SimpleNamespace(
fp16=False,
keep_fp32_master=True,
backend=SimpleNamespace(fp16=False, keep_fp32_master=True),
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
sglang=make_sglang_config(true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name),
),
)
model = apply_fp32_master(model, policy.sync_dtype_resolver)
@@ -169,11 +169,9 @@ def test_qwen3_ref_model_uses_fp32_master_storage(monkeypatch):
config = _tiny_config()
args = SimpleNamespace(
attn_implementation="eager",
fp16=False,
keep_fp32_master=True,
backend=SimpleNamespace(attn_implementation="eager", fp16=False, keep_fp32_master=True),
true_on_policy_mode=True,
sglang_true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name,
sglang=make_sglang_config(true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1.name),
)
actor = object.__new__(actor_module.FSDPTrainRayActor)
actor.args = args
@@ -173,10 +173,6 @@ def test_enable_follows_use_routing_replay():
assert routing_replay.enable(Namespace(use_routing_replay=False, ci_test=False)) is False
def test_enable_defaults_false_when_arg_absent():
assert routing_replay.enable(Namespace(ci_test=False)) is False
def test_enable_turns_on_the_replay_check_only_under_ci_test():
routing_replay.enable(Namespace(use_routing_replay=True, ci_test=True))
assert routing_replay_manager.enable_check_replay_result is True
+3 -1
View File
@@ -6,7 +6,9 @@ from miles.backends.fsdp_utils import actor as actor_module
def test_save_model_delegates_to_checkpoint(monkeypatch):
actor = object.__new__(actor_module.FSDPTrainRayActor)
actor.args = SimpleNamespace(debug_rollout_only=False, save="/tmp/checkpoint", async_save=False)
actor.args = SimpleNamespace(
debug_rollout_only=False, backend=SimpleNamespace(save="/tmp/checkpoint", async_save=False)
)
save = Mock()
monkeypatch.setattr(actor_module.checkpoint, "save", save)
+1
View File
@@ -22,6 +22,7 @@ def test_fsdp_train_debug_rollout_only_returns_a_normal_output(monkeypatch):
"""A debug-rollout-only FSDP step trains nothing yet answers the driver with a NORMAL output."""
actor = object.__new__(actor_module.FSDPTrainRayActor)
actor.args = Namespace(offload_train=False, debug_rollout_only=True)
actor._config_snapshot_train_recorded = True
actor.train_parallel_config = make_train_parallel_config(dp_size=1)
actor._heartbeat = Mock()
actor._train_core = Mock()
@@ -29,6 +29,8 @@ from pathlib import Path
import torch
import torch.distributed as dist
from tests.fast.fixtures.args_fixtures import ConfigNamespace
from miles.backends.training_utils.parallel import GroupInfo, ParallelState, set_parallel_state
ARTIFACTS_REPO = "https://github.com/yueming-yuan/miles-artifacts.git"
@@ -104,27 +106,41 @@ _ARGS_DEFAULTS = dict(
custom_pg_loss_reducer_function_path=None,
use_opsm=False,
opsm_delta=0.1,
calculate_per_token_loss=False,
eps_clip_c=None,
dump_details=None,
multi_lora=False,
custom_loss_function_path=None,
# value_loss_function
value_clip=0.2,
# loss_function dispatcher
global_batch_size=1, # overridden by make_inputs
use_dynamic_global_batch_size=False,
recompute_loss_function=False,
)
# fields a trainer reads from its backend namespace, kept flat in the stored snapshot args
_BACKEND_ARGS_DEFAULTS = dict(
bf16=False,
fp16=False,
calculate_per_token_loss=False,
vocab_size=None,
global_batch_size=1, # overridden by make_inputs
)
def make_args(**overrides) -> Namespace:
d = {**_ARGS_DEFAULTS, **overrides}
return Namespace(**d)
return args_from_dict(overrides)
def args_to_dict(args: Namespace) -> dict:
return vars(args)
flat = {name: value for name, value in vars(args).items() if name not in {"backend", "train_backend"}}
return flat | vars(args.backend)
def args_from_dict(d: dict) -> Namespace:
return Namespace(**{**_ARGS_DEFAULTS, **d})
values = {**_ARGS_DEFAULTS, **_BACKEND_ARGS_DEFAULTS, **d}
backend = {name: values.pop(name) for name in _BACKEND_ARGS_DEFAULTS}
return ConfigNamespace(**values, train_backend="megatron", backend=Namespace(**backend))
# ---------------------------------------------------------------------------
@@ -154,7 +154,7 @@ def _get_sum_of_sample_mean(batch, args, parallel_state):
batch["total_lengths"],
batch["response_lengths"],
batch["loss_masks"],
args.calculate_per_token_loss,
args.backend.calculate_per_token_loss,
args.qkv_format,
batch.get("max_seq_lens", None),
)
@@ -17,12 +17,12 @@ def test_loss_passes_return_independent_detached_outputs(monkeypatch, recompute)
monkeypatch.setattr(loss_module, "get_sum_of_sample_mean", lambda *_args, **_kwargs: None)
monkeypatch.setattr(tinker_losses, "_target_logprobs", lambda _args, _batch, logits: [logits])
args = Namespace(
calculate_per_token_loss=False,
qkv_format="thd",
loss_type="policy_loss",
recompute_loss_function=recompute,
use_dynamic_global_batch_size=True,
global_batch_size=1,
multi_lora=True,
backend=Namespace(calculate_per_token_loss=False, global_batch_size=1),
)
completed = []
for sample_index in (7, 11):
@@ -102,12 +102,12 @@ def test_nonzero_objectives_and_gradients(monkeypatch, recompute, loss_fn, confi
monkeypatch.setattr(loss_module, "get_sum_of_sample_mean", lambda *_args, **_kwargs: None)
monkeypatch.setattr(tinker_losses, "_target_logprobs", lambda _args, _batch, logits: [logits[:3], logits[3:]])
args = Namespace(
calculate_per_token_loss=False,
qkv_format="thd",
loss_type="policy_loss",
recompute_loss_function=recompute,
use_dynamic_global_batch_size=True,
global_batch_size=4,
multi_lora=True,
backend=Namespace(calculate_per_token_loss=False, global_batch_size=4),
)
# Both advantage signs cross each clipping boundary; the last two tokens are masked.
ratios = torch.tensor([0.5, 0.5, 1, 1, 5, 5, 0.5, 5], dtype=torch.float64)
@@ -12,6 +12,7 @@ from miles.backends.training_utils.loss import compute_advantages_and_returns
from miles.backends.training_utils.loss_hub import losses as losses_module
from miles.backends.training_utils.loss_hub.logit_processors import get_log_probs_and_entropy
from miles.backends.training_utils.loss_hub.losses import policy_loss_function
from miles.utils.args.custom_function import CustomFunctionConfig
from miles.utils.ft_utils.process_group_utils import GroupInfo
from .loss_test_utils import deep_clone, make_args, make_batch, make_inputs, make_parallel_state, make_rollout_data
@@ -39,7 +40,7 @@ def _run_policy_loss(args, batch, inputs, *, skip_actor_forward_only):
batch["total_lengths"],
batch["response_lengths"],
batch["loss_masks"],
args.calculate_per_token_loss,
args.backend.calculate_per_token_loss,
args.qkv_format,
batch.get("max_seq_lens"),
)
@@ -233,7 +234,7 @@ def test_policy_loss_rejects_missing_old_policy_log_probs(
batch["total_lengths"],
batch["response_lengths"],
batch["loss_masks"],
args.calculate_per_token_loss,
args.backend.calculate_per_token_loss,
args.qkv_format,
)
@@ -339,7 +340,7 @@ def test_mismatch_metrics_keep_actor_log_probs_as_training_source(process_group,
parallel_state = make_parallel_state()
parallel_state.tp = GroupInfo(rank=0, size=1, group=dist.group.WORLD)
args = make_args(
custom_tis_function_path="tests.fake_tis",
custom_tis_function_path=CustomFunctionConfig(path="tests.fake_tis"),
entropy_coef=0.0,
get_mismatch_metrics=True,
observe_training_entropy=False,
@@ -26,6 +26,7 @@ class TestGetLogProbsAndEntropy:
log_probs_chunk_size=-1,
allgather_cp=False,
debug_unified_grad_fused_logprob=True,
train_backend="fsdp",
)
with torch.no_grad():
@@ -9,8 +9,9 @@ class TestCheckKl:
def test_namespaced_policy_metrics_still_trigger_the_kl_checker(self) -> None:
"""A policy namespace must not hide an out-of-tolerance PPO KL value."""
args = Namespace(
multi_latent_attention=False,
backend=Namespace(multi_latent_attention=False),
trainer_model_id="alpha",
lora_rank=0,
use_rollout_routing_replay=False,
)
@@ -26,10 +26,9 @@ def _parallel_state(*, dp_size: int) -> ParallelState:
def _static_args() -> Namespace:
return Namespace(
qkv_format="thd",
global_batch_size=256,
use_dynamic_global_batch_size=False,
use_dynamic_batch_size=False,
micro_batch_size=8,
backend=Namespace(global_batch_size=256, micro_batch_size=8),
)
@@ -37,6 +36,7 @@ def _scheduled_args() -> Namespace:
return Namespace(
qkv_format="thd",
global_batch_size=256,
backend=Namespace(global_batch_size=256),
use_dynamic_global_batch_size=False,
use_dynamic_batch_size=True,
max_tokens_per_gpu=64,
@@ -32,10 +32,8 @@ def _args(qkv_format: str) -> Namespace:
enable_witness=False,
qkv_format=qkv_format,
data_pad_size_multiplier=16,
compress_ratios=[],
true_on_policy_mode=False,
bf16=False,
fp16=False,
backend=Namespace(compress_ratios=[], bf16=False, fp16=False),
)
@@ -154,11 +154,10 @@ def test_get_log_probs_and_entropy_applies_per_response_sampling_support(monkeyp
qkv_format="thd",
rollout_temperature=1.0,
true_on_policy_mode=True,
bf16=False,
fp16=False,
log_probs_chunk_size=-1,
vocab_size=4,
allgather_cp=False,
train_backend="megatron",
backend=SimpleNamespace(bf16=False, fp16=False, vocab_size=4),
debug_unified_grad_fused_logprob=False,
)
logits = torch.tensor(
@@ -10,13 +10,22 @@ from miles.backends.training_utils import log_utils
def test_true_on_policy_rollout_logprob_dtype_follows_training_precision():
assert (
data_utils._rollout_logprob_dtype(Namespace(true_on_policy_mode=True, bf16=True, fp16=False)) is torch.bfloat16
data_utils._rollout_logprob_dtype(
Namespace(true_on_policy_mode=True, backend=Namespace(bf16=True, fp16=False))
)
is torch.bfloat16
)
assert (
data_utils._rollout_logprob_dtype(Namespace(true_on_policy_mode=True, bf16=False, fp16=True)) is torch.float16
data_utils._rollout_logprob_dtype(
Namespace(true_on_policy_mode=True, backend=Namespace(bf16=False, fp16=True))
)
is torch.float16
)
assert (
data_utils._rollout_logprob_dtype(Namespace(true_on_policy_mode=False, bf16=True, fp16=False)) is torch.float32
data_utils._rollout_logprob_dtype(
Namespace(true_on_policy_mode=False, backend=Namespace(bf16=True, fp16=False))
)
is torch.float32
)
@@ -3,12 +3,14 @@ from types import SimpleNamespace
import pytest
import torch
from tests.fast.fixtures.args_fixtures import ConfigNamespace
from miles.backends.training_utils.loss_hub import losses as loss_utils
from miles.utils.args.custom_function import CustomFunctionConfig
def _make_args(*, use_rollout_logprobs: bool) -> Namespace:
return Namespace(
return ConfigNamespace(
use_rollout_logprobs=use_rollout_logprobs,
use_sampling_support_replay=False,
skip_actor_forward_only=False,
@@ -18,9 +20,11 @@ def _make_args(*, use_rollout_logprobs: bool) -> Namespace:
use_tis=False,
eps_clip=0.2,
eps_clip_high=0.2,
eps_clip_c=None,
dump_details=None,
custom_tis_function_path=None,
custom_pg_loss_reducer_function_path=None,
calculate_per_token_loss=False,
backend=Namespace(calculate_per_token_loss=False),
qkv_format="thd",
entropy_coef=0.0,
use_kl_loss=False,
@@ -210,7 +214,7 @@ def test_kl_loss_does_not_backpropagate_through_reference_scores(monkeypatch):
def test_custom_tis_can_ignore_missing_trainer_scored_log_probs(monkeypatch):
args = _make_args(use_rollout_logprobs=True)
args.use_tis = True
args.custom_tis_function_path = "tests.custom_tis"
args.custom_tis_function_path = CustomFunctionConfig(path="tests.custom_tis")
batch = _make_batch(
old_log_probs=torch.tensor([0.10, 0.20]),
rollout_log_probs=torch.tensor([0.12, 0.22]),
@@ -4,6 +4,7 @@ from typing import Any
import pytest
import torch
from tests.fast.fixtures.sglang_config_fixtures import make_sglang_config
def _plan(
@@ -71,7 +72,7 @@ class TestRemoteTransferPlanParallelism:
p2p_transfer_utils, monkeypatch, pp_rank=1, pp_size=2, gathered_dp_rank=3, gathered_dp_size=4
)
plan = p2p_transfer_utils.RemoteTransferPlan(SimpleNamespace(sglang_pp_size=1))
plan = p2p_transfer_utils.RemoteTransferPlan(SimpleNamespace(sglang=make_sglang_config(pp_size=1)))
assert (plan._pp_rank, plan._pp_size) == (1, 2)
assert (plan._gathered_dp_rank, plan._gathered_dp_size) == (3, 4)
@@ -85,7 +86,7 @@ class TestRemoteTransferPlanParallelism:
)
with pytest.raises(NotImplementedError):
p2p_transfer_utils.RemoteTransferPlan(SimpleNamespace(sglang_pp_size=2))
p2p_transfer_utils.RemoteTransferPlan(SimpleNamespace(sglang=make_sglang_config(pp_size=2)))
class TestPlanP2P:
@@ -91,7 +91,6 @@ class TestConnectRolloutEnginesFromDistributed:
):
ray_mock._private.services.get_node_ip_address.return_value = "10.0.0.1"
result = connect_rollout_engines_from_distributed(
Namespace(),
"miles-pp_0",
engines,
engine_gpu_counts=[2, 4, 1],
@@ -112,38 +111,17 @@ class TestConnectRolloutEnginesFromDistributed:
"group_name": "miles-pp_0",
}
def test_missing_gpu_counts_fall_back_to_the_uniform_engine_size(self) -> None:
"""Callers without discovered counts retain the uniform engine layout."""
started = threading.Semaphore(0)
release = threading.Event()
engines = [_GatedEngine(started, release) for _ in range(3)]
group = MagicMock(name="nccl_group")
def join(**kwargs):
for _ in engines:
assert started.acquire(timeout=30), "an engine had not been asked before the local join"
release.set()
return group
def test_missing_gpu_counts_are_rejected_before_any_engine_is_asked(self) -> None:
"""The inference controller always reports per-engine GPU counts, so a caller without them is a bug."""
engines = [_AcceptingEngine(), _AcceptingEngine()]
with (
patch(f"{_BROADCAST_MODULE}.ray") as ray_mock,
patch(f"{_BROADCAST_MODULE}.init_process_group", side_effect=join) as init_process_group,
patch(f"{_BROADCAST_MODULE}.init_process_group") as init_process_group,
pytest.raises(AssertionError),
):
ray_mock._private.services.get_node_ip_address.return_value = "10.0.0.1"
result = connect_rollout_engines_from_distributed(
Namespace(rollout_num_gpus_per_engine=2),
"miles-pp_0",
engines,
)
connect_rollout_engines_from_distributed("miles-pp_0", engines)
assert result is group
master_port = engines[0].calls[0][1][1]
assert [engine.calls for engine in engines] == [
[("init_weights_update_group", ("10.0.0.1", master_port, 1, 7, "miles-pp_0"), {"backend": "nccl"})],
[("init_weights_update_group", ("10.0.0.1", master_port, 3, 7, "miles-pp_0"), {"backend": "nccl"})],
[("init_weights_update_group", ("10.0.0.1", master_port, 5, 7, "miles-pp_0"), {"backend": "nccl"})],
]
assert init_process_group.call_args.kwargs["world_size"] == 7
init_process_group.assert_not_called()
def test_an_engine_that_refuses_the_group_fails_the_connect(self) -> None:
"""The submitted joins are awaited, so a refusing engine surfaces instead of being dropped."""
@@ -154,9 +132,9 @@ class TestConnectRolloutEnginesFromDistributed:
ray_mock._private.services.get_node_ip_address.return_value = "10.0.0.1"
with pytest.raises(RuntimeError, match="engine refused the group"):
connect_rollout_engines_from_distributed(
Namespace(rollout_num_gpus_per_engine=2),
"miles-pp_0",
[_AcceptingEngine(), _AcceptingEngine(), _RefusingEngine()],
engine_gpu_counts=[2, 2, 2],
)
@@ -189,7 +167,6 @@ class TestUpdateWeightFromDistributedConnect:
)
connect.assert_called_once_with(
protocol.args,
"miles-pp_0",
engines,
engine_gpu_counts=[2, 4],
@@ -205,7 +182,6 @@ class TestDisconnectRolloutEnginesFromDistributed:
dist_mock.destroy_process_group.side_effect = RuntimeError("nccl teardown failed")
with pytest.raises(RuntimeError, match="nccl teardown failed"):
disconnect_rollout_engines_from_distributed(
Namespace(),
"miles-pp_0",
MagicMock(name="nccl_group"),
engines,
@@ -90,6 +90,7 @@ class TestConnect:
patch(f"{_TENSOR_MODULE}.disconnect_rollout_engines_from_distributed") as disconnect,
):
dist_mock.get_rank.return_value = rank
dist_mock.get_world_size.return_value = 8
protocol.connect(
engines,
engine_gpu_counts=_ENGINE_GPU_COUNTS,
@@ -120,7 +121,7 @@ class TestConnect:
assert protocol.use_distribute is True
assert protocol.rollout_engines == engines[:3]
assert protocol.distributed_rollout_engines == engines[3:]
connect.assert_called_once_with(protocol.args, "miles", engines[3:], engine_gpu_counts=[2])
connect.assert_called_once_with("miles", engines[3:], engine_gpu_counts=[2])
disconnect.assert_not_called()
assert protocol._model_update_groups is connect.return_value
assert protocol._ipc_engine is engines[expected_engine_index]
+4 -2
View File
@@ -13,9 +13,9 @@ from tests.fast.charts.utils import (
with_object_names,
)
from tests.fast.utils.external_utils.command_utils.helm_backend.launcher.values import utils as values_utils
from tests.fast.utils.external_utils.command_utils.helm_backend.launcher.values.utils import build_values_as_launched
from miles.utils.external_utils.colocate_pairing.pods import _GATE_NAME, release_patch
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.builder import build_values
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.helm_values_types import (
_PLATFORM_OWNED_ENV_VARS,
)
@@ -317,7 +317,9 @@ class TestTheNamesTheChartWritesAreTheNamesTheSchemaReserves:
def _sub_node_engine_argv() -> list[str]:
specs = [values_utils.engine(num_cells=2, gpus_per_engine=4), values_utils.trainer(num_cells=1, gpus_per_cell=8)]
plan = values_utils.LAYOUT.model_copy(update={"colocate": True})
return build_values(specs, plan).as_values()["run"]["inferenceEngines"][0]["command"]
return build_values_as_launched(specs, plan, scaling=values_utils.SCALING).as_values()["run"]["inferenceEngines"][
0
]["command"]
@requires_helm
+7 -8
View File
@@ -20,11 +20,12 @@ from tests.fast.charts.utils import (
requires_helm,
)
from tests.fast.launch_scripts.sh_harness import REPO_ROOT, SANDBOX_PLACEHOLDER
from tests.fast.utils.external_utils.command_utils.helm_backend.launcher.values.utils import build_values_as_launched
from miles.ray.specs.entrypoint import compute_specs
from miles.utils.args.configs.scaling import ScalingConfig
from miles.utils.arguments import parse_args
from miles.utils.external_utils.command_utils.common import rsync_cmd
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.builder import build_values
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.misc import LaunchPlan
from miles.utils.external_utils.model_args_utils import load_model_args
from miles.utils.test_utils.snapshot import assert_matches_snapshot
@@ -195,14 +196,11 @@ def _dump_values(values: dict[str, Any]) -> str:
_NO_WRAP = 1 << 30
def synthetic_specs() -> list[Any]:
with override_argv(SCENARIO_ARGV):
return compute_specs(parse_args())
def synthetic_run_values() -> dict[str, Any]:
return build_values(
synthetic_specs(),
with override_argv(SCENARIO_ARGV):
args = parse_args()
return build_values_as_launched(
compute_specs(args),
LaunchPlan(
run_id=RUN_ID,
release=RUN_RELEASE_NAME,
@@ -214,6 +212,7 @@ def synthetic_run_values() -> dict[str, Any]:
colocate=True,
prepare_cmd=PREPARE_CMD,
),
scaling=ScalingConfig.slice_from(args),
).as_values()
@@ -13,20 +13,26 @@ from tests.fast.charts.utils import (
sole_container_of,
with_object_names,
)
from tests.fast.utils.external_utils.command_utils.helm_backend.launcher.values.utils import (
SCALING,
build_values_as_launched,
)
from tests.fast.utils.workers.fake_specs import FakeServeSpec
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.builder import build_values
from miles.utils.args.runtime_base import BaseLeafConfig
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.misc import LaunchPlan
from miles.utils.workers.worker_spec import BaseServeSpec, SchedulingSpec
from miles.utils.workers.worker_spec import DEFAULT_RPC_PORT_INFO, SchedulingSpec
def _rollout_executor() -> BaseServeSpec:
return BaseServeSpec(
def _rollout_executor() -> FakeServeSpec:
return FakeServeSpec(
name="rollout-executor",
port_infos=[],
env_var=lambda context: {},
scheduling=SchedulingSpec(num_cells=1, num_workers_per_cell=1, num_gpus_per_worker=0, num_cpus_per_worker=1),
port_infos=[DEFAULT_RPC_PORT_INFO],
fixed_scheduling=SchedulingSpec(
num_cells=1, num_workers_per_cell=1, num_gpus_per_worker=0, num_cpus_per_worker=1
),
args=BaseLeafConfig(),
worker_class="miles.ray.rollout.rollout_executor.RolloutExecutor",
ctor_kwargs=lambda context: {},
)
@@ -159,7 +165,7 @@ class TestStaticWorkers:
class TestGeneratedStaticWorkerShape:
def test_accepts_the_pool_the_launcher_writes_on_every_entry(self):
"""The values builder stamps pool_id on all three sections, so a schema without it rejects every run."""
generated = build_values(
generated = build_values_as_launched(
[_rollout_executor()],
LaunchPlan(
run_id=RUN_ID,
@@ -169,6 +175,7 @@ class TestGeneratedStaticWorkerShape:
orchestrator_command=["python", "train.py"],
worker_argv=["--cluster-backend", "kubernetes"],
),
scaling=SCALING,
).as_values()
entries = generated["run"]["staticWorkers"]
+2
View File
@@ -162,6 +162,8 @@ def _make_args(dump_dir: Path, *, num_prompts: int, n_samples_per_prompt: int) -
reward_key=None,
qkv_format="thd",
enable_sample_ownership_checker=False,
use_rollout_routing_replay=False,
use_rollout_indexer_replay=False,
)
+3 -2
View File
@@ -3,13 +3,14 @@ import logging
import pytest
from miles.dashboard.args import add_dashboard_arguments, collector_config_from_args, validate_dashboard_args
from miles.dashboard.args import collector_config_from_args, validate_dashboard_args
from miles.dashboard.sglang_scraper import DEFAULT_METRIC_WHITELIST
from miles.utils.args.configs.dashboard import DashboardConfig
def parse(argv):
parser = argparse.ArgumentParser()
add_dashboard_arguments(parser)
DashboardConfig.add_arguments(parser=parser)
return parser.parse_args(argv)
+5 -1
View File
@@ -1,4 +1,5 @@
import os
import sys
from pathlib import Path
import pytest
@@ -14,11 +15,14 @@ _PROXY_ENV_VARS: tuple[str, ...] = ("http_proxy", "https_proxy", "HTTP_PROXY", "
@pytest.fixture
def scenario_harness(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> ScenarioHarness:
def scenario_harness(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path, request: pytest.FixtureRequest
) -> ScenarioHarness:
for name in [name for name in os.environ if name.startswith(SCRIPT_ENV_VAR_PREFIX)]:
monkeypatch.delenv(name)
for name in _PROXY_ENV_VARS:
monkeypatch.delenv(name, raising=False)
monkeypatch.setattr(sys, "argv", [str(request.path)])
monkeypatch.setenv(f"{SCRIPT_ENV_VAR_PREFIX}RUN_ID", SCENARIO_RUN_ID)
monkeypatch.setenv("MILES_TEST_DUMPS_ROOT", str(tmp_path / "dumps"))
@@ -13,6 +13,7 @@ from tests.e2e.deploy.conftest_deploy.hot_restart.driver import ScheduledFreeze
from tests.e2e.deploy.conftest_deploy.hot_restart.scenario_hot_restart_deterministic import compute_checkpoint_dir
from tests.utils.deploy.hot_restart.evidence import HotRestartRecord
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
from miles.ray.rollout.rollout_executor import compute_rollout_checkpoint_dir
from miles.utils.audit_utils.event_logger import checkpoint as event_logger_checkpoint
from miles.utils.audit_utils.event_logger.logger import EVENTS_DIRNAME, EventLogger
@@ -45,13 +46,16 @@ class _Run:
@property
def megatron_args(self) -> Namespace:
return Namespace(
args = Namespace(
save=str(self.checkpoint_dir),
load=str(self.checkpoint_dir),
requested_load=str(self.checkpoint_dir),
megatron_config=None,
save_debug_event_data=str(self.events_dir),
use_critic=False,
)
args.raw_megatron = resolve_megatron_config(args, base_args={})
return args
def train(self, *rollout_ids: int) -> None:
for rollout_id in rollout_ids:
@@ -13,6 +13,7 @@ from tests.e2e.deploy.conftest_deploy.hot_restart.scenario_hot_restart_determini
from tests.utils.deploy.hot_restart.evidence import HotRestartRecord
from miles.backends.megatron_utils.checkpoint_tracker import read_checkpoint_tracker_iteration
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
from miles.ray.rollout.rollout_executor import compute_rollout_checkpoint_dir
from miles.utils.audit_utils.event_logger import checkpoint as event_logger_checkpoint
from miles.utils.audit_utils.event_logger.logger import EVENTS_DIRNAME, EventLogger
@@ -42,13 +43,16 @@ class _Run:
@property
def megatron_args(self) -> Namespace:
return Namespace(
args = Namespace(
save=str(self.checkpoint_dir),
load=str(self.checkpoint_dir),
requested_load=str(self.checkpoint_dir),
megatron_config=None,
save_debug_event_data=str(self.events_dir),
use_critic=False,
)
args.raw_megatron = resolve_megatron_config(args, base_args={})
return args
def train(self, *rollout_ids: int) -> None:
for rollout_id in rollout_ids:
@@ -92,10 +92,9 @@ class TestTheSoakOfOneDeployment:
assert form.event_log is soak["event_log"]
assert launch.config is soak["config"]
assert form.launch_spec.config == soak["config"]
assert (
f"{form.launch_spec.train_args} --deploy-component {launch.config.deploy_component.value}"
== launch.request.train_args
)
relaunched = f"{form.launch_spec.train_args} --deploy-component {launch.config.deploy_component.value}"
snapshot_name = launch.value_of("--config-snapshot-name")
assert launch.request.train_args == f"{relaunched} --config-snapshot-name {snapshot_name}"
def test_the_observer_watches_the_release_checkpoints_and_events_of_this_run(
self, harness: ScenarioHarness
@@ -0,0 +1,31 @@
from pathlib import Path
import pytest
from tests.e2e.long.verifiers_contract import _assert_contract_report
class TestVerifiersSdkContractReport:
def test_three_completed_cases_are_accepted(self, tmp_path: Path) -> None:
"""All three real SDK witnesses must complete successfully."""
report = tmp_path / "report.xml"
report.write_text("<testsuites><testsuite>" + "<testcase/>" * 3 + "</testsuite></testsuites>")
_assert_contract_report(report)
@pytest.mark.parametrize("status", ["skipped", "failure", "error"])
def test_incomplete_or_unsuccessful_cases_are_rejected(self, tmp_path: Path, status: str) -> None:
"""A zero pytest exit code cannot hide a skipped SDK witness."""
report = tmp_path / "report.xml"
report.write_text(f"<testsuite><testcase/><testcase/><testcase><{status}/></testcase></testsuite>")
with pytest.raises(AssertionError, match="did not pass"):
_assert_contract_report(report)
@pytest.mark.parametrize("count", [0, 2, 4])
def test_missing_or_extra_cases_are_rejected(self, tmp_path: Path, count: int) -> None:
"""A changed collection cannot silently remove a required SDK witness."""
report = tmp_path / "report.xml"
report.write_text("<testsuite>" + "<testcase/>" * count + "</testsuite>")
with pytest.raises(AssertionError, match="Expected three"):
_assert_contract_report(report)
+29
View File
@@ -0,0 +1,29 @@
from pathlib import Path
import pytest
from tests.fast.utils.command_recorder import patch_helper, record_commands
from miles.utils.external_utils.command_utils.ray_backend.backend import RayCommandBackend
@pytest.fixture
def eval_launch_commands(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> list[str]:
commands = record_commands(monkeypatch)
patch_helper(monkeypatch, "convert_checkpoint", lambda self, **kwargs: None)
patch_helper(monkeypatch, "hf_download_dataset", lambda self, *args, **kwargs: None)
patch_helper(monkeypatch, "_check_has_nvlink", lambda self: False, backend_class=RayCommandBackend)
for name in (
"RAY_ADDRESS",
"WANDB_API_KEY",
"NCCL_NVLS_ENABLE",
"http_proxy",
"https_proxy",
"HTTP_PROXY",
"HTTPS_PROXY",
):
monkeypatch.delenv(name, raising=False)
monkeypatch.setenv("MILES_SCRIPT_CLUSTER_BACKEND", "ray")
monkeypatch.setenv("MILES_SCRIPT_ENABLE_RAY_SUBMIT", "1")
monkeypatch.setenv("MILES_SCRIPT_OUTPUT_DIR", str(tmp_path / "output"))
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
return commands
@@ -0,0 +1,43 @@
import json
import runpy
import shlex
from pathlib import Path
import pytest
from miles.utils.audit_utils.config_snapshot.generated_values import GENERATED_VALUES_ENV_VAR
from miles.utils.external_utils import command_utils
from miles.utils.test_utils.snapshot import SNAPSHOT_RECORD_DIR_ENV_VAR
@pytest.mark.parametrize("snapshot_enabled", [False, True])
def test_original_eval_modes_register_owned_directories_before_worker_launch(
eval_launch_commands: list[str], monkeypatch: pytest.MonkeyPatch, tmp_path: Path, snapshot_enabled: bool
) -> None:
"""The original script transports each mode's allocated directory only for snapshot recording."""
if snapshot_enabled:
monkeypatch.setenv(SNAPSHOT_RECORD_DIR_ENV_VAR, str(tmp_path / "records"))
else:
monkeypatch.delenv(SNAPSHOT_RECORD_DIR_ENV_VAR, raising=False)
script = Path(command_utils.repo_base_dir) / "tests/e2e/megatron/test_qwen3_4b_fully_async_eval.py"
runpy.run_path(str(script), run_name="__main__")
submissions = [command for command in eval_launch_commands if "ray job submit" in command]
assert len(submissions) == 3
for mode, command in zip(("shared", "fleet", "external"), submissions, strict=True):
words = shlex.split(command)
environment = json.loads(
next(word.split("=", 1)[1] for word in words if word.startswith("--runtime-env-json="))
)["env_vars"]
if not snapshot_enabled:
assert GENERATED_VALUES_ENV_VAR not in environment
continue
values = json.loads(environment[GENERATED_VALUES_ENV_VAR])
[allocated] = [
value
for value in values
if value["kind"] == "temporary_directory" and value["name"] == f"fully_async_eval_{mode}"
]
assert allocated["value"].startswith(f"/dev/shm/miles_eval_{mode}_")
if mode != "shared":
assert words[words.index("--eval-hf-dir") + 1] == allocated["value"]
+18 -8
View File
@@ -6,6 +6,7 @@ from typing import Any
import pytest
import yaml
from tests.fast.charts.conftest import vendored_dependencies
from tests.fast.charts.utils import RUN_CHART_DIR, documents_of, requires_helm
from tests.fast.e2e.external_rollout_script import load_external_rollout_script
from tests.fast.launch_scripts.sh_harness import REPO_ROOT
@@ -17,10 +18,12 @@ from miles.utils.external_utils.command_utils.helm_backend.launcher import comma
from miles.utils.external_utils.command_utils.helm_backend.launcher.command_wrapper import Helm
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.misc import LaunchPlan
from miles.utils.external_utils.model_args_utils import shell_safe_model_args
from miles.utils.workers.serving.utils import override_argv
from miles.utils.object_store_config import MOONCAKE_MASTER_ADDRESS_KEY
from miles.utils.workers.serving.utils import override_argv, parse_orchestrator_argv, parse_serve_worker_config
from miles.utils.workers.types import ClusterBackend
script = load_external_rollout_script()
_ = vendored_dependencies
NAMESPACE = "rl"
RUN_ID = "260101-000000-000"
@@ -44,7 +47,6 @@ MODEL_CONFIG_JSON = """\
}
"""
EXTERNAL_ROLLOUT_FLAG = "--rollout-external-engine-addrs"
CONTROLLER_POOL = "inference-controller"
TRAINER_POOL = "trainer-engine-actor"
ENGINE_POOL_PREFIX = "inference-engine"
@@ -139,9 +141,9 @@ def launch(monkeypatch, sandbox: Path) -> _Launch:
planned: list[LaunchPlan] = []
build_values = entrypoint.build_values
def record_plan(specs, plan):
def record_plan(specs, plan, *, scaling, static_connections):
planned.append(plan)
return build_values(specs, plan)
return build_values(specs, plan, scaling=scaling, static_connections=static_connections)
monkeypatch.setattr(entrypoint, "build_values", record_plan)
@@ -202,7 +204,7 @@ class TestTheScriptOwnArgvSelectsTheExternalPath:
args = parse_args()
assert args.rollout_external
assert args.custom_inference_engine_provider_path == STATIC_ENGINE_PROVIDER
assert args.custom_inference_engine_provider_path.path == STATIC_ENGINE_PROVIDER
@requires_helm
@@ -217,9 +219,17 @@ class TestTheScriptOwnArgvSurvivesTheWholeLauncher:
def test_the_master_the_pods_dial_is_the_one_this_release_installs(self, monkeypatch, tmp_path):
"""The script names a loopback address, which is nothing at all from another pod."""
launched = launch(monkeypatch, tmp_path)
orchestrator = parse_orchestrator_argv(launched.values["run"]["orchestrator"]["command"])
controller_command = named_pool_entry(launched.values, CONTROLLER_POOL)["command"]
controller = parse_serve_worker_config(controller_command[controller_command.index("--config") + 1])
assert launched.values["run"]["mooncake"]["enabled"] is True
assert "127.0.0.1" not in " ".join(launched.plan.worker_argv)
dialed = {
payload["mooncake_store_init_kwargs"][MOONCAKE_MASTER_ADDRESS_KEY]
for payload in (orchestrator.args, controller.args)
}
assert len(dialed) == 1, f"the orchestrator and a served pod dial different mooncake masters: {dialed}"
assert not any(address.startswith("127.0.0.1") for address in dialed)
def test_the_run_declares_no_inference_engine_pool_of_its_own(self, monkeypatch, tmp_path):
"""External rollout means miles provisions none, and one rendered anyway would take gpus and idle."""
@@ -245,8 +255,8 @@ class TestTheScriptOwnArgvSurvivesTheWholeLauncher:
)
).addrs
start = command.index(EXTERNAL_ROLLOUT_FLAG) + 1
assert command[start : start + len(addrs)] == addrs
worker_config = parse_serve_worker_config(command[command.index("--config") + 1])
assert worker_config.args["rollout_external_engine_addrs"] == addrs
def test_the_engines_the_script_wrote_are_installed_with_the_run(self, monkeypatch, tmp_path):
"""The whole point of the kubernetes half: the engines ride along in the release that trains against them."""
@@ -4,7 +4,9 @@ from typing import Any
import pytest
from tests.e2e import conftest_multi_policy as e2e
from tests.fast.utils.env_report.conftest import make_args
from miles.backends.megatron_utils.megatron_config import MegatronArgsNamespace
from miles.utils.audit_utils.event_logger.models import (
EnvReport,
EnvReportArgsDump,
@@ -13,6 +15,7 @@ from miles.utils.audit_utils.event_logger.models import (
MetricEvent,
)
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity, TrainProcessIdentity
from miles.utils.env_report.collector import _dump_args
MEGATRON_CONFIG = dict(
trainers=[
@@ -46,6 +49,12 @@ def _make_report(*, model_id: str | None, rank: int = 0, values: dict[str, Any])
def _reports_of(model_id: str, **overrides: Any) -> list[EnvReportEvent]:
[trainer] = [entry for entry in MEGATRON_CONFIG["trainers"] if entry["model_id"] == model_id]
values = {**trainer["overrides"], "trainer_model_id": model_id, **overrides}
values = _dump_args(
make_args(
**{name: value for name, value in values.items() if name != "num_layers"},
backend=MegatronArgsNamespace(num_layers=values["num_layers"]),
)
).values
return [_make_report(model_id=model_id, rank=rank, values=values) for rank in (0, 1) for _ in range(2)]
@@ -52,7 +52,7 @@ def _tool(shot_width=640, display_width=1920, focused=True):
def _run(coro):
return asyncio.get_event_loop().run_until_complete(coro)
return asyncio.run(coro)
def test_press_sequence_is_one_call_per_key():
@@ -19,7 +19,7 @@ from examples.experimental.hud.sglang_compat import _SglangTokenIds
def _run(coro):
return asyncio.get_event_loop().run_until_complete(coro)
return asyncio.run(coro)
class _CannedTransport:
@@ -3,6 +3,7 @@ import sys
from argparse import Namespace
from contextlib import asynccontextmanager
from types import SimpleNamespace
from typing import Any
import pytest
@@ -29,7 +30,9 @@ from examples.experimental.verifiers.verifiers_rollout import (
)
from tests.fast.train_parallel_config_utils import make_train_parallel_config
from miles.backends.sglang_utils.sglang_config import SglangConfig
from miles.rollout.base_types import BaseRolloutFn, RolloutFnConstructorInput
from miles.utils.args.custom_view import ImmutableNamespace
from miles.utils.types import Sample
@@ -48,7 +51,7 @@ def _args(**overrides) -> Namespace:
"sglang_router_ip": "127.0.0.1",
"sglang_router_policy": "round_robin",
"sglang_router_port": 30000,
"sglang_tokenizer_path": None,
"sglang": SglangConfig(models=[], base_args={"tokenizer_path": None, "enable_deterministic_inference": False}),
}
values.update(overrides)
return Namespace(**values)
@@ -191,7 +194,8 @@ def test_renderer_identity_is_inferred_from_standard_checkpoint_paths(checkpoint
assert _renderer_identity(checkpoint) == expected
def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monkeypatch):
def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monkeypatch: pytest.MonkeyPatch) -> None:
"""Local tokenizer files retain the checkpoint's registered renderer identity."""
renderers = pytest.importorskip("renderers", minversion="0.1.8")
checkpoint = "/cache/models--Qwen--Qwen3-4B-Instruct-2507/snapshots/revision"
seen = {}
@@ -224,7 +228,7 @@ def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monke
monkeypatch.setattr("renderers.base.load_tokenizer", load_tokenizer)
runtime = SimpleNamespace(TrainClient=BaseTrainClient)
args = _args(sglang_tokenizer_path="/models/custom-tokenizer")
args = _args(sglang=SglangConfig(models=[], base_args={"tokenizer_path": "/models/custom-tokenizer"}))
client = _train_client(runtime, args, checkpoint, pool_size=3)
pool = client._renderer_pool(checkpoint, chat_template_kwargs={"enable_thinking": False})
@@ -239,6 +243,33 @@ def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monke
}
def test_canonical_tokenizer_selects_tool_renderer_for_ambiguous_local_checkpoint(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Structured tokenizer identity disambiguates local Qwen checkpoints for tools."""
renderers = pytest.importorskip("renderers", minversion="0.1.8")
runtime = pytest.importorskip("verifiers.v1.clients.train")
checkpoint = "/root/models/Qwen3-0.6B"
loaded_sources: list[str] = []
def load_tokenizer(source: str) -> SimpleNamespace:
loaded_sources.append(source)
return SimpleNamespace(name_or_path=source, convert_tokens_to_ids=lambda token: 1, unk_token_id=0)
monkeypatch.setattr("renderers.base.load_tokenizer", load_tokenizer)
args = ImmutableNamespace.model_validate(
vars(_args(sglang=SglangConfig(models=[], base_args={"tokenizer_path": "Qwen/Qwen3-0.6B"})))
)
assert _renderer_identity(checkpoint) is None
client = _train_client(runtime, args, checkpoint, pool_size=1)
pool = client._renderer_pool(checkpoint)
assert isinstance(pool, renderers.RendererPool)
assert pool.supports_tools is True
assert loaded_sources == ["Qwen/Qwen3-0.6B"]
@pytest.mark.asyncio
async def test_train_client_reports_unsupported_tool_renderer_as_configuration_error():
pytest.importorskip("renderers", minversion="0.1.8")
@@ -473,35 +504,45 @@ def test_group_reward_train_count_is_ignored_for_eval_only_runs():
@pytest.mark.asyncio
async def test_verifiers_episode_owns_group_reward_computation():
@pytest.mark.parametrize("deterministic", [False, True])
async def test_verifiers_episode_owns_group_reward_computation(deterministic: bool) -> None:
"""Episodes own rewards and receive distinct seeds only when deterministic."""
pytest.importorskip("verifiers", minversion="0.2.0")
pytest.importorskip("renderers", minversion="0.1.8")
traces = [_trace(id="a", reward=0.0), _trace(id="b", reward=0.0)]
runtime = _import_verifiers()
ctx = runtime.ModelContext(client=object(), model="test-model", sampling=runtime.SamplingConfig(sampling_seed=19))
rollouts = [SimpleNamespace(ctx=ctx), SimpleNamespace(ctx=ctx)]
class Episode:
rollouts = []
def __init__(self) -> None:
self.rollouts = rollouts
async def run(self, semaphore):
async def run(self, semaphore: asyncio.Semaphore) -> list[SimpleNamespace]:
assert semaphore is not None
traces[0].reward = -1.0
traces[1].reward = 1.0
return traces
class Environment:
def episode(self, task, ctx, n):
def episode(self, task: str, ctx: Any, n: int) -> Episode:
assert task == "task"
assert ctx == "ctx"
assert ctx.model == "test-model"
assert n == 2
return Episode()
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args(sglang_enable_deterministic_inference=False)
adapter.args = _args(sglang=SglangConfig(models=[], base_args={"enable_deterministic_inference": deterministic}))
adapter.env = Environment()
adapter.ctx = "ctx"
adapter.ctx = ctx
adapter.model = "test-model"
result = await adapter._run_task_group("task", 2, asyncio.Semaphore(2), seed_base=0)
result = await adapter._run_task_group("task", 2, asyncio.Semaphore(2), seed_base=41)
assert [trace.reward for trace in result] == [-1.0, 1.0]
assert [rollout.ctx.sampling.sampling_seed for rollout in rollouts] == ([41, 42] if deterministic else [19, 19])
assert ctx.sampling.sampling_seed == 19
assert all(rollout.ctx.client is ctx.client and rollout.ctx.model == ctx.model for rollout in rollouts)
def test_sampling_config_preserves_miles_minimum_tokens():
@@ -540,7 +581,7 @@ def test_eval_args_clear_training_prompt_cap_and_preserve_other_fallbacks():
rollout_max_response_len=8,
)
eval_args = _make_eval_args(args)
eval_args = _make_eval_args(ImmutableNamespace.model_validate(vars(args)))
assert eval_args.rollout_max_context_len == 128
assert eval_args.rollout_max_prompt_len is None
@@ -6,8 +6,10 @@ from types import SimpleNamespace
import pytest
from examples.multi_policy import solver_verifier
from examples.multi_policy.solver_verifier import _Verdict
from tests.fast.fixtures.args_fixtures import parser_defaults, resolve_parse_boundary_configs
from tests.fast.fixtures.megatron_config_fixtures import encode_megatron_config
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput
from miles.rollout.generate_hub import single_turn
from miles.utils.types import Sample
@@ -39,19 +41,16 @@ class _FakeTokenizer:
def _make_input(*, prompt: str | list[dict[str, str]], label: str) -> GenerateFnInput:
args = Namespace(
megatron_config=encode_megatron_config("solver", "verifier"),
use_critic=False,
sglang_model_routers={"solver": ("solver-host", 1111), "verifier": ("verifier-host", 2222)},
sglang_router_policy="round_robin",
sglang_speculative_algorithm=None,
rollout_max_response_len=16,
rollout_max_context_len=None,
use_rollout_routing_replay=False,
use_rollout_indexer_replay=False,
use_sampling_support_replay=False,
lora_rank=0,
lora_adapter_path=None,
**{
**parser_defaults(),
"megatron_config": encode_megatron_config("solver", "verifier"),
"sglang_model_routers": {"solver": ("solver-host", 1111), "verifier": ("verifier-host", 2222)},
"sglang_router_policy": "round_robin",
"rollout_max_response_len": 16,
"rollout_num_gpus": 2,
}
)
resolve_parse_boundary_configs(args)
state = SimpleNamespace(args=args, tokenizer=_FakeTokenizer(), processor=None)
sample = Sample(group_index=3, index=7, prompt=prompt, label=label)
return GenerateFnInput(state=state, sample=sample, sampling_params={}, evaluation=False)
@@ -334,6 +333,7 @@ class TestGenerate:
monkeypatch.setattr(solver_verifier, "single_turn_generate", fake)
input = _make_input(prompt=[dict(role="user", content="What is 9 + 9?")], label="#### 18")
input.args.megatron_config = encode_megatron_config("solver")
input.args.raw_megatron = resolve_megatron_config(input.args, base_args={})
with pytest.raises(AssertionError, match="pairs one solver policy with one verifier policy"):
await solver_verifier.generate(input)
+121 -2
View File
@@ -3,20 +3,30 @@ from __future__ import annotations
import argparse
import contextlib
import functools
import os
import sys
from argparse import Namespace
from collections.abc import Iterator
from typing import Any
from typing import Any, TypeVar
from unittest.mock import patch
from miles.utils.arguments import get_miles_extra_args_provider
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
from miles.backends.sglang_utils.sglang_config import SglangConfig
from miles.utils.args.configs.router import RouterConfig
from miles.utils.args.runtime import AllConfig, TrainerConfig
from miles.utils.arguments import _compute_init_expected_num_cells, get_miles_extra_args_provider, parse_args
from miles.utils.run_uuid import RUN_UUID_LENGTH
_ConfigT = TypeVar("_ConfigT", AllConfig, TrainerConfig)
# megatron's own parser adds these and miles' code reads them, but a unit test builds only the miles
# extras, so nothing else would put them on the namespace
_TRAIN_BACKEND_DEFAULTS: dict[str, Any] = dict(
disable_param_buffers_cpu_backup=False,
fp16=False,
lr_warmup_iters=None,
load=None,
num_layers=None,
)
# declared with no default and resolved after parsing, so the raw parser value is one no production
@@ -26,10 +36,19 @@ _RESOLVED_AFTER_PARSING: dict[str, Any] = dict(
offload_rollout=False,
eval_uses_snapshots=False,
starts_inference_engines=True,
use_critic=False,
rollout_external=False,
multi_lora=False,
use_sampling_support_replay=False,
run_uuid="0" * RUN_UUID_LENGTH,
)
class ConfigNamespace(argparse.Namespace):
def __iter__(self) -> Iterator[tuple[str, Any]]:
return iter(vars(self).items())
@contextlib.contextmanager
def _with_relaxed_parser_required_args(parser: argparse.ArgumentParser) -> Iterator[None]:
required = [action for action in parser._actions if action.required]
@@ -52,3 +71,103 @@ def parser_defaults() -> dict[str, Any]:
with _with_relaxed_parser_required_args(parser), patch.object(sys, "argv", ["test"]):
parsed, _ = parser.parse_known_args([])
return {**_TRAIN_BACKEND_DEFAULTS, **vars(parsed), **_RESOLVED_AFTER_PARSING}
def resolve_parse_boundary_configs(args: Namespace) -> Namespace:
args.raw_megatron = resolve_megatron_config(args, base_args={})
args.sglang, args.sglang_scaling = SglangConfig.parse_args(args)
vars(args).update(RouterConfig.from_args(args))
args.init_expected_num_cells = _compute_init_expected_num_cells(
args, sglang=args.sglang, sglang_scaling=args.sglang_scaling
)
return args
_COMMON_TEST_ARGV = [
"--rollout-batch-size",
"2",
"--num-rollout",
"1",
"--actor-num-gpus-per-node",
"1",
"--micro-batch-size",
"1",
]
_MEGATRON_TEST_ARGV = [
"--train-backend",
"megatron",
*_COMMON_TEST_ARGV,
"--num-layers",
"1",
"--hidden-size",
"128",
"--num-attention-heads",
"2",
]
_FSDP_TEST_ARGV = ["--train-backend", "fsdp", *_COMMON_TEST_ARGV]
def parse_megatron_test_config(*argv: str) -> AllConfig:
return _parse_test_config([*_MEGATRON_TEST_ARGV, *argv])
def parse_fsdp_test_config(*argv: str) -> AllConfig:
return _parse_test_config([*_FSDP_TEST_ARGV, *argv])
def _parse_test_config(argv: list[str]) -> AllConfig:
environment = {"RANK": "0", "WORLD_SIZE": "1", "LOCAL_RANK": "0", "MILES_SCRIPT_ENV_REPORT": ""}
with patch.object(sys, "argv", ["test", *argv]), patch.dict(os.environ, environment):
return parse_args()
def make_trainer_args(*, train_backend: str = "megatron", **values: Any) -> ConfigNamespace:
from miles.backends.fsdp_utils.config import FsdpArgsNamespace
from miles.backends.megatron_utils.megatron_config import MegatronArgsNamespace
from miles.utils.args.configs.backend_fields import TrainerBackendTraitConfig
from miles.utils.args.runtime import TrainerConfig
values = {**parser_defaults(), **values, "train_backend": train_backend}
trainer_fields = TrainerConfig.model_fields.keys() - TrainerBackendTraitConfig.model_fields.keys()
backend_cls = MegatronArgsNamespace if train_backend == "megatron" else FsdpArgsNamespace
backend = backend_cls(**{name: value for name, value in values.items() if name not in trainer_fields})
trainer = {name: value for name, value in values.items() if name in trainer_fields}
return ConfigNamespace(**trainer, backend=backend)
def with_backend_values(args: ConfigNamespace, **values: Any) -> ConfigNamespace:
backend = type(args.backend)(**(vars(args.backend) | values))
return ConfigNamespace(**(vars(args) | {"backend": backend}))
def make_trainer_config(**values: Any) -> Any:
from miles.utils.args.runtime import TrainerConfig
args = make_trainer_args(**values)
return TrainerConfig.model_construct(**vars(args))
def replace_config_values(config: _ConfigT, **updates: Any) -> _ConfigT:
fields = type(config).model_fields
backend_updates = {name: value for name, value in updates.items() if name not in fields}
top_level = {name: value for name, value in updates.items() if name in fields}
if backend_updates:
top_level.update(_replace_backend_values(config, **backend_updates))
return config.model_copy(update=top_level)
def _replace_backend_values(config: AllConfig | TrainerConfig, **updates: Any) -> dict[str, Any]:
if isinstance(config, TrainerConfig):
field, backend = "backend", config.backend
elif config.train_backend == "megatron":
field, backend = "raw_megatron", None
else:
field, backend = "raw_fsdp", config.raw_fsdp
values = dict(config.raw_megatron.base_args) if backend is None else vars(backend)
unknown = updates.keys() - values.keys()
assert not unknown, f"{sorted(unknown)} are neither {type(config).__name__} fields nor backend arguments"
if backend is None:
return {field: config.raw_megatron.model_copy(update={"base_args": values | updates})}
return {field: type(backend)(**(values | updates))}
+7 -3
View File
@@ -18,11 +18,15 @@ class FakeBackendCapability(BackendCapability):
self.cells_provider = cells_provider
self.static_provider = static_provider
self.operations = cell_operations
self.requested_pool_ids: list[list[str]] = []
self.requested_pool_ids: list[list[str] | None] = []
self.requested_categories: list[str | None] = []
self.requested_static_pool_ids: list[str] = []
def dynamic_worker_provider(self, *, pool_ids: Sequence[str]) -> BaseWorkerProvider:
self.requested_pool_ids.append(list(pool_ids))
def dynamic_worker_provider(
self, *, pool_ids: Sequence[str] | None, category: str | None = None
) -> BaseWorkerProvider:
self.requested_pool_ids.append(None if pool_ids is None else list(pool_ids))
self.requested_categories.append(category)
assert self.cells_provider is not None, "this capability was built without a cells provider"
return self.cells_provider
+16 -10
View File
@@ -10,12 +10,15 @@ from unittest.mock import patch
import pytest
from miles.ray.rollout.rollout_executor import _compute_rollout_function_config
from miles.rollout.base_types import GenerateFnInput
from miles.rollout.inference_rollout.compatibility import load_generate_function
from miles.rollout.inference_rollout.inference_rollout_common import GenerateState
from miles.rollout.session.config import compute_session_server_config
from miles.rollout.session.server import SessionServer
from miles.rollout.session.types import SessionServerInstance
from miles.utils.args.component_rollout import InferenceRuntimeMutState
from miles.utils.args.custom_view import ImmutableNamespace
from miles.utils.async_utils import run
from miles.utils.http_utils import find_available_port, init_http_client
from miles.utils.misc import SingletonMeta
@@ -94,7 +97,7 @@ def make_sample(
@dataclass
class GenerateEnv:
args: Namespace
args: ImmutableNamespace
mock_server: Any
@@ -153,7 +156,6 @@ def make_args(
rollout_max_context_len: int | None = None,
chat_template_path: str | None = None,
num_layers: int | None = None,
moe_router_topk: int | None = None,
) -> Namespace:
argv = [
"pytest",
@@ -207,21 +209,25 @@ def make_args(
from miles.utils.arguments import parse_args
with patch("sys.argv", argv):
args = parse_args()
# R3 decode shape overrides — not CLI flags (derived from the model config
# in production). Applied here, before with_session_server copies args into
# the worker namespace, because sample assembly runs inside the worker.
if num_layers is not None:
args.num_layers = num_layers
if moe_router_topk is not None:
args.moe_router_topk = moe_router_topk
def override_num_layers(parsed: Namespace) -> None:
if num_layers is not None:
parsed.num_layers = num_layers
with patch("sys.argv", argv):
args = parse_args(preprocess_args=override_num_layers)
args.inference_runtime_mut_state.set_(InferenceRuntimeMutState(engine_count=1, gpu_count=1))
init_http_client(args)
return args
def compute_generate_args(args: Namespace) -> ImmutableNamespace:
return _compute_rollout_function_config(args, args.rollout_function_path)
@contextmanager
def _noop_port(port: int):
"""No-op context manager that just yields the given port."""
@@ -309,7 +315,7 @@ def generation_env(request, variant):
if is_agentic:
mock_tools.AGENTIC_MAX_TURNS = args_kwargs.get("generate_max_turns")
mock_tools.AGENTIC_RETURN_METADATA = args_kwargs.get("agentic_return_metadata")
yield GenerateEnv(args=args, mock_server=mock_server)
yield GenerateEnv(args=compute_generate_args(args), mock_server=mock_server)
mock_tools.AGENTIC_MAX_TURNS = None
mock_tools.AGENTIC_RETURN_METADATA = None
@@ -0,0 +1,152 @@
import json
import sys
from collections.abc import Callable, Iterator
from pathlib import Path
from types import ModuleType, SimpleNamespace
from typing import Any
import pytest
import torch
from safetensors.torch import save_file
from tests.fast.fixtures.args_fixtures import make_trainer_config
from miles.utils.args.custom_function import CustomFunctionConfig
from miles.utils.function_registry import function_registry
from miles_plugins.models.inkling import lora
@pytest.fixture
def inkling_provider_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> Iterator[SimpleNamespace]:
from megatron.core import parallel_state
from miles.backends.megatron_utils import model
from miles.backends.training_utils import parallel
layers = ModuleType("miles_plugins.models.inkling.layers")
layers.__dict__.update(
{
name: type(name, (_UnusedInklingLayer,), {})
for name in ("InklingDenseMLP", "InklingSelfAttention", "InklingSharedExperts")
}
)
monkeypatch.setitem(sys.modules, layers.__name__, layers)
provider_path = "miles_plugins.models.inkling.model.inkling_model_provider"
args = make_trainer_config(
custom_model_provider_path=CustomFunctionConfig(path=provider_path),
model_name="inkling",
megatron_to_hf_mode="raw",
hf_checkpoint=str(tmp_path),
load=None,
pretrained_checkpoint=str(tmp_path),
moe_use_upcycling=False,
lora_rank=2,
lora_alpha=4,
lora_dropout=0,
lora_A_init_method="xavier",
lora_type="lora",
lora_adapter_path=None,
debug_disable_optimizer=True,
stream_optimizer_state_to_disk=False,
enable_witness=False,
optimizer="muon",
muon_split_qkv=True,
lr=0.001,
use_gloo_process_groups=False,
)
monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_rank", lambda: 0)
monkeypatch.setattr(parallel_state, "get_tensor_model_parallel_world_size", lambda: 1)
monkeypatch.setattr(torch.distributed, "get_rank", lambda: 0)
rank = SimpleNamespace(rank=0)
monkeypatch.setattr(
parallel, "get_parallel_state", lambda: SimpleNamespace(pp=rank, tp=rank, cp=rank, intra_dp=rank)
)
monkeypatch.setattr(model, "get_model", _build_local_model)
monkeypatch.setattr(model, "is_first_replica_megatron_main_rank", lambda: True)
monkeypatch.setattr(
model, "get_megatron_muon_optimizer", lambda *, config, **kwargs: SimpleNamespace(config=config)
)
monkeypatch.setattr(model, "get_optimizer_param_scheduler", lambda args, optimizer: None)
monkeypatch.setattr(model, "check_peak_gpu_memory_after_load", lambda args: None)
monkeypatch.setattr(model, "clear_memory", lambda: None)
monkeypatch.setattr(model, "check_model_hashes", lambda args, model, iteration: None)
monkeypatch.setattr(lora, "_UNPADDED_VOCAB_CACHE", [None])
with (
function_registry.temporary(provider_path, _tiny_provider),
function_registry.temporary("models.other.provider", _tiny_provider),
):
yield SimpleNamespace(args=args, module=model)
@pytest.fixture
def inkling_adapter_model(inkling_provider_env: SimpleNamespace) -> torch.nn.Module:
model = _TinyInklingModel()
lora._apply_lm_head_lora(model, inkling_provider_env.args, scale=2, dropout=0, a_init="xavier")
return model
@pytest.fixture
def inkling_tower_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> SimpleNamespace:
from miles.backends.megatron_utils.update_weight import hf_weight_iterator
tensors = {"visual.weight": torch.ones(2), "audio.weight": torch.full((2,), 3), "language.weight": torch.zeros(2)}
save_file(tensors, str(tmp_path / "model.safetensors"))
(tmp_path / "model.safetensors.index.json").write_text(
json.dumps({"weight_map": {name: "model.safetensors" for name in tensors}})
)
monkeypatch.setattr(hf_weight_iterator, "_MM_TOWER_CACHE", None)
monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu")
return SimpleNamespace(module=hf_weight_iterator, checkpoint=tmp_path, tensors=tensors)
@pytest.fixture
def inkling_reload_env(
inkling_provider_env: SimpleNamespace,
inkling_adapter_model: torch.nn.Module,
monkeypatch: pytest.MonkeyPatch,
) -> SimpleNamespace:
env = SimpleNamespace(
args=inkling_provider_env.args,
module=inkling_provider_env.module,
model=inkling_adapter_model,
optimizer=_OptimizerMasters(inkling_adapter_model),
native_optimizer_restored=False,
)
monkeypatch.setattr(
env.module, "load_checkpoint", lambda *args, **kwargs: (0, False, env.native_optimizer_restored)
)
return env
class _UnusedInklingLayer(torch.nn.Module):
def __init__(self) -> None:
raise AssertionError("The tiny Inkling fixture has no decoder layers")
class _TinyInklingModel(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.config = SimpleNamespace(
hidden_size=2, sequence_parallel=False, inkling=SimpleNamespace(logits_mup_width_multiplier=None)
)
self.decoder = SimpleNamespace(layers=[])
self.output_layer = torch.nn.Linear(2, 3, bias=False)
self.post_process = True
def _tiny_provider(*, pre_process: bool, post_process: bool, args: Any) -> torch.nn.Module:
return _TinyInklingModel()
def _build_local_model(provider: Callable[[], torch.nn.Module], model_type: Any) -> list[torch.nn.Module]:
return [provider()]
class _OptimizerMasters:
def __init__(self, model: torch.nn.Module) -> None:
self.model = model
self.masters: dict[str, torch.Tensor] = {}
def reload_model_params(self) -> None:
self.masters = {name: tensor.detach().clone() for name, tensor in self.model.named_parameters()}
+8 -1
View File
@@ -6,7 +6,10 @@ from tests.fast.utils.workers.worker_provider.kubernetes import fake_pod_api
from tests.fast.utils.workers.worker_provider.kubernetes.core.test_provider import FakePodApi
from tests.fast.utils.workers.worker_provider.kubernetes.run_specs import _RELEASE, make_engine_spec, make_router_spec
from miles.backends.sglang_utils.sglang_config import SglangScalingConfig
from miles.utils.args.configs.scaling import ScalingConfig
from miles.utils.workers.backend_capability.kubernetes import KubernetesBackendCapability
from miles.utils.workers.connection_config import build_static_conn_config
from miles.utils.workers.worker_provider.kubernetes.helm import naming
from miles.utils.workers.worker_provider.kubernetes.helm.builder import compute_helm_backend_capability
@@ -17,4 +20,8 @@ ROUTER_HOST = naming.static_worker_host(_RELEASE, "inference-router-0", 0)
def install_workers(*, pods: list[Any] | None = None) -> KubernetesBackendCapability:
fake_pod_api.install(FakePodApi(pods=list(pods or [])))
return compute_helm_backend_capability(specs=[make_router_spec(), make_engine_spec()])
static_connections = build_static_conn_config(
specs=[make_router_spec(), make_engine_spec()],
scaling=ScalingConfig(sglang_scaling=SglangScalingConfig(groups={})),
)
return compute_helm_backend_capability(config=static_connections)
+3 -1
View File
@@ -19,6 +19,7 @@ from miles.rollout.session.server import SessionServer
from miles.rollout.session.types import SessionServerInstance
from miles.router.config import compute_miles_router_config
from miles.router.router import MilesRouter
from miles.utils.args.component_rollout import InferenceRuntimeMutState
from miles.utils.arguments import parse_args
from miles.utils.function_registry import load_function
from miles.utils.http_utils import find_available_port, init_http_client
@@ -81,13 +82,14 @@ def _build_args(*, data_path: str, router_port: int, extra_argv: list[str] | Non
] + (extra_argv or [])
with patch("sys.argv", argv):
args = parse_args()
args.inference_runtime_mut_state.set_(InferenceRuntimeMutState(engine_count=1, gpu_count=1))
init_http_client(args)
return args
@contextmanager
def _with_miles_router(args: Namespace) -> Iterator[UvicornThreadServer]:
config = compute_miles_router_config(args, host=args.sglang_router_ip, port=args.sglang_router_port, num_engines=1)
config = compute_miles_router_config(args, host=args.sglang_router_ip, port=args.sglang_router_port)
router = MilesRouter(config, verbose=False)
server = UvicornThreadServer(router.app, host=config.host, port=config.port)
try:
+120
View File
@@ -0,0 +1,120 @@
from collections.abc import AsyncIterator
from functools import partial
from types import SimpleNamespace
import httpx
import pytest
import serve_tinker
import uvicorn
from miles.ray.train.init_request import TrainerControllerInitRequest
from miles.utils import http_utils
from miles.utils.args.component_rollout import InferenceRuntimeImmutState, InferenceRuntimeMutState
@pytest.fixture
async def tinker_startup(monkeypatch: pytest.MonkeyPatch) -> AsyncIterator[SimpleNamespace]:
controller = _InferenceController()
trainer = _Trainer()
requests: list[httpx.Request] = []
args = SimpleNamespace(
multi_lora=True,
load="base-model",
hf_checkpoint="base-model",
tinker_checkpoint_root="checkpoints",
max_tokens_per_gpu=None,
inference_runtime_mut_state=InferenceRuntimeMutState(),
eval_uses_snapshots=False,
sglang_server_concurrency=8,
use_distributed_post=False,
num_rollout=1,
wandb_run_id=None,
mlflow_run_id=None,
tinker_base_model=None,
multi_lora_n_adapters=1,
lora_alpha=16,
lora_rank=8,
tinker_lora_groups=["attn"],
sglang_router_ip="router",
sglang_router_port=30000,
actor_num_nodes=1,
actor_num_gpus_per_node=2,
raw_megatron=SimpleNamespace(
base_args={
"tensor_model_parallel_size": 1,
"pipeline_model_parallel_size": 1,
"context_parallel_size": 1,
}
),
tinker_server_host="127.0.0.1",
tinker_server_port=10613,
)
hf_config = SimpleNamespace(max_position_embeddings=4096, vocab_size=128)
actor_config = SimpleNamespace(role=serve_tinker.ACTOR_ROLE, trainer_id="actor")
monkeypatch.setattr(serve_tinker, "load_hf_config", lambda _: SimpleNamespace(get_text_config=lambda: hf_config))
monkeypatch.setattr(serve_tinker.ArgvOrchestratorStartupInfo, "create", lambda _: object())
monkeypatch.setattr(serve_tinker, "init_orchestration_script", lambda _, *, disposer: object())
monkeypatch.setattr(serve_tinker, "compute_router_providers", lambda _, *, capability: [])
monkeypatch.setattr(serve_tinker, "resolve_router_addrs", _resolve_router_addrs)
monkeypatch.setattr(serve_tinker, "create_inference_controller_handle", lambda *, capability: controller)
monkeypatch.setattr(serve_tinker, "compute_trainer_configs", lambda _: [actor_config])
monkeypatch.setattr(serve_tinker, "create_trainer_handles", lambda _, **kwargs: {"actor": trainer})
monkeypatch.setattr(serve_tinker.uvicorn, "Server", _Server)
monkeypatch.setattr(http_utils, "_http_client", None)
monkeypatch.setattr(http_utils, "_client_concurrency", 0)
monkeypatch.setattr(http_utils, "_distributed_post_enabled", False)
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(
200,
json={"meta_info": {"output_token_logprobs": [[-0.25, 42]], "finish_reason": {"type": "length"}}},
)
monkeypatch.setattr(
http_utils.httpx, "AsyncClient", partial(httpx.AsyncClient, transport=httpx.MockTransport(respond))
)
try:
yield SimpleNamespace(args=args, controller=controller, trainer=trainer, requests=requests)
finally:
if http_utils._http_client is not None:
await http_utils._http_client.aclose()
class _InferenceController:
def __init__(self) -> None:
self.initialized = False
self.disposed = False
self.state_error: Exception | None = None
async def init(self) -> None:
self.initialized = True
async def get_inference_runtime_immut_state(self) -> InferenceRuntimeImmutState:
assert self.initialized
if self.state_error is not None:
raise self.state_error
return InferenceRuntimeImmutState(engine_count=2, gpu_count=4, eval_engine_count=1)
async def dispose(self) -> None:
self.disposed = True
class _Trainer:
async def init(self, request: TrainerControllerInitRequest) -> None:
pass
async def dispose(self) -> None:
pass
class _Server:
def __init__(self, config: uvicorn.Config) -> None:
pass
async def serve(self) -> None:
pass
async def _resolve_router_addrs(args: SimpleNamespace, *, router_providers: list) -> None:
pass
+25 -1
View File
@@ -1,10 +1,34 @@
from argparse import Namespace
from typing import Any
from tests.fast.fixtures.args_fixtures import parser_defaults
from miles.backends.sglang_utils.sglang_api_client import WorkerType
from miles.backends.sglang_utils.sglang_config import ModelConfig, ServerGroupConfig, SglangConfig
from miles.backends.sglang_utils.sglang_config import (
ModelConfig,
ServerGroupConfig,
SglangConfig,
SglangScalingConfig,
_compute_raw_sglang_config,
)
def resolve_sglang_config(args: Namespace) -> SglangConfig:
config, _ = resolve_sglang_config_and_scaling(args)
return config
def resolve_sglang_config_and_scaling(args: Namespace) -> tuple[SglangConfig, SglangScalingConfig]:
return SglangConfig.resolve(raw=_compute_raw_sglang_config(args), args=args, base_args={})
def make_sglang_config(**base_args: Any) -> SglangConfig:
group = ServerGroupConfig(worker_type=WorkerType.REGULAR, num_gpus_per_engine=1, needs_offload=False)
model = ModelConfig(name="default", model_path=None, server_groups=[group], update_weights=True)
return SglangConfig(models=[model], base_args=base_args)
def with_parser_defaults_and_sglang_config(values: dict[str, Any]) -> dict[str, Any]:
values = {**parser_defaults(), **values}
config, _ = SglangConfig.parse_args(Namespace(**values))
return values | {"sglang": config}
@@ -2,7 +2,9 @@ from types import SimpleNamespace
import pytest
import torch
from pydantic import TypeAdapter
from miles.utils.args.configs.lora import LoraConfig
from miles.utils.lora.hf_lora_targets import resolve_hf_lora_targets
from miles.utils.lora.utils import get_adapter_target_modules
from miles_plugins.models.inkling.lora import _export_dense_mlp, _export_experts, resolve_inkling_adapter_targets
@@ -36,7 +38,9 @@ def test_hf_mlp_selection_matches_existing_native_export(multimodal):
hf_prefix="language_model.layers.1.mlp.experts.",
**{f"w{projection}_{factor}": tensor for projection in (1, 2, 3) for factor in ("A", "B")},
)
plan = _export_dense_mlp(dense, _LocalGather()) + _export_experts(experts, _LocalGather())
plan = _export_dense_mlp(dense, _LocalGather(), hf_checkpoint="/unused") + _export_experts(
experts, _LocalGather(), hf_checkpoint="/unused"
)
weights = {name: value() if callable(value) else value for name, value in plan}
assert "language_model.layers.0.mlp.gate_up_proj" in get_adapter_target_modules(weights)
assert weights["language_model.layers.0.mlp.gate_up_proj.lora_A.weight"] is tensor
@@ -50,3 +54,11 @@ def test_legacy_config_selects_the_same_native_adapters():
native = dict(model_type="inkling_text", mlp_layer_types=["dense", "sparse"], n_shared_experts=1)
assert resolve_hf_lora_targets(legacy) == resolve_hf_lora_targets(native)
resolve_inkling_adapter_targets(legacy, resolve_hf_lora_targets(legacy))
def test_native_adapter_selector_is_a_valid_configured_adapter_target():
"""The resolved native selector must validate as the configuration's adapter targets."""
config = dict(model_type="inkling_text", mlp_layer_types=["dense", "sparse"], n_shared_experts=0)
targets = resolve_inkling_adapter_targets(config, resolve_hf_lora_targets(config))
annotation = LoraConfig.model_fields["lora_adapter_targets"].annotation
assert TypeAdapter(annotation).validate_python(targets) == targets
+60 -7
View File
@@ -3,16 +3,20 @@ from __future__ import annotations
import textwrap
from argparse import ArgumentParser, Namespace
from collections.abc import AsyncIterator
from typing import Any
from typing import Any, TypeVar
from unittest.mock import MagicMock
import pytest
import ray
from sglang_router.launch_router import RouterArgs
from tests.fast.fixtures.args_fixtures import parser_defaults, resolve_parse_boundary_configs
from tests.fast.fixtures.args_fixtures import parser_defaults
from miles.ray.specs.inference import inference_controller_worker_name
from miles.utils import object_store
from miles.utils.args.component_rollout import InferenceRuntimeImmutState, InferenceRuntimeMutState
from miles.utils.args.configs.router import RouterConfig
from miles.utils.args.custom_function import CustomFunctionConfig
from miles.utils.args.runtime import AllConfig, InferenceControllerConfig, RolloutConfig
from miles.utils.types import Sample
@@ -88,7 +92,6 @@ def make_args(**overrides: Any) -> Namespace:
critic_num_gpus_per_node=0,
use_critic=False,
megatron_config=None,
critic_train_only=False,
# sglang router
sglang_router_ip=None,
sglang_router_port=None,
@@ -142,7 +145,6 @@ def make_args(**overrides: Any) -> Namespace:
seed=42,
fp16=False,
use_rollout_indexer_replay=False,
env_report=None,
env_report_interval_seconds=3600.0,
# checkpoint / data source
hf_checkpoint="/fake/model",
@@ -175,14 +177,65 @@ def make_args(**overrides: Any) -> Namespace:
ci_assert_prefill_lag_max=None,
# dumper (sglang debug dumper integration)
dumper_enable=False,
dumper_inference=False,
)
defaults.update(router_defaults)
defaults["inference_runtime_mut_state"] = InferenceRuntimeMutState()
defaults.update(overrides)
defaults.setdefault("starts_inference_engines", not defaults["debug_train_only"] or defaults["eval_num_gpus"] > 0)
if defaults["debug_train_only"]:
defaults["rollout_num_gpus"] = 0
return Namespace(**{**parser_defaults(), **defaults})
return resolve_parse_boundary_configs(Namespace(**{**parser_defaults(), **defaults}))
def make_rollout_config(**overrides: Any) -> RolloutConfig:
"""The typed config the rollout executor receives, sliced from ``make_args``."""
return _slice_config(RolloutConfig, make_args(**overrides))
def make_inference_controller_config(args: Namespace) -> InferenceControllerConfig:
"""The typed config the inference controller receives, sliced from a ``make_args`` namespace."""
return _slice_config(InferenceControllerConfig, args)
def _slice_config(config_class: type[_LeafConfigT], args: Namespace) -> _LeafConfigT:
values = vars(args) | _RESOLVED_ROLLOUT_FIELDS | RouterConfig.from_args(args)
for name in _CUSTOM_FUNCTION_FIELDS:
if isinstance(path := values[name], str):
values[name] = CustomFunctionConfig(path=path)
return config_class.model_validate({name: values[name] for name in config_class.model_fields if name in values})
_LeafConfigT = TypeVar("_LeafConfigT", RolloutConfig, InferenceControllerConfig)
_RESOLVED_ROLLOUT_FIELDS: dict[str, Any] = dict(
ci_enable_metrics_capture=False,
eval_datasets=[],
ckpt_step=None,
num_layers=None,
raw_fsdp=None,
lora_A_init_method="xavier",
lora_B_init_method="zero",
)
_CUSTOM_FUNCTION_FIELDS = [
name
for name, field in AllConfig.model_fields.items()
if field.annotation in {CustomFunctionConfig, CustomFunctionConfig | None}
]
class FakeInferenceTopologyProvider:
"""Serves the inference controller the rollout executor asks for the observed engine topology."""
def __init__(self, state: InferenceRuntimeImmutState) -> None:
self._state = state
def get_handle(self, worker_name: str) -> FakeInferenceTopologyProvider:
assert worker_name == inference_controller_worker_name()
return self
async def get_inference_runtime_immut_state(self) -> InferenceRuntimeImmutState:
return self._state
def make_sample(
@@ -8,11 +8,16 @@ from unittest.mock import MagicMock
import pytest
import ray
from tests.fast.ray.rollout.conftest import make_args, make_samples_grouped
from tests.fast.ray.rollout.conftest import (
FakeInferenceTopologyProvider,
make_args,
make_rollout_config,
make_samples_grouped,
)
from tests.fast.train_parallel_config_utils import make_train_parallel_config
from miles.ray.rollout.debug_data import save_debug_rollout_data
from miles.ray.rollout.rollout_executor import RolloutExecutor
from miles.ray.rollout.rollout_executor import RolloutExecutor, _compute_rollout_function_config
from miles.rollout.base_types import (
BaseRolloutFn,
RolloutFnEvalInput,
@@ -21,6 +26,8 @@ from miles.rollout.base_types import (
RolloutFnTrainOutput,
)
from miles.rollout.checkpoint_eval import CheckpointEvalFn
from miles.utils.args.component_rollout import InferenceRuntimeImmutState
from miles.utils.args.custom_function import CustomFunctionConfig
from miles.utils.types import WeightVersionSpan, WeightVersionsPerCall
from miles.utils.weight_version import max_rollouts_without_published_weight_version
@@ -76,14 +83,16 @@ async def _make_executor(args):
args=args,
router_providers=[_NeverUsedProvider()],
session_server_provider=None,
inference_controller_provider=_NeverUsedProvider(),
inference_controller_provider=FakeInferenceTopologyProvider(
InferenceRuntimeImmutState(engine_count=8, gpu_count=8)
),
)
await executor.init()
return executor
def _make_test_args(**overrides):
return make_args(
return make_rollout_config(
sglang_router_ip="127.0.0.1",
sglang_router_port=30000,
use_wandb=False,
@@ -110,8 +119,7 @@ class TestProcessSetup:
self, ray_local_mode, patch_low_level, http_client_calls
):
"""A snapshot-eval fleet may exist in this mode; init_http_client itself decides whether there is anything to talk to."""
args = _make_test_args()
args.debug_train_only = True
args = _make_test_args(debug_train_only=True)
await _make_executor(args)
@@ -132,8 +140,7 @@ class TestProcessSetup:
self, ray_local_mode, patch_low_level, own_args_resolutions
):
"""No engines and no session servers exist in this mode, so there is nothing to wait for."""
args = _make_test_args()
args.debug_train_only = True
args = _make_test_args(debug_train_only=True)
await _make_executor(args)
@@ -152,10 +159,9 @@ class TestRolloutFunctionConstruction:
"""Replaying dumped rollout data must not build the rollout functions."""
import miles.ray.rollout.rollout_executor as rexec
args = _make_test_args()
args.debug_train_only = True
args.load_debug_rollout_data = str(tmp_path / "rollout-{rollout_id}.pt")
args.rollout_num_gpus = None
args = _make_test_args(
debug_train_only=True, load_debug_rollout_data=str(tmp_path / "rollout-{rollout_id}.pt")
)
monkeypatch.delenv("MILES_USE_LEGACY_ROLLOUT_V1", raising=False)
def fail_if_loaded(*args, **kwargs):
@@ -178,8 +184,7 @@ class TestRolloutFunctionConstruction:
"""Without replay data, debug_train_only still builds both rollout functions."""
import miles.ray.rollout.rollout_executor as rexec
args = _make_test_args()
args.debug_train_only = True
args = _make_test_args(debug_train_only=True)
monkeypatch.delenv("MILES_USE_LEGACY_ROLLOUT_V1", raising=False)
loaded_paths: list[str] = []
@@ -200,8 +205,7 @@ class TestRolloutFunctionConstruction:
class TestGenerate:
async def test_invokes_rollout_fn_with_correct_input_and_returns_dp_split(self, ray_local_mode, patch_low_level):
"""generate passes a train input and returns the samples split per dp rank."""
args = _make_test_args()
args.global_batch_size = 8
args = _make_test_args(global_batch_size=8)
executor = await _make_executor(args)
executor.set_train_parallel_config(make_train_parallel_config(dp_size=2))
@@ -233,8 +237,7 @@ class TestGenerate:
async def test_rejects_samples_generated_under_the_default_weight_version(self, ray_local_mode, patch_low_level):
"""A batch carrying the sglang never-updated version must fail get(), not reach training."""
args = _make_test_args()
args.global_batch_size = 8
args = _make_test_args(global_batch_size=8)
executor = await _make_executor(args)
executor.set_train_parallel_config(make_train_parallel_config(dp_size=2))
@@ -251,8 +254,7 @@ class TestGenerate:
async def test_a_frozen_weight_version_eventually_fails_the_run(self, ray_local_mode, patch_low_level):
"""A driver that never forwards the version must fail loudly, not silently disable staleness filtering."""
args = _make_test_args()
args.global_batch_size = 8
args = _make_test_args(global_batch_size=8)
executor = await _make_executor(args)
executor.set_train_parallel_config(make_train_parallel_config(dp_size=2))
@@ -266,8 +268,7 @@ class TestGenerate:
async def test_publishing_a_version_every_step_keeps_the_run_alive(self, ray_local_mode, patch_low_level):
"""The supported driver publishes after each update, which must never trip the staleness assert."""
args = _make_test_args()
args.global_batch_size = 8
args = _make_test_args(global_batch_size=8)
executor = await _make_executor(args)
executor.set_train_parallel_config(make_train_parallel_config(dp_size=2))
@@ -281,8 +282,7 @@ class TestGenerate:
async def test_does_not_touch_the_inference_side(self, ray_local_mode, patch_low_level):
"""The controller is a driver-side object the executor cannot reach, so generate must not need it."""
args = _make_test_args()
args.global_batch_size = 4
args = _make_test_args(global_batch_size=4)
executor = await _make_executor(args)
executor.set_train_parallel_config(make_train_parallel_config(dp_size=1))
@@ -480,8 +480,7 @@ class TestEval:
async def test_skipped_in_debug_train_only_mode(self, ray_local_mode, patch_low_level):
"""debug_train_only short-circuits eval before the rollout function runs."""
args = _make_test_args()
args.debug_train_only = True
args = _make_test_args(debug_train_only=True)
executor = await _make_executor(args)
@@ -527,8 +526,8 @@ class TestRolloutFunctionLoading:
executor = await _make_executor(args)
assert executor.generate_rollout.path == "pkg.train_fn"
assert executor.eval_generate_rollout.path == "pkg.eval_fn"
assert executor.generate_rollout.path == CustomFunctionConfig(path="pkg.train_fn")
assert executor.eval_generate_rollout.path == CustomFunctionConfig(path="pkg.eval_fn")
@pytest.mark.asyncio
@@ -548,7 +547,9 @@ class TestCustomHooks:
monkeypatch.setattr(
rexec,
"load_function",
lambda path: conversion_hook if path == "pkg.convert" else (lambda *a, **kw: None),
lambda path: (
conversion_hook if path == CustomFunctionConfig(path="pkg.convert") else (lambda *a, **kw: None)
),
)
args = _make_test_args(global_batch_size=4, custom_convert_samples_to_train_data_path="pkg.convert")
@@ -580,7 +581,7 @@ class TestCustomHooks:
monkeypatch.setattr(
rexec,
"load_function",
lambda path: reward_hook if path == "pkg.reward" else (lambda *a, **kw: None),
lambda path: reward_hook if path == CustomFunctionConfig(path="pkg.reward") else (lambda *a, **kw: None),
)
args = _make_test_args(global_batch_size=4, custom_reward_post_process_path="pkg.reward")
@@ -754,7 +755,8 @@ class TestLegacyRolloutProtocol:
await executor.get(rollout_id=11)
assert calls == [(args, 11, executor.data_source, False)]
fn_args = _compute_rollout_function_config(args, args.rollout_function_path)
assert calls == [(fn_args, 11, executor.data_source, False)]
async def test_eval_keeps_the_legacy_call_signature_with_the_evaluation_flag(
self, ray_local_mode, patch_low_level
@@ -775,7 +777,8 @@ class TestLegacyRolloutProtocol:
await executor.eval(rollout_id=12)
assert calls == [(args, 12, executor.data_source, True)]
fn_args = _compute_rollout_function_config(args, args.eval_function_path)
assert calls == [(fn_args, 12, executor.data_source, True)]
class _RecordingMetricChecker:
@@ -915,9 +918,10 @@ class TestCheckpointWithoutARolloutFunction:
"""A run without --load names no checkpoint directory, so there is nothing for the data source to read."""
monkeypatch.delenv("MILES_USE_LEGACY_ROLLOUT_V1", raising=False)
executor = await _make_executor(
_checkpoint_args(tmp_path, load_debug_rollout_data="/nonexistent/rollout_{rollout_id}.pt")
_make_test_args(
save=str(tmp_path), load=None, load_debug_rollout_data="/nonexistent/rollout_{rollout_id}.pt"
)
)
executor.args.load = None
executor.data_source = MagicMock()
await executor.load(0)
@@ -1039,8 +1043,7 @@ class TestCheckpointOfADistinctEvalRolloutFunction:
self, ray_local_mode, patch_low_level, tmp_path
):
"""A run without --load names no checkpoint directory, so neither instance has anything to restore."""
executor = await _make_executor(_checkpoint_args(tmp_path))
executor.args.load = None
executor = await _make_executor(_make_test_args(save=str(tmp_path), load=None))
executor.data_source = MagicMock()
calls: list[tuple[str, str]] = []
executor.generate_rollout = _RecordingRolloutFn("train", calls)
+11 -5
View File
@@ -2,10 +2,10 @@ from __future__ import annotations
import pydantic
import pytest
from tests.fast.fixtures.sglang_config_fixtures import resolve_sglang_config, resolve_sglang_config_and_scaling
from tests.fast.ray.rollout.conftest import make_args, make_sglang_config_yaml
from miles.backends.sglang_utils.sglang_api_client import WorkerType
from miles.backends.sglang_utils.sglang_config import resolve_sglang_config
# ----------------------------- resolve_sglang_config matrix -----------------------------
@@ -16,14 +16,20 @@ def _resolve_yaml(tmp_path, yaml_text: str, **args_overrides):
return resolve_sglang_config(make_args(sglang_config=str(cfg_path), **args_overrides))
def _resolve_yaml_and_scaling(tmp_path, yaml_text: str, **args_overrides):
cfg_path = tmp_path / "config.yaml"
cfg_path.write_text(yaml_text)
return resolve_sglang_config_and_scaling(make_args(sglang_config=str(cfg_path), **args_overrides))
class TestResolveSglangConfigPaths:
def test_default_path_when_no_yaml_or_prefill(self):
args = make_args(rollout_num_gpus=8, sglang_config=None, prefill_num_servers=None)
cfg = resolve_sglang_config(args)
cfg, scaling = resolve_sglang_config_and_scaling(args)
assert len(cfg.models) == 1
assert cfg.models[0].name == "default"
assert cfg.models[0].server_groups[0].worker_type == WorkerType.REGULAR
assert sum(g.num_gpus for m in cfg.models for g in m.server_groups) == 8
assert sum(g.num_gpus for groups in scaling.groups.values() for g in groups) == 8
def test_prefill_num_servers_path(self):
args = make_args(
@@ -49,7 +55,7 @@ class TestResolveSglangConfigPaths:
def test_yaml_path_multi_model_actor_plus_reference(self, tmp_path):
# 8 gpu actor + 4 gpu ref = 12 → must match args.rollout_num_gpus
cfg = _resolve_yaml(
cfg, scaling = _resolve_yaml_and_scaling(
tmp_path,
"sglang:\n"
" - name: actor\n"
@@ -68,7 +74,7 @@ class TestResolveSglangConfigPaths:
rollout_num_gpus=12,
)
assert [m.name for m in cfg.models] == ["actor", "ref"]
assert sum(g.num_gpus for m in cfg.models for g in m.server_groups) == 12
assert sum(g.num_gpus for groups in scaling.groups.values() for g in groups) == 12
# ----------------------------- server group validation matrix ---------------
@@ -1,18 +1,19 @@
from __future__ import annotations
from argparse import Namespace
from pathlib import Path
import pytest
from tests.fast.ray.rollout.conftest import make_args, make_sglang_config_yaml
from tests.fast.fixtures.args_fixtures import parse_megatron_test_config
from tests.fast.ray.rollout.conftest import make_sglang_config_yaml
from tests.fast.utils.workers.fake_ray import FakeRayCluster, FakeRayModule
from miles.ray.placement_group import PlacementGroupInfo
from miles.ray.specs.inference import specs_inference_engine
from miles.ray.specs.inference import InferenceEngineSpec
from miles.utils.args.runtime import AllConfig
from miles.utils.workers.naming import compute_worker_name
from miles.utils.workers.ray_worker_manager import RayWorkerManager
from miles.utils.workers.types import WorkerCommBackend
from miles.utils.workers.worker_spec import BaseCommandSpec, LaunchCommandContext, NamedHostAndPorts
from miles.utils.workers.worker_spec import LaunchCommandContext, NamedHostAndPorts
@pytest.fixture
@@ -26,7 +27,7 @@ def fake_ray_cluster(monkeypatch: pytest.MonkeyPatch) -> FakeRayCluster:
return cluster
def _make_args(*, tmp_path: Path, worker_types: list[str], num_gpus: int, num_gpus_per_engine: int) -> Namespace:
def _make_args(*, tmp_path: Path, worker_types: list[str], num_gpus: int, num_gpus_per_engine: int) -> AllConfig:
config_path = tmp_path / "sglang.yaml"
config_path.write_text(
make_sglang_config_yaml(
@@ -36,30 +37,32 @@ def _make_args(*, tmp_path: Path, worker_types: list[str], num_gpus: int, num_gp
]
)
)
return make_args(
sglang_config=str(config_path),
rollout_num_gpus=num_gpus * len(worker_types),
use_session_server=False,
return parse_megatron_test_config(
"--sglang-config", str(config_path), "--rollout-num-gpus", str(num_gpus * len(worker_types))
)
async def _launch_engines(args: Namespace) -> dict[str, LaunchCommandContext]:
async def _launch_engines(args: AllConfig) -> dict[str, LaunchCommandContext]:
"""Run the real launch pipeline and return, per worker name, the context its launch command got."""
contexts: dict[str, LaunchCommandContext] = {}
def _recording_spec(spec: BaseCommandSpec) -> BaseCommandSpec:
def _record(ctx: LaunchCommandContext) -> str:
class _RecordingEngineSpec(InferenceEngineSpec):
def launch_command(self, ctx: LaunchCommandContext) -> str:
worker_name = compute_worker_name(
pool_id=spec.name, cell_index=ctx.cell_index, worker_in_cell_index=ctx.worker_in_cell_index
pool_id=self.name, cell_index=ctx.cell_index, worker_in_cell_index=ctx.worker_in_cell_index
)
contexts[worker_name] = ctx
return f"launch {worker_name}"
return spec.model_copy(update={"launch_command": _record})
specs = [_recording_spec(spec) for spec in specs_inference_engine(args)]
specs = [
_RecordingEngineSpec(**{name: getattr(spec, name) for name in InferenceEngineSpec.model_fields})
for config in InferenceEngineSpec.slice_configs(args)
for spec in InferenceEngineSpec.create(config)
]
num_slots = sum(
spec.scheduling.num_cells * spec.scheduling.num_workers_per_cell * spec.scheduling.num_gpu_slots_per_worker
(scheduling := spec.scheduling(args)).num_cells
* scheduling.num_workers_per_cell
* scheduling.num_gpu_slots_per_worker
for spec in specs
)
+60 -21
View File
@@ -9,6 +9,7 @@ from types import SimpleNamespace
import pytest
from tests.fast.ray.rollout.conftest import make_args as _make_args
from tests.fast.ray.rollout.conftest import make_rollout_config
import miles.ray.rollout.eval_fleet as eval_fleet_mod
from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient
@@ -20,20 +21,28 @@ from miles.ray.rollout.eval_fleet import (
)
from miles.ray.rollout.rollout_server import RolloutServer
from miles.rollout.checkpoint_eval import EvalSkip
from miles.utils.args.component_rollout import InferenceRuntimeMutState
from miles.utils.context_lock import ContextLock
from miles.utils.workers.rpc.client.misc import RpcWorkerCallError, ServerRestartedError
from miles.utils.workers.worker_handle import WorkerUnreachableError
from miles.utils.workers.worker_spec import HostAndPort
_FLEET_ARGS = dict(
eval_num_gpus=1,
eval_num_gpus_per_engine=1,
sglang_model_routers={"default": ("10.0.0.1", 30000), "eval": ("10.0.0.2", 31000)},
)
_FLEET_INFO = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[1, 1])
def make_args(**overrides):
defaults = dict(
eval_num_gpus=1,
eval_num_gpus_per_engine=1,
sglang_model_routers={"default": ("10.0.0.1", 30000), "eval": ("10.0.0.2", 31000)},
)
defaults.update(overrides)
return _make_args(**defaults)
return _make_args(**{**_FLEET_ARGS, **overrides})
def make_config(**overrides):
return make_rollout_config(**{**_FLEET_ARGS, **overrides})
class FakeEngine:
@@ -53,8 +62,9 @@ class FakeEngine:
class FakeEvalServer:
def __init__(self, engines):
def __init__(self, engines, *, engine_gpu_counts=()):
self._engines = engines
self._engine_gpu_counts = list(engine_gpu_counts)
self.context_lock = ContextLock("FakeEvalServer")
self.router_ip = "10.0.0.2"
self.router_port = 31000
@@ -64,6 +74,11 @@ class FakeEvalServer:
assert self.context_lock.held_in_current_context, "api_clients is read under the server's lock"
return list(self._engines)
@property
def engine_gpu_counts(self):
assert self.context_lock.held_in_current_context, "engine_gpu_counts is read under the server's lock"
return list(self._engine_gpu_counts)
class FlakyEngine(FakeEngine):
def __init__(self, log, *, failures: int):
@@ -102,17 +117,18 @@ def router_ready(monkeypatch):
monkeypatch.setattr(eval_fleet_mod.InferenceControllerEvalFleet, "_wait_router_ready", noop_router_ready)
def make_fleet(args, engines):
return InferenceControllerEvalFleet(args, srv=FakeEvalServer(engines))
def make_fleet(args, engines, *, engine_gpu_counts=()):
return InferenceControllerEvalFleet(args, srv=FakeEvalServer(engines, engine_gpu_counts=engine_gpu_counts))
class TestEvalFleetInfo:
def test_describes_the_fleet_its_router_serves(self):
@pytest.mark.asyncio
async def test_describes_the_fleet_its_router_serves(self):
"""The description the executor retargets its eval args to comes from the server, not its own args."""
fleet = make_fleet(make_args(eval_num_gpus=4, eval_num_gpus_per_engine=2), [])
fleet = make_fleet(make_args(eval_num_gpus=4, eval_num_gpus_per_engine=2), [], engine_gpu_counts=[2, 2])
assert fleet.info == EvalFleetInfo(
router=HostAndPort(host="10.0.0.2", port=31000), num_gpus=4, num_gpus_per_engine=2
assert await fleet.info() == EvalFleetInfo(
router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[2, 2]
)
@@ -284,9 +300,13 @@ class TestRouterProbe:
class FakeInferenceController:
def __init__(self, pins: list[EvalFleetPin]):
def __init__(self, pins: list[EvalFleetPin], *, info: EvalFleetInfo | None = _FLEET_INFO):
self.calls: list[dict] = []
self._pins = pins
self._info = info
async def get_eval_fleet_info(self) -> EvalFleetInfo | None:
return self._info
async def pin_eval_fleet(self, *, checkpoint_dir: str, weight_version: str) -> EvalFleetPin:
self.calls.append(dict(checkpoint_dir=checkpoint_dir, weight_version=weight_version))
@@ -311,7 +331,7 @@ class FakeControllerProvider:
@pytest.fixture
def fleet_states(monkeypatch):
built = []
monkeypatch.setattr(eval_fleet_mod, "GenerateState", lambda args: built.append(args) or f"fake-state-{len(built)}")
monkeypatch.setattr(eval_fleet_mod, "GenerateState", lambda args: built.append(args) or SimpleNamespace(args=args))
return built
@@ -321,8 +341,8 @@ def make_session(controller, *, info=None):
def make_session_over(provider, *, info=None):
return RolloutExecutorEvalFleet(
make_args(),
info=info or EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), num_gpus=2, num_gpus_per_engine=1),
make_config(),
info=info or _FLEET_INFO,
inference_controller_provider=provider,
)
@@ -334,7 +354,7 @@ class TestRolloutExecutorEvalFleet:
(state_args,) = fleet_states
assert (state_args.sglang_router_ip, state_args.sglang_router_port) == ("10.0.0.2", 31000)
assert (state_args.rollout_num_gpus, state_args.rollout_num_gpus_per_engine) == (2, 1)
assert state_args.inference_runtime_mut_state == InferenceRuntimeMutState(engine_count=2, gpu_count=2)
async def test_pins_over_rpc_and_returns_the_cached_state(self, fleet_states):
"""Pinning is the controller's call; the state is built once and handed back per point."""
@@ -348,8 +368,27 @@ class TestRolloutExecutorEvalFleet:
dict(checkpoint_dir="/snap/step_5", weight_version="5"),
dict(checkpoint_dir="/snap/step_6", weight_version="6"),
]
assert first == second == "fake-state-1"
assert len(fleet_states) == 1
assert first is second
assert [state.args for state in (first,)] == fleet_states
assert first.args.inference_runtime_mut_state == InferenceRuntimeMutState(engine_count=2, gpu_count=2)
async def test_a_pin_records_the_engine_topology_the_controller_now_reports(self, fleet_states):
"""The eval fleet may have resized since the state was built, and the state sizes its semaphore off it."""
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[4, 4, 4])
session = make_session(FakeInferenceController([EvalFleetPin(skip_reason=None)], info=info))
state = await session.pin("/snap/step_5", "5")
assert state.args.inference_runtime_mut_state == InferenceRuntimeMutState(engine_count=3, gpu_count=12)
async def test_a_controller_without_an_eval_fleet_skips_the_point(self, fleet_states):
"""A pinned point whose fleet vanished has no engines to generate against, so it is skipped."""
session = make_session(FakeInferenceController([EvalFleetPin(skip_reason=None)], info=None))
with pytest.raises(EvalSkip) as exc:
await session.pin("/snap/step_5", "5")
assert exc.value.reason == "controller_unreachable"
async def test_a_remote_skip_stays_an_attributable_skip(self, fleet_states):
"""The reason the controller skipped for must survive the wire as EvalSkip."""
@@ -291,17 +291,20 @@ class TestStaticInferenceEngineWorkerProvider:
offsets = [info.meta["gpu_offset"] for info in provider.cell_infos]
assert offsets == [0, 2, 4]
async def test_a_fleet_with_a_different_gpu_total_is_rejected(self, monkeypatch):
"""--rollout-num-gpus sizes the placement group and the router, so a fleet that is smaller
than it claims must fail at startup instead of hanging in NCCL."""
async def test_a_fleet_whose_gpu_total_differs_from_the_arguments_is_described_as_observed(self, monkeypatch):
"""The run sizes weight transfer from the observed layout, so each engine keeps the gpus it reported."""
args = _make_args(["host1:8000", "host2:8000"], num_gpus_per_engine=4)
payloads = {
"http://host1:8000": _regular_payload(num_gpus=4),
"http://host2:8000": _regular_payload(num_gpus=2),
}
with pytest.raises(AssertionError, match="6 gpus in total"):
await self._make_provider(monkeypatch, args, payloads)
provider = await self._make_provider(monkeypatch, args, payloads)
assert [(info.meta["num_gpus_per_engine"], info.meta["gpu_offset"]) for info in provider.cell_infos] == [
(4, 0),
(2, 4),
]
async def test_a_pd_fleet_without_the_router_flag_is_rejected(self, monkeypatch):
"""The router was already launched non-PD, so serving a PD fleet behind it would misroute."""
@@ -445,16 +448,15 @@ def _engine(*, url: str, num_gpus: int) -> _ExternalEngineInfo:
class TestARaggedExternalFleet:
def test_engines_that_do_not_each_report_the_per_engine_argument_are_refused(self):
"""Engines of unequal size cannot each hold --rollout-num-gpus-per-engine gpus."""
def test_engines_of_unequal_size_pass_the_argument_check(self):
"""The per-engine argument no longer describes the fleet, so engines of their own size are accepted."""
args = _make_args(["host1:8000", "host2:8000"], num_gpus_per_engine=2, rollout_num_gpus=6)
engines = [_engine(url="http://host1:8000", num_gpus=4), _engine(url="http://host2:8000", num_gpus=2)]
with pytest.raises(AssertionError, match="rollout-num-gpus-per-engine"):
_assert_engines_match_args(args, engines=engines)
_assert_engines_match_args(args, engines=engines)
async def test_a_fleet_of_unequal_engines_is_refused_before_the_run_is_handed_an_engine(self, monkeypatch):
"""Discovery refuses the fleet, so nothing downstream ever sees it."""
async def test_a_fleet_of_unequal_engines_hands_the_run_each_engine_at_its_own_size(self, monkeypatch):
"""Each discovered engine becomes a cell carrying the gpus it reported and the offset after the previous one."""
args = _make_args(["host1:8000", "host2:8000"], num_gpus_per_engine=2, rollout_num_gpus=6)
payloads = {
"http://host1:8000": _regular_payload(num_gpus=4),
@@ -463,9 +465,12 @@ class TestARaggedExternalFleet:
_install_payloads(monkeypatch, payloads)
provider = StaticInferenceEngineWorkerProvider(args=args)
await provider.init()
with pytest.raises(AssertionError, match="rollout-num-gpus-per-engine"):
await provider.init()
assert [(info.meta["num_gpus_per_engine"], info.meta["gpu_offset"]) for info in provider.cell_infos] == [
(4, 0),
(2, 4),
]
class TestAPdFleetRoles:
@@ -1,5 +1,4 @@
import asyncio
from argparse import Namespace
from types import SimpleNamespace
from typing import Any
from unittest.mock import patch
@@ -67,7 +66,7 @@ async def _discovered_server(
)
def _connect(*, engine_gpu_counts: list[int], rollout_num_gpus_per_engine: int) -> list[dict]:
def _connect(*, engine_gpu_counts: list[int]) -> list[dict]:
calls: list[dict] = []
engines = [_RecordingEngine(calls) for _ in engine_gpu_counts]
async_utils = SimpleNamespace(submit=lambda coro: coro, wait_futures=lambda futures: None)
@@ -78,7 +77,6 @@ def _connect(*, engine_gpu_counts: list[int], rollout_num_gpus_per_engine: int)
patch(f"{_BROADCAST_MODULE}.init_process_group"),
):
connect_rollout_engines_from_distributed(
Namespace(rollout_num_gpus_per_engine=rollout_num_gpus_per_engine),
"miles-pp_0",
engines,
engine_gpu_counts=engine_gpu_counts,
@@ -89,7 +87,7 @@ def _connect(*, engine_gpu_counts: list[int], rollout_num_gpus_per_engine: int)
class TestExternalPdFleetWeightUpdateLayout:
@pytest.mark.asyncio
async def test_the_discovered_gpu_counts_reach_the_update_group_unchanged(self, monkeypatch):
"""Discovered engine sizes reach the group instead of its deliberately different fallback."""
"""Discovered engine sizes reach the update group unchanged."""
urls = ["prefill:8000", "decode:8000"]
payloads = {
"http://prefill:8000": _payload(num_gpus=2, disaggregation_mode="prefill"),
@@ -109,7 +107,7 @@ class TestExternalPdFleetWeightUpdateLayout:
assert srv.api_clients == ["client-0", "client-2"]
counts = srv.engine_gpu_counts
calls = _connect(engine_gpu_counts=counts, rollout_num_gpus_per_engine=1)
calls = _connect(engine_gpu_counts=counts)
assert [call["rank"] for call in calls] == [1, 3]
assert {call["world_size"] for call in calls} == {5}
@@ -167,7 +165,7 @@ class TestExternalRegularEngineWeightUpdateLayout:
counts = srv.engine_gpu_counts
assert counts == [2]
calls = _connect(engine_gpu_counts=counts, rollout_num_gpus_per_engine=1)
calls = _connect(engine_gpu_counts=counts)
assert [call["rank"] for call in calls] == [1]
assert {call["world_size"] for call in calls} == {3}
@@ -18,7 +18,7 @@ from miles.ray.rollout.inference_controller import (
)
from miles.ray.rollout.rollout_server import RolloutServer
from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata
from miles.ray.specs.inference import compute_engine_pool_ids, compute_router_pool_id, specs_inference_engine
from miles.ray.specs.inference import InferenceEngineSpec, compute_router_pool_id
from miles.utils.context_lock import ContextLock
from miles.utils.ft_utils.health_checker import ActivenessTracker
from miles.utils.workers.registration.hub import RegistrationHub
@@ -27,7 +27,7 @@ from miles.utils.workers.rpc.client.handle import RpcWorkerHandle
from miles.utils.workers.rpc.common.metadata import collect_rpc_method_specs
from miles.utils.workers.worker_info import WorkerInfo
from miles.utils.workers.worker_provider.base import BaseWorkerProvider, CellInfo, CellReconcileFn, StopWatchFn
from miles.utils.workers.worker_spec import HostAndPort, NamedHostAndPorts, WorkerMetaContext
from miles.utils.workers.worker_spec import HostAndPort, NamedHostAndPorts
_RUN_UUID = "run-uuid-1"
@@ -248,6 +248,10 @@ class _RecordingEvalFleet:
return None
def _engine_pool_ids(args: Namespace) -> list[str]:
return [spec.name for spec in InferenceEngineSpec.create(args)]
class _FakeWorkerProvider(BaseWorkerProvider):
def __init__(self, cell_infos: list[CellInfo], *, pool_ids: list[str] | None = None) -> None:
self._cell_infos = cell_infos
@@ -281,9 +285,12 @@ class _FakeWorkerProvider(BaseWorkerProvider):
class _RecordingInferenceControllerEvalFleet:
def __init__(self, info: EvalFleetInfo):
self.info = info
self._info = info
self.pins: list[dict] = []
async def info(self) -> EvalFleetInfo:
return self._info
async def pin(self, checkpoint_dir: str, weight_version: str) -> EvalFleetPin:
self.pins.append(dict(checkpoint_dir=checkpoint_dir, weight_version=weight_version))
return EvalFleetPin(skip_reason=None)
@@ -686,7 +693,7 @@ class TestInitSubscription:
monkeypatch.setattr(inference_controller_module, "create_rollout_servers", _fake_create_rollout_servers)
monkeypatch.setattr(inference_controller_module, "resolve_router_addrs", _fake_resolve_router_addrs)
args = make_args()
provider = _OrderRecordingProvider([], pool_ids=compute_engine_pool_ids(args))
provider = _OrderRecordingProvider([], pool_ids=_engine_pool_ids(args))
await _init_controller(args, engine_provider=provider)
@@ -696,12 +703,12 @@ class TestInitSubscription:
async def test_init_watches_the_engine_provider_it_was_handed(self, monkeypatch: pytest.MonkeyPatch):
"""The pools are the provider's own, so the controller may only open a watch on what it was given."""
args = make_args()
provider = _FakeWorkerProvider([], pool_ids=compute_engine_pool_ids(args))
provider = _FakeWorkerProvider([], pool_ids=_engine_pool_ids(args))
_patch_init(monkeypatch, servers={"default": _RecordingServer()})
await _init_controller(args, engine_provider=provider)
assert provider.watched_pool_ids == compute_engine_pool_ids(args)
assert provider.watched_pool_ids == _engine_pool_ids(args)
assert compute_router_pool_id(0) not in provider.watched_pool_ids
assert "session-server" not in provider.watched_pool_ids
@@ -710,7 +717,7 @@ class TestInitSubscription:
"""A router cell carries no engine meta, so reading one as engine meta would kill startup; the
controller is safe because it subscribes to the engine pools alone."""
args = make_args()
assert compute_router_pool_id(0) not in compute_engine_pool_ids(args)
assert compute_router_pool_id(0) not in _engine_pool_ids(args)
router_info = CellInfo(
cell_id="inference-router-0-0",
@@ -720,8 +727,8 @@ class TestInitSubscription:
workers_hash="pseudo-hash-router",
meta={},
)
engine_info = _make_cell_info(model_id="default", pool_id=compute_engine_pool_ids(args)[0])
provider = _FakeWorkerProvider([router_info, engine_info], pool_ids=compute_engine_pool_ids(args))
engine_info = _make_cell_info(model_id="default", pool_id=_engine_pool_ids(args)[0])
provider = _FakeWorkerProvider([router_info, engine_info], pool_ids=_engine_pool_ids(args))
srv = _RecordingServer()
_patch_init(monkeypatch, servers={"default": srv})
@@ -743,7 +750,7 @@ class TestEngineMetaContract:
" num_gpus_per_engine: 2\n"
)
args = make_args(sglang_config=str(config_path), rollout_num_gpus=4, sglang_api_key="from-args")
(spec,) = specs_inference_engine(args)
(spec,) = InferenceEngineSpec.create(args)
info = CellInfo(
cell_id="inference-engine-0-0-1",
@@ -751,7 +758,7 @@ class TestEngineMetaContract:
alive=True,
worker_names=["inference-engine-0-0-1-0"],
workers_hash="pseudo-hash-0",
meta=spec.meta(WorkerMetaContext(cell_index=1)),
meta=spec.static_meta.resolve(cell_index=1),
)
assert _compute_server_cell_meta_from_info(info) == ServerCellMetadata(
@@ -1444,7 +1451,7 @@ class TestEvalFleetSurface:
def test_the_fleet_description_survives_the_wire(self):
"""The executor retargets its eval args to what it decodes, so every field must round-trip."""
serializer = collect_rpc_method_specs(InferenceController)["get_eval_fleet_info"].serializer
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), num_gpus=2, num_gpus_per_engine=1)
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[1, 1])
assert serializer.decode_result(serializer.encode_result(info)) == info
assert serializer.decode_result(serializer.encode_result(None)) is None
@@ -1479,7 +1486,7 @@ class TestEvalFleetSurface:
async def test_the_fleet_answers_and_pins_through_the_controller(self):
"""The fleet lives beside its engines: the executor only ever addresses it through the controller."""
controller = _make_controller({})
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), num_gpus=2, num_gpus_per_engine=1)
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[1, 1])
controller._eval_fleet = _RecordingInferenceControllerEvalFleet(info)
assert await controller.get_eval_fleet_info() == info
@@ -5,7 +5,7 @@ from pathlib import Path
import pytest
import torch
from tests.fast.ray.rollout.conftest import make_args, make_sample
from tests.fast.ray.rollout.conftest import FakeInferenceTopologyProvider, make_args, make_rollout_config, make_sample
from tests.fast.train_parallel_config_utils import make_train_parallel_config
from miles.backends.megatron_utils.ft.types import TrainStepOutcome
@@ -25,6 +25,7 @@ from miles.rollout.base_types import (
from miles.rollout.data_source import RolloutDataSource
from miles.rollout.inference_rollout import inference_rollout_common
from miles.rollout.inference_rollout.inference_rollout_common import GenerateState
from miles.utils.args.component_rollout import InferenceRuntimeImmutState, InferenceRuntimeMutState
from miles.utils.audit_utils.event_analyzer.rules.sample_ownership.models import SampleOwnershipViolation
from miles.utils.audit_utils.event_logger.logger import EventLogger, read_events, set_event_logger
from miles.utils.audit_utils.event_logger.models import (
@@ -43,6 +44,10 @@ from miles.utils.workers.worker_spec import HostAndPort
class FakeInferenceController:
def __init__(self) -> None:
self.pins: list[tuple[str, str]] = []
self.info: EvalFleetInfo | None = None
async def get_eval_fleet_info(self) -> EvalFleetInfo | None:
return self.info
async def pin_eval_fleet(self, *, checkpoint_dir: str, weight_version: str) -> EvalFleetPin:
self.pins.append((checkpoint_dir, weight_version))
@@ -189,7 +194,7 @@ class TestSetEvalFleetInfo:
provider = FakeInferenceControllerProvider(controller)
eval_function = FakeEvalFunction()
executor = RolloutExecutor.__new__(RolloutExecutor)
executor.args = Namespace(
executor.args = make_rollout_config(
chat_template_path=None,
custom_eval_rollout_log_function_path=None,
custom_generate_function_path=None,
@@ -217,11 +222,8 @@ class TestSetEvalFleetInfo:
executor.eval_generate_rollout = eval_function
executor.last_get_rollout_id_of_model_id = {None: 9}
executor._metric_checker = None
info = EvalFleetInfo(
router=HostAndPort(host="10.0.0.2", port=31000),
num_gpus=2,
num_gpus_per_engine=1,
)
info = EvalFleetInfo(router=HostAndPort(host="10.0.0.2", port=31000), engine_gpu_counts=[1, 1])
controller.info = info
await executor.set_eval_fleet_info(info)
await executor._eval_checkpoint(
@@ -243,8 +245,9 @@ class TestSetEvalFleetInfo:
assert isinstance(first.generate_state, GenerateState)
assert first.generate_state.args.sglang_router_ip == info.router.host
assert first.generate_state.args.sglang_router_port == info.router.port
assert first.generate_state.args.rollout_num_gpus == info.num_gpus
assert first.generate_state.args.rollout_num_gpus_per_engine == info.num_gpus_per_engine
assert first.generate_state.args.inference_runtime_mut_state == InferenceRuntimeMutState(
engine_count=2, gpu_count=2
)
assert second.generate_state is None
@@ -383,6 +386,9 @@ class TestOutputSnapshotReplay:
@staticmethod
def _configure_async_executor(executor: RolloutExecutor, *, args: Namespace, rollout_fn: BaseRolloutFn) -> None:
executor.args = args
executor._inference_controller_provider = FakeInferenceTopologyProvider(
InferenceRuntimeImmutState(engine_count=8, gpu_count=8)
)
executor.use_legacy_rollout_v1 = False
executor.generate_rollout = rollout_fn
executor.eval_generate_rollout = rollout_fn
@@ -1,11 +1,12 @@
from unittest.mock import MagicMock
import pytest
from tests.fast.ray.rollout.conftest import make_args
from tests.fast.ray.rollout.conftest import FakeInferenceTopologyProvider, make_args
from miles.ray.rollout import rollout_executor as executor_module
from miles.ray.rollout.rollout_executor import RolloutExecutor
from miles.rollout.session.types import SessionServerInstance
from miles.utils.args.component_rollout import InferenceRuntimeImmutState, InferenceRuntimeMutState
from miles.utils.init_once import InitOnce, InitState
pytestmark = pytest.mark.asyncio
@@ -30,7 +31,9 @@ class TestInitRunsExactlyOnce:
args=args,
router_providers=[],
session_server_provider=MagicMock(),
inference_controller_provider=MagicMock(),
inference_controller_provider=FakeInferenceTopologyProvider(
InferenceRuntimeImmutState(engine_count=1, gpu_count=1)
),
)
async def resolve_router(args, *, router_providers) -> dict:
@@ -53,6 +56,7 @@ class TestInitRunsExactlyOnce:
await executor.init()
assert args.session_server_instances == [SessionServerInstance(addr="eval-session:5000")]
assert args.inference_runtime_mut_state == InferenceRuntimeMutState(engine_count=1, gpu_count=1)
async def test_a_constructed_executor_reports_itself_uninitialized(self):
"""The constructor the run really uses is what has to leave the guard clear."""
@@ -6,12 +6,14 @@ from typing import Any
import pytest
from tests.fast.fixtures.driver_fakes import FakeObjectStore
from tests.fast.ray.rollout.conftest import FakeInferenceTopologyProvider
from tests.fast.train_parallel_config_utils import make_train_parallel_config
from miles.ray.rollout import rollout_executor as rollout_executor_module
from miles.ray.rollout.output_snapshotter import _RolloutExecutorOutputSnapshotter
from miles.ray.rollout.rollout_executor import RolloutExecutor
from miles.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainInput
from miles.utils.args.component_rollout import InferenceRuntimeImmutState, InferenceRuntimeMutState
from miles.utils.dp_schedule import TrainParallelConfig
from miles.utils.timer import Timer
from miles.utils.weight_version import (
@@ -50,8 +52,14 @@ def _make_executor() -> RolloutExecutor:
debug_train_only=False,
debug_skip_weight_update=False,
lora_rank=0,
lora_adapter_path=None,
update_weights_interval=1,
ci_inject_missing_prefetched_batch_bug=False,
starts_inference_engines=True,
inference_runtime_mut_state=InferenceRuntimeMutState(),
)
executor._inference_controller_provider = FakeInferenceTopologyProvider(
InferenceRuntimeImmutState(engine_count=2, gpu_count=2)
)
executor._output_snapshotter = _RolloutExecutorOutputSnapshotter(args=executor.args)
executor.data_source = Namespace()
+13 -24
View File
@@ -5,13 +5,10 @@ from pathlib import Path
from types import SimpleNamespace
import pytest
from tests.fast.fixtures.sglang_config_fixtures import resolve_sglang_config, resolve_sglang_config_and_scaling
from tests.fast.ray.rollout.conftest import make_args
from miles.backends.sglang_utils.sglang_config import (
_compute_megatron_num_gpus,
_compute_rollout_offset,
resolve_sglang_config,
)
from miles.backends.sglang_utils.sglang_config import _compute_megatron_num_gpus, _compute_rollout_offset
from miles.ray.rollout.cell_state import CellAddrInfo, StateServing
from miles.ray.rollout.rollout_server import RolloutServer, create_rollout_servers
from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata
@@ -31,19 +28,19 @@ class TestRolloutServerPureFunctions:
" num_gpus: 4\n"
" num_gpus_per_engine: 1\n"
)
args = make_args(sglang_config=str(cfg_path), rollout_num_gpus=8)
with pytest.raises(AssertionError, match="total GPUs"):
resolve_sglang_config(args)
resolve_sglang_config(make_args(sglang_config=str(cfg_path), rollout_num_gpus=8))
def test_eval_fleet_inherits_rollout_engine_settings(self):
"""The eval model carries only what makes it an eval fleet; the rest is inherited."""
args = make_args(eval_num_gpus=2, eval_num_gpus_per_engine=2)
config = resolve_sglang_config(args)
config, scaling = resolve_sglang_config_and_scaling(args)
[eval_model] = [m for m in config.models if m.name == "eval"]
assert eval_model.update_weights is False
[group] = eval_model.server_groups
assert (group.num_gpus, group.num_gpus_per_engine) == (2, 2)
[group_scaling] = scaling.groups["eval"]
assert (group_scaling.num_gpus, group.num_gpus_per_engine) == (2, 2)
# Eval samples never feed training, so the replay side-channels are forced off.
assert group.overrides["enable_return_routed_experts"] is False
assert group.overrides["enable_return_indexer_topk"] is False
@@ -88,10 +85,10 @@ class TestRolloutServerPureFunctions:
def test_debug_train_only_builds_only_the_eval_model(self):
args = make_args(debug_train_only=True, eval_num_gpus=8, eval_num_gpus_per_engine=1)
config = resolve_sglang_config(args)
config, scaling = resolve_sglang_config_and_scaling(args)
assert [model.name for model in config.models] == ["eval"]
assert sum(group.num_gpus for model in config.models for group in model.server_groups) == 8
assert sum(group.num_gpus for groups in scaling.groups.values() for group in groups) == 8
def test_debug_train_only_without_eval_fleet_builds_no_model(self):
args = make_args(debug_train_only=True, eval_num_gpus=0)
@@ -164,23 +161,17 @@ class TestRolloutServerPureFunctions:
)
assert _compute_rollout_offset(args) == 0
def test_compute_rollout_offset_critic_train_only(self):
args = make_args(
colocate=False,
debug_train_only=False,
debug_rollout_only=False,
critic_train_only=True,
critic_num_nodes=1,
critic_num_gpus_per_node=4,
)
assert _compute_rollout_offset(args) == 4
def test_compute_rollout_offset_refuses_the_removed_critic_train_only_field(self):
args = make_args(colocate=False, debug_train_only=False, debug_rollout_only=False)
args.critic_train_only = True
with pytest.raises(AssertionError, match="critic_train_only is not supported"):
_compute_rollout_offset(args)
def test_compute_rollout_offset_shared_actor_critic(self):
args = make_args(
colocate=False,
debug_train_only=False,
debug_rollout_only=False,
critic_train_only=False,
use_critic=True,
actor_num_nodes=1,
actor_num_gpus_per_node=8,
@@ -195,7 +186,6 @@ class TestRolloutServerPureFunctions:
actor_num_gpus_per_node=8,
use_critic=False,
debug_rollout_only=False,
critic_train_only=False,
)
assert _compute_megatron_num_gpus(args) == 16
@@ -207,7 +197,6 @@ class TestRolloutServerPureFunctions:
critic_num_nodes=1,
critic_num_gpus_per_node=4,
debug_rollout_only=False,
critic_train_only=False,
)
assert _compute_megatron_num_gpus(args) == 8
@@ -8,7 +8,7 @@ from tests.fast.ray.rollout.conftest import make_args
from miles.ray.rollout.rollout_server import RolloutServer, create_rollout_servers
from miles.ray.rollout.router_manager import resolve_router_addrs
from miles.ray.specs.inference import compute_engine_pool_id, specs_inference_engine
from miles.ray.specs.inference import InferenceEngineSpec, compute_engine_pool_id
from miles.utils.context_lock import ContextLock
from miles.utils.workers.worker_provider.base import BaseWorkerProvider
from miles.utils.workers.worker_spec import HostAndPort, NamedHostAndPorts
@@ -93,11 +93,11 @@ def _make_args_with_config(models: list[dict], tmp_path: Path) -> Namespace:
def _expected_num_cells_from_specs(args: Namespace) -> dict[int, int]:
specs_by_name = {spec.name: spec for spec in specs_inference_engine(args)}
specs_by_name = {spec.name: spec for spec in InferenceEngineSpec.create(args)}
counts: dict[int, int] = {}
for name, spec in specs_by_name.items():
model_idx = int(name.split("-")[-2])
counts[model_idx] = counts.get(model_idx, 0) + spec.scheduling.num_cells
counts[model_idx] = counts.get(model_idx, 0) + spec.scheduling(args).num_cells
return counts

Some files were not shown because too many files have changed in this diff Show More