mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Absorb the orchestration scripts' non-training machinery (#3027)
This commit is contained in:
@@ -0,0 +1,19 @@
|
||||
from argparse import Namespace
|
||||
|
||||
from ray.actor import ActorHandle
|
||||
|
||||
from miles.ray.wiring import launch_worker_manager
|
||||
from miles.utils import object_store
|
||||
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
|
||||
from miles.utils.debug_utils.periodic_py_spy import maybe_start_periodic_pyspy_dump
|
||||
from miles.utils.logging_utils import configure_logger
|
||||
from miles.utils.tracking_utils.tracking import init_tracking
|
||||
|
||||
|
||||
def init_orchestration_script(args: Namespace) -> ActorHandle | None:
|
||||
configure_logger(args, source=SimpleProcessIdentity(component="main"))
|
||||
maybe_start_periodic_pyspy_dump()
|
||||
init_tracking(args)
|
||||
worker_manager = launch_worker_manager(args)
|
||||
object_store.init_instance(args, contribute_segment=False)
|
||||
return worker_manager
|
||||
@@ -24,6 +24,7 @@ TRAIN_ONLY_SUBCOMMAND = "train"
|
||||
ORCHESTRATION_SCRIPTS = ("train.py", "train_async.py", "train_multi_lora_async.py")
|
||||
|
||||
BACKEND_CAPABILITY_FN = "create_backend_capability"
|
||||
ORCHESTRATION_INIT_FN = "init_orchestration_script"
|
||||
|
||||
UPPER_LAYER_MODULES = (
|
||||
"kubernetes",
|
||||
@@ -50,9 +51,9 @@ UPPER_LAYER_NAMES = (
|
||||
UPPER_LAYER_EXEMPTIONS = {
|
||||
"miles/ray/specs": "the composition root of a worker process: a spec says what its worker is built from",
|
||||
"miles/ray/wiring.py": "the glue layer holding the driver process's single fork between the backends",
|
||||
"train.py": "orchestration script: its first lines are the driver process's composition root",
|
||||
"train_async.py": "orchestration script: its first lines are the driver process's composition root",
|
||||
"train_multi_lora_async.py": "orchestration script: its first lines are the driver process's composition root",
|
||||
"miles/utils/orchestration_utils.py": (
|
||||
"the shared driver composition root that launches the selected worker manager"
|
||||
),
|
||||
"miles/utils/workers/worker_provider": "the infrastructure that owns every provider implementation",
|
||||
"miles/utils/workers/serving/serve_inner.py": "the composition root of a served worker process",
|
||||
"miles/utils/workers/ray_worker_manager.py": "the composition root of a worker process an actor wraps",
|
||||
|
||||
@@ -67,11 +67,7 @@ def _install_driver_fakes(
|
||||
) -> None:
|
||||
events.append(f"update_weights:{rollout_id}")
|
||||
|
||||
monkeypatch.setattr(train_driver, "configure_logger", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(train_driver, "maybe_start_periodic_pyspy_dump", lambda: None)
|
||||
monkeypatch.setattr(train_driver, "launch_worker_manager", lambda _args: None)
|
||||
monkeypatch.setattr(train_driver.object_store, "init_instance", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(train_driver, "init_tracking", lambda _args: None)
|
||||
monkeypatch.setattr(train_driver, "init_orchestration_script", lambda _args: None)
|
||||
monkeypatch.setattr(train_driver, "create_rollout_components", create_rollout_components)
|
||||
monkeypatch.setattr(train_driver, "create_training_models", create_training_models)
|
||||
monkeypatch.setattr(train_driver, "maybe_start_mini_ft_controller", lambda _args: None)
|
||||
|
||||
@@ -12,8 +12,11 @@ from tests.fast.fixtures.driver_fakes import (
|
||||
)
|
||||
|
||||
from miles.backends.megatron_utils.ft.types import TrainStepOutcome, TrainStepOutput
|
||||
from miles.ray import placement_group as placement_group_mod
|
||||
from miles.utils import object_store
|
||||
|
||||
from miles.ray import placement_group as placement_group_mod
|
||||
|
||||
|
||||
def _make_args(**overrides: Any) -> SimpleNamespace:
|
||||
args = SimpleNamespace(
|
||||
@@ -61,6 +64,7 @@ def _install_driver_fakes(
|
||||
actor_model=FakeTrainingModel(events, "actor"),
|
||||
critic_model=FakeTrainingModel(events, "critic") if args.use_critic else None,
|
||||
api_server_calls=[],
|
||||
cell_operations=object(),
|
||||
)
|
||||
|
||||
async def create_rollout_components(_args: SimpleNamespace) -> tuple[Any, Any, int]:
|
||||
@@ -74,18 +78,20 @@ def _install_driver_fakes(
|
||||
) -> None:
|
||||
events.append(f"update_weights:{rollout_id}")
|
||||
|
||||
monkeypatch.setattr(train_async_driver, "configure_logger", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(train_async_driver, "maybe_start_periodic_pyspy_dump", lambda: None)
|
||||
monkeypatch.setattr(train_async_driver, "launch_worker_manager", lambda _args: None)
|
||||
monkeypatch.setattr(train_async_driver.object_store, "init_instance", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(train_async_driver, "init_tracking", lambda _args: None)
|
||||
monkeypatch.setattr(train_async_driver, "init_orchestration_script", lambda _args: None)
|
||||
monkeypatch.setattr(train_async_driver, "create_rollout_components", create_rollout_components)
|
||||
monkeypatch.setattr(train_async_driver, "create_training_models", create_training_models)
|
||||
monkeypatch.setattr(train_async_driver, "maybe_start_mini_ft_controller", lambda _args: None)
|
||||
monkeypatch.setattr(train_async_driver, "update_weights", update_weights)
|
||||
monkeypatch.setattr(train_async_driver, "remove_rollout_data_refs", lambda *_args, **_kwargs: None)
|
||||
# the driver reaches the server through maybe_start_api_server, whose gate the tests exercise
|
||||
monkeypatch.setattr(
|
||||
train_async_driver, "start_api_server", lambda **kwargs: components.api_server_calls.append(kwargs)
|
||||
placement_group_mod, "start_api_server", lambda **kwargs: components.api_server_calls.append(kwargs)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
placement_group_mod,
|
||||
"get_backend_capability",
|
||||
lambda _args: SimpleNamespace(cell_operations=lambda: components.cell_operations),
|
||||
)
|
||||
return components
|
||||
|
||||
@@ -100,7 +106,8 @@ class TestApiServer:
|
||||
await train_async_driver.train(args)
|
||||
|
||||
(call,) = components.api_server_calls
|
||||
assert call["actor_model"] is components.actor_model
|
||||
assert list(call["trainer_models"]) == ["actor"]
|
||||
assert call["trainer_models"]["actor"] is components.actor_model
|
||||
assert call["inference_controller"] is components.inference_controller
|
||||
assert call["port"] == 8123
|
||||
assert call["ft_components"] == ["rollout"]
|
||||
|
||||
@@ -4,10 +4,8 @@ register_cpu_ci(est_time=30, suite="stage-a-cpu", labels=[])
|
||||
|
||||
import asyncio
|
||||
from argparse import Namespace
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import train_multi_policy as multi_policy_driver
|
||||
@@ -15,9 +13,7 @@ from tests.fast.fixtures.args_fixtures import parser_defaults
|
||||
from tests.fast.fixtures.megatron_config_fixtures import encode_megatron_config
|
||||
from train_multi_policy import train_multi_policy
|
||||
|
||||
from miles.ray import placement_group
|
||||
from miles.utils.multi_policy.checkpoint_state import MultiPolicyCheckpointState
|
||||
from miles.utils.multi_policy.parker import Parker
|
||||
from miles.utils.multi_policy.utils import TrainerInfo
|
||||
|
||||
|
||||
@@ -42,8 +38,6 @@ def _make_args(**overrides: Any) -> Namespace:
|
||||
def _make_trainers(model_ids, handles=None, start_rollout_ids=None) -> dict[str, TrainerInfo]:
|
||||
handles = {model_id: AsyncMock() for model_id in model_ids} if handles is None else handles
|
||||
start_rollout_ids = start_rollout_ids or {}
|
||||
for handle in handles.values():
|
||||
_let_a_follower_yield(handle)
|
||||
return {
|
||||
model_id: TrainerInfo(model_id=model_id, start_rollout_id=start_rollout_ids.get(model_id, 0), handle=handle)
|
||||
for model_id, handle in handles.items()
|
||||
@@ -79,17 +73,14 @@ async def _run(
|
||||
def _stub_driver_environment(monkeypatch):
|
||||
"""Everything the driver reaches outside its own loop: cluster, tracking and logging."""
|
||||
for name in (
|
||||
"configure_logger",
|
||||
"maybe_start_periodic_pyspy_dump",
|
||||
"init_tracking",
|
||||
"init_orchestration_script",
|
||||
"define_policy_metric_groups",
|
||||
"launch_worker_manager",
|
||||
"maybe_start_api_server",
|
||||
"maybe_start_mini_ft_controller",
|
||||
"validate_multi_policy_args",
|
||||
"assert_consistent_restore",
|
||||
):
|
||||
monkeypatch.setattr(multi_policy_driver, name, lambda *a, **kw: None)
|
||||
monkeypatch.setattr(multi_policy_driver.object_store, "init_instance", lambda *a, **kw: None)
|
||||
monkeypatch.setattr(multi_policy_driver, "create_trainers", AsyncMock(return_value={}))
|
||||
monkeypatch.setattr(multi_policy_driver, "create_rollout_components", AsyncMock())
|
||||
|
||||
@@ -104,42 +95,7 @@ def _stub_update_weights(monkeypatch):
|
||||
monkeypatch.setattr(multi_policy_driver, "update_weights", AsyncMock())
|
||||
|
||||
|
||||
async def _slow_train(rollout_id: int, rollout_data_ref, **kwargs) -> None:
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
|
||||
async def _train_that_never_returns(rollout_id: int, rollout_data_ref: Any, **kwargs: Any) -> None:
|
||||
await asyncio.Event().wait()
|
||||
|
||||
|
||||
def _let_a_follower_yield(handle) -> None:
|
||||
async def yield_to_the_leader(rollout_id: int, rollout_data_ref, **kwargs) -> None:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
if isinstance(handle.train, AsyncMock) and handle.train.side_effect is None:
|
||||
handle.train.side_effect = yield_to_the_leader
|
||||
|
||||
|
||||
class TestInitialWeightPublication:
|
||||
async def test_api_server_receives_every_configured_trainer_handle(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Each configured trainer id must expose its model's live handle through the shared API server."""
|
||||
start_api_server = Mock()
|
||||
monkeypatch.setattr(placement_group, "start_api_server", start_api_server)
|
||||
monkeypatch.setattr(
|
||||
placement_group,
|
||||
"get_backend_capability",
|
||||
lambda _args: SimpleNamespace(cell_operations=lambda: object()),
|
||||
)
|
||||
|
||||
context = await _run(_make_args(num_rollout=0, api_server_port=18080))
|
||||
|
||||
start_api_server.assert_called_once()
|
||||
assert start_api_server.call_args.kwargs["trainer_models"] == {
|
||||
"a-actor": context["trainers"]["a"],
|
||||
"b-actor": context["trainers"]["b"],
|
||||
}
|
||||
assert start_api_server.call_args.kwargs["inference_controller"] is context["inference_controller"]
|
||||
|
||||
async def test_every_policy_compares_its_engines_against_its_own_trainer(self):
|
||||
"""--ci-test asks for this comparison, and running it for one policy would leave the others unchecked."""
|
||||
context = await _run(_make_args(num_rollout=0, check_weight_update_equal=True))
|
||||
@@ -147,17 +103,6 @@ class TestInitialWeightPublication:
|
||||
compared = [call.kwargs["model_id"] for call in context["inference_controller"].check_weights.await_args_list]
|
||||
assert sorted(compared) == ["a", "b"]
|
||||
|
||||
async def test_each_policy_stamps_its_startup_sync_with_its_own_restore_point(self):
|
||||
"""The policies resume at their own rollouts, so one global start id would misattribute their weights."""
|
||||
await _run(_make_args(num_rollout=0), start_rollout_ids=dict(a=3, b=7))
|
||||
|
||||
stamped = {
|
||||
call.kwargs["trainer_model_id"]: call.args[0].start_rollout_id
|
||||
for call in multi_policy_driver.update_weights.await_args_list
|
||||
if "rollout_id" not in call.kwargs
|
||||
}
|
||||
assert stamped == dict(a=3, b=7)
|
||||
|
||||
async def test_a_run_that_does_not_ask_for_the_comparison_does_not_pay_for_it(self):
|
||||
"""The comparison walks every parameter, so it stays off unless the run turns it on."""
|
||||
context = await _run(_make_args(num_rollout=0))
|
||||
@@ -298,43 +243,6 @@ class TestRunPolicies:
|
||||
|
||||
assert len(trainers["b"].train.await_args_list) >= 2
|
||||
|
||||
async def test_a_debug_run_stops_the_leader_after_its_own_rounds(self):
|
||||
"""--debug-exit-after-rollout counts from where the policy resumed, not from rollout zero."""
|
||||
trainers = {"a": AsyncMock(), "b": AsyncMock()}
|
||||
|
||||
await _run(
|
||||
_make_args(num_rollout=10, debug_exit_after_rollout=1), trainers=trainers, start_rollout_ids=dict(a=0, b=5)
|
||||
)
|
||||
|
||||
assert [call.args[0] for call in trainers["a"].train.await_args_list] == [0]
|
||||
|
||||
async def test_a_debug_run_leaves_the_followers_running(self):
|
||||
"""A follower's rounds are the leader's to end, so honouring the flag there would retire it
|
||||
from the checkpoint every remaining round waits at."""
|
||||
trainers = {"a": AsyncMock(), "b": AsyncMock()}
|
||||
trainers["a"].train = _slow_train
|
||||
|
||||
await _run(_make_args(num_rollout=10, debug_exit_after_rollout=1), trainers=trainers)
|
||||
|
||||
assert len(trainers["b"].train.await_args_list) >= 2
|
||||
|
||||
async def test_the_run_ends_when_the_leader_runs_out_of_rounds(self):
|
||||
"""The leader owns --num-rollout; a follower resuming further back must not extend the run."""
|
||||
trainers = {"a": AsyncMock(), "b": AsyncMock()}
|
||||
|
||||
await _run(_make_args(num_rollout=2), trainers=trainers, start_rollout_ids=dict(a=0, b=0))
|
||||
|
||||
assert [call.args[0] for call in trainers["a"].train.await_args_list] == [0, 1]
|
||||
|
||||
async def test_a_follower_is_never_the_one_that_ends_the_run(self):
|
||||
"""Followers train unbounded rounds, so the run must not stop because one of them reached num_rollout."""
|
||||
trainers = {"a": AsyncMock(), "b": AsyncMock()}
|
||||
trainers["a"].train = _slow_train
|
||||
|
||||
await _run(_make_args(num_rollout=2), trainers=trainers)
|
||||
|
||||
assert len(trainers["b"].train.await_args_list) >= 2
|
||||
|
||||
|
||||
class TestSaving:
|
||||
async def test_the_leader_parks_everybody_and_records_where_they_stood(self, tmp_path):
|
||||
@@ -434,94 +342,3 @@ class TestSaving:
|
||||
trainers["a"].save_model.assert_not_awaited()
|
||||
trainers["b"].save_model.assert_not_awaited()
|
||||
context["rollout_executor"].save.assert_not_awaited()
|
||||
|
||||
|
||||
class TestARunThatCancellationCannotEnd:
|
||||
async def test_a_follower_that_absorbs_cancellation_still_stops(self):
|
||||
"""A follower loops without bound, so ending the run must not depend on cancellation reaching it."""
|
||||
absorbed = False
|
||||
|
||||
async def _train(rollout_id: int, rollout_data_ref: Any, **kwargs: Any) -> None:
|
||||
nonlocal absorbed
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
except asyncio.CancelledError:
|
||||
if absorbed:
|
||||
raise
|
||||
absorbed = True
|
||||
|
||||
async def _get(rollout_id: int, trainer_model_id: str | None = None) -> dict:
|
||||
await asyncio.sleep(0.01)
|
||||
return dict(data_ref=None)
|
||||
|
||||
trainers = {"a": AsyncMock(), "b": AsyncMock()}
|
||||
trainers["b"].train = _train
|
||||
rollout_executor = AsyncMock()
|
||||
rollout_executor.get = _get
|
||||
|
||||
await asyncio.wait_for(
|
||||
_run(_make_args(num_rollout=1), trainers=trainers, rollout_executor=rollout_executor), timeout=10
|
||||
)
|
||||
|
||||
assert absorbed
|
||||
|
||||
|
||||
class TestEvalDispatch:
|
||||
async def test_eval_fires_once_per_point_however_many_policies_train(self):
|
||||
"""One shared-engine eval scores every policy through the generate chain; a per-policy dispatch would double it."""
|
||||
context = await _run(_make_args(num_rollout=2, eval_interval=2))
|
||||
|
||||
eval_rollout_ids = [call.args[0] for call in context["rollout_executor"].eval.await_args_list]
|
||||
assert eval_rollout_ids == [0, 1]
|
||||
assert context["inference_controller"].prepare_eval.await_count == 2
|
||||
|
||||
async def test_a_run_without_an_eval_interval_never_evaluates(self):
|
||||
"""--eval-interval is the only opt-in; a surprise eval pauses production for the whole test split."""
|
||||
context = await _run(_make_args(num_rollout=2))
|
||||
|
||||
context["rollout_executor"].eval.assert_not_awaited()
|
||||
context["inference_controller"].prepare_eval.assert_not_awaited()
|
||||
|
||||
async def test_skip_eval_before_train_drops_only_the_starting_point(self):
|
||||
"""The flag exists to skip the expensive untrained point, not to turn eval off."""
|
||||
context = await _run(_make_args(num_rollout=2, eval_interval=2, skip_eval_before_train=True))
|
||||
|
||||
eval_rollout_ids = [call.args[0] for call in context["rollout_executor"].eval.await_args_list]
|
||||
assert eval_rollout_ids == [1]
|
||||
|
||||
async def test_eval_holds_every_follower_parked(self, monkeypatch):
|
||||
"""A follower pushing weights mid-eval would swap its engines' weights under the running sweep."""
|
||||
held_during_eval = []
|
||||
|
||||
class SpyParker(Parker):
|
||||
holding = False
|
||||
|
||||
@asynccontextmanager
|
||||
async def with_all_parked(self):
|
||||
async with super().with_all_parked():
|
||||
SpyParker.holding = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
SpyParker.holding = False
|
||||
|
||||
async def _eval(rollout_id: int) -> None:
|
||||
held_during_eval.append(SpyParker.holding)
|
||||
|
||||
rollout_executor = AsyncMock()
|
||||
rollout_executor.eval = AsyncMock(side_effect=_eval)
|
||||
monkeypatch.setattr(multi_policy_driver, "Parker", SpyParker)
|
||||
|
||||
await _run(
|
||||
_make_args(num_rollout=2, eval_interval=2, skip_eval_before_train=True),
|
||||
rollout_executor=rollout_executor,
|
||||
)
|
||||
|
||||
assert held_during_eval == [True]
|
||||
|
||||
async def test_a_resumed_run_does_not_re_evaluate_the_untrained_model(self):
|
||||
"""The rollout-0 point describes the base checkpoint; a resume from rollout 1 is past it."""
|
||||
context = await _run(_make_args(num_rollout=2, eval_interval=2), start_rollout_ids={"a": 1, "b": 1})
|
||||
|
||||
eval_rollout_ids = [call.args[0] for call in context["rollout_executor"].eval.await_args_list]
|
||||
assert eval_rollout_ids == [1]
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
from argparse import Namespace
|
||||
|
||||
import pytest
|
||||
|
||||
import miles.utils.orchestration_utils as orchestration_utils
|
||||
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
|
||||
|
||||
|
||||
class TestInitOrchestrationScript:
|
||||
def test_initializes_the_shared_driver_machinery_in_order(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Every driver must initialize the shared machinery once, in dependency order, and return its manager."""
|
||||
args = Namespace(run="test")
|
||||
worker_manager = object()
|
||||
calls: list[str] = []
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_configure_logger(actual_args: Namespace, *, source: SimpleProcessIdentity) -> None:
|
||||
calls.append("configure_logger")
|
||||
captured["logger_args"] = actual_args
|
||||
captured["source"] = source
|
||||
|
||||
def fake_maybe_start_periodic_pyspy_dump() -> None:
|
||||
calls.append("maybe_start_periodic_pyspy_dump")
|
||||
|
||||
def fake_init_tracking(actual_args: Namespace) -> None:
|
||||
calls.append("init_tracking")
|
||||
captured["tracking_args"] = actual_args
|
||||
|
||||
def fake_launch_worker_manager(actual_args: Namespace) -> object:
|
||||
calls.append("launch_worker_manager")
|
||||
captured["worker_manager_args"] = actual_args
|
||||
return worker_manager
|
||||
|
||||
def fake_init_object_store(actual_args: Namespace, *, contribute_segment: bool) -> None:
|
||||
calls.append("object_store.init_instance")
|
||||
captured["object_store_args"] = actual_args
|
||||
captured["contribute_segment"] = contribute_segment
|
||||
|
||||
monkeypatch.setattr(orchestration_utils, "configure_logger", fake_configure_logger)
|
||||
monkeypatch.setattr(
|
||||
orchestration_utils,
|
||||
"maybe_start_periodic_pyspy_dump",
|
||||
fake_maybe_start_periodic_pyspy_dump,
|
||||
)
|
||||
monkeypatch.setattr(orchestration_utils, "init_tracking", fake_init_tracking)
|
||||
monkeypatch.setattr(orchestration_utils, "launch_worker_manager", fake_launch_worker_manager)
|
||||
monkeypatch.setattr(orchestration_utils.object_store, "init_instance", fake_init_object_store)
|
||||
|
||||
result = orchestration_utils.init_orchestration_script(args)
|
||||
|
||||
assert calls == [
|
||||
"configure_logger",
|
||||
"maybe_start_periodic_pyspy_dump",
|
||||
"init_tracking",
|
||||
"launch_worker_manager",
|
||||
"object_store.init_instance",
|
||||
]
|
||||
assert captured["logger_args"] is args
|
||||
assert captured["source"] == SimpleProcessIdentity(component="main")
|
||||
assert captured["tracking_args"] is args
|
||||
assert captured["worker_manager_args"] is args
|
||||
assert captured["object_store_args"] is args
|
||||
assert captured["contribute_segment"] is False
|
||||
assert result is worker_manager
|
||||
@@ -18,6 +18,9 @@ EXCLUDED_DIRS = (REPO_ROOT / "tests",)
|
||||
RAY_USING_MODULES = {
|
||||
"miles/ray/placement_group.py": "launcher closure: placement groups are how ray is asked to schedule",
|
||||
"miles/utils/ray_utils.py": "launcher closure: node lookup and pinning options for the launcher's own calls",
|
||||
"miles/utils/orchestration_utils.py": (
|
||||
"launcher closure: the shared driver composition root returns its ray worker manager"
|
||||
),
|
||||
"miles/utils/workers/ray_worker_manager.py": "launcher closure: it is the launcher",
|
||||
"miles/utils/workers/ray_worker_handle.py": "launcher closure: the handle of the ray communication mode itself",
|
||||
"miles/utils/workers/worker_provider/ray.py": "launcher closure: it reads the launcher's own bookkeeping",
|
||||
|
||||
@@ -11,28 +11,20 @@ from miles.ray.placement_group import (
|
||||
update_weights,
|
||||
)
|
||||
from miles.ray.rollout.eval_dispatch import EvalDispatcher
|
||||
from miles.ray.wiring import launch_worker_manager
|
||||
from miles.utils import object_store
|
||||
from miles.utils.arguments import parse_args
|
||||
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
|
||||
from miles.utils.data import remove_rollout_data_refs, remove_train_output_refs
|
||||
from miles.utils.debug_utils.periodic_py_spy import maybe_start_periodic_pyspy_dump
|
||||
from miles.utils.ft_utils.mini_ft_controller import maybe_start_mini_ft_controller
|
||||
from miles.utils.logging_utils import configure_logger
|
||||
from miles.utils.lora.utils import lora_rollout_enabled
|
||||
from miles.utils.misc import should_run_periodic_action
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
|
||||
from miles.utils.orchestration_utils import init_orchestration_script
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def train(args):
|
||||
assert not args.fully_async, "--fully-async requires the async driver: run train_async.py"
|
||||
configure_logger(args, source=SimpleProcessIdentity(component="main"))
|
||||
maybe_start_periodic_pyspy_dump()
|
||||
init_tracking(args)
|
||||
_worker_manager = launch_worker_manager(args)
|
||||
object_store.init_instance(args, contribute_segment=False)
|
||||
_worker_manager = init_orchestration_script(args)
|
||||
|
||||
if args.colocate_memory_peak_device == "gpu":
|
||||
assert (
|
||||
@@ -48,7 +40,6 @@ async def train(args):
|
||||
actor_model, critic_model = await create_training_models(args, rollout_executor)
|
||||
|
||||
maybe_start_api_server(args, trainer_models={"actor": actor_model}, inference_controller=inference_controller)
|
||||
|
||||
maybe_start_mini_ft_controller(args)
|
||||
|
||||
# always update weight first so that sglang has the loaded weights from training.
|
||||
|
||||
+3
-12
@@ -9,17 +9,13 @@ from miles.ray.placement_group import (
|
||||
update_weights,
|
||||
)
|
||||
from miles.ray.rollout.eval_dispatch import EvalDispatcher
|
||||
from miles.ray.wiring import launch_worker_manager
|
||||
from miles.utils import object_store
|
||||
from miles.utils.arguments import parse_args, validate_async_off_policy_correction
|
||||
from miles.utils.async_utils import eager_create_task
|
||||
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
|
||||
from miles.utils.data import remove_rollout_data_refs, remove_train_output_refs
|
||||
from miles.utils.debug_utils.periodic_py_spy import maybe_start_periodic_pyspy_dump
|
||||
from miles.utils.ft_utils.mini_ft_controller import maybe_start_mini_ft_controller
|
||||
from miles.utils.logging_utils import configure_logger
|
||||
from miles.utils.misc import should_run_periodic_action
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
|
||||
from miles.utils.orchestration_utils import init_orchestration_script
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -28,11 +24,7 @@ logger = logging.getLogger(__name__)
|
||||
async def train(args):
|
||||
assert not args.colocate, "Colocation is not supported for async training."
|
||||
validate_async_off_policy_correction(args)
|
||||
configure_logger(args, source=SimpleProcessIdentity(component="main"))
|
||||
maybe_start_periodic_pyspy_dump()
|
||||
init_tracking(args)
|
||||
_worker_manager = launch_worker_manager(args)
|
||||
object_store.init_instance(args, contribute_segment=False)
|
||||
_worker_manager = init_orchestration_script(args)
|
||||
|
||||
# create the rollout manager, with sglang engines inside.
|
||||
# need to initialize rollout manager first to calculate num_rollout
|
||||
@@ -42,7 +34,6 @@ async def train(args):
|
||||
actor_model, critic_model = await create_training_models(args, rollout_executor)
|
||||
|
||||
maybe_start_api_server(args, trainer_models={"actor": actor_model}, inference_controller=inference_controller)
|
||||
|
||||
maybe_start_mini_ft_controller(args)
|
||||
|
||||
# always update weight first so that sglang has the loaded weights from training.
|
||||
|
||||
+3
-12
@@ -7,15 +7,10 @@ from pathlib import Path
|
||||
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
|
||||
from miles.ray.placement_group import create_rollout_components, maybe_start_api_server, update_weights
|
||||
from miles.ray.specs.train import compute_trainer_configs
|
||||
from miles.ray.wiring import launch_worker_manager
|
||||
from miles.utils import object_store
|
||||
from miles.utils.arguments import parse_args
|
||||
from miles.utils.async_utils import wait_cancelling_pending_on_first_completion
|
||||
from miles.utils.audit_utils.process_identity import SimpleProcessIdentity
|
||||
from miles.utils.data import remove_rollout_data_refs
|
||||
from miles.utils.debug_utils.periodic_py_spy import maybe_start_periodic_pyspy_dump
|
||||
from miles.utils.ft_utils.mini_ft_controller import maybe_start_mini_ft_controller
|
||||
from miles.utils.logging_utils import configure_logger
|
||||
from miles.utils.misc import should_run_periodic_action
|
||||
from miles.utils.multi_policy.checkpoint_state import MultiPolicyCheckpointState
|
||||
from miles.utils.multi_policy.parker import Parker
|
||||
@@ -26,7 +21,8 @@ from miles.utils.multi_policy.utils import (
|
||||
define_policy_metric_groups,
|
||||
validate_multi_policy_args,
|
||||
)
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking, init_tracking
|
||||
from miles.utils.orchestration_utils import init_orchestration_script
|
||||
from miles.utils.tracking_utils.tracking import finish_tracking
|
||||
from miles.utils.workers.worker_handle import BaseWorkerHandle
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -35,12 +31,8 @@ logger = logging.getLogger(__name__)
|
||||
async def train_multi_policy(args) -> None:
|
||||
megatron_config = resolve_megatron_config(args)
|
||||
validate_multi_policy_args(args, megatron_config=megatron_config)
|
||||
configure_logger(args, source=SimpleProcessIdentity(component="main"))
|
||||
maybe_start_periodic_pyspy_dump()
|
||||
init_tracking(args)
|
||||
_worker_manager = init_orchestration_script(args)
|
||||
define_policy_metric_groups(megatron_config)
|
||||
_worker_manager = launch_worker_manager(args)
|
||||
object_store.init_instance(args, contribute_segment=False)
|
||||
|
||||
inference_controller, rollout_executor, num_rollout_per_epoch = await create_rollout_components(args)
|
||||
|
||||
@@ -55,7 +47,6 @@ async def train_multi_policy(args) -> None:
|
||||
},
|
||||
inference_controller=inference_controller,
|
||||
)
|
||||
|
||||
maybe_start_mini_ft_controller(args)
|
||||
|
||||
for model_id, trainer in trainers.items():
|
||||
|
||||
Reference in New Issue
Block a user