mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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'")
|
||||
|
||||
@@ -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}"
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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]),
|
||||
|
||||
+3
-2
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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))}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()}
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user