Absorb the orchestration scripts' non-training machinery (#3027)

This commit is contained in:
fzyzcjy
2026-09-26 20:21:13 +08:00
committed by GitHub
parent 5d9b53fca2
commit 95cf1849ef
10 changed files with 117 additions and 237 deletions
+19
View File
@@ -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
+4 -3
View File
@@ -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",
+1 -5
View File
@@ -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)
+14 -7
View File
@@ -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"]
+3 -186
View File
@@ -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",
+3 -12
View File
@@ -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
View File
@@ -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
View File
@@ -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():