mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Build the trainer handles once, before anything drives them (#2985)
This commit is contained in:
@@ -7,7 +7,7 @@ import ray
|
||||
from ray.util.placement_group import PlacementGroup, placement_group
|
||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||
|
||||
from miles.backends.megatron_utils.megatron_config import compute_trainer_args
|
||||
from miles.backends.megatron_utils.megatron_config import MegatronTrainerConfig, compute_trainer_args
|
||||
from miles.ray.rollout.inference_controller import UpdatableEngines
|
||||
from miles.ray.rollout.router_manager import resolve_router_addrs, wait_session_server_ready
|
||||
from miles.ray.specs.inference import (
|
||||
@@ -21,7 +21,6 @@ from miles.ray.specs.train import (
|
||||
CRITIC_ROLE,
|
||||
TRAINER_CONTROLLER_ADDRS_FLAG,
|
||||
compute_trainer_configs,
|
||||
compute_trainer_ids,
|
||||
create_trainer_controller_handle,
|
||||
external_trainer_controller_addrs,
|
||||
)
|
||||
@@ -162,8 +161,16 @@ class TrainerInfo(NamedTuple):
|
||||
|
||||
|
||||
# TODO: move (when reorganizing files)
|
||||
async def create_training_model(args, *, trainer_id: str) -> TrainerInfo:
|
||||
handle = create_trainer_controller_handle(args, capability=get_backend_capability(args), trainer_id=trainer_id)
|
||||
def create_trainer_handles(args, *, trainer_configs: list[MegatronTrainerConfig]) -> dict[str, BaseWorkerHandle]:
|
||||
capability = get_backend_capability(args)
|
||||
return {
|
||||
config.trainer_id: create_trainer_controller_handle(args, capability=capability, trainer_id=config.trainer_id)
|
||||
for config in trainer_configs
|
||||
}
|
||||
|
||||
|
||||
# TODO: move (when reorganizing files)
|
||||
async def create_training_model(args, *, handle: BaseWorkerHandle, trainer_id: str) -> TrainerInfo:
|
||||
restored_rollout_ids = await handle.init(args)
|
||||
assert len(set(restored_rollout_ids)) == 1, f"trainer {trainer_id!r} restored {restored_rollout_ids}"
|
||||
[restored_rollout_id] = set(restored_rollout_ids)
|
||||
@@ -185,12 +192,15 @@ async def create_training_model(args, *, trainer_id: str) -> TrainerInfo:
|
||||
async def create_training_models(
|
||||
args, rollout_executor: BaseWorkerHandle
|
||||
) -> tuple[BaseWorkerHandle, BaseWorkerHandle | None]:
|
||||
await wait_external_trainers(args)
|
||||
|
||||
trainer_configs = compute_trainer_configs(args)
|
||||
handles = create_trainer_handles(args, trainer_configs=trainer_configs)
|
||||
await wait_external_trainers(args, handles=handles)
|
||||
|
||||
[actor_config] = [config for config in trainer_configs if config.role == ACTOR_ROLE]
|
||||
actor_info = await create_training_model(
|
||||
compute_trainer_args(args, actor_config), trainer_id=actor_config.trainer_id
|
||||
compute_trainer_args(args, actor_config),
|
||||
handle=handles[actor_config.trainer_id],
|
||||
trainer_id=actor_config.trainer_id,
|
||||
)
|
||||
|
||||
critic_configs = [config for config in trainer_configs if config.role == CRITIC_ROLE]
|
||||
@@ -198,7 +208,9 @@ async def create_training_models(
|
||||
if args.use_critic:
|
||||
[critic_config] = critic_configs
|
||||
critic_info = await create_training_model(
|
||||
compute_trainer_args(args, critic_config), trainer_id=critic_config.trainer_id
|
||||
compute_trainer_args(args, critic_config),
|
||||
handle=handles[critic_config.trainer_id],
|
||||
trainer_id=critic_config.trainer_id,
|
||||
)
|
||||
assert critic_info.restored_rollout_id == actor_info.restored_rollout_id, (
|
||||
f"the actor restored to rollout {actor_info.restored_rollout_id} but its critic to "
|
||||
@@ -218,23 +230,17 @@ async def create_training_models(
|
||||
|
||||
|
||||
# TODO: move (when reorganizing files)
|
||||
async def wait_external_trainers(args) -> None:
|
||||
async def wait_external_trainers(args, *, handles: dict[str, BaseWorkerHandle]) -> None:
|
||||
"""Wait for every independently deployed trainer controller, and refuse one that another run deployed."""
|
||||
if args.trainer_controller_addrs is None:
|
||||
return
|
||||
|
||||
trainer_ids = compute_trainer_ids(args)
|
||||
addrs = external_trainer_controller_addrs(args, trainer_ids=trainer_ids)
|
||||
addrs = external_trainer_controller_addrs(args, trainer_ids=list(handles))
|
||||
logger.info(f"Waiting for the independently deployed trainer controllers at {addrs}")
|
||||
await wait_static_addrs_ready(addrs.values())
|
||||
|
||||
capability = get_backend_capability(args)
|
||||
handles = [
|
||||
create_trainer_controller_handle(args, capability=capability, trainer_id=trainer_id)
|
||||
for trainer_id in trainer_ids
|
||||
]
|
||||
identities = await asyncio.gather(*[handle.get_deployment_identity() for handle in handles])
|
||||
for trainer_id, identity in zip(trainer_ids, identities, strict=True):
|
||||
identities = await asyncio.gather(*[handle.get_deployment_identity() for handle in handles.values()])
|
||||
for trainer_id, identity in zip(handles, identities, strict=True):
|
||||
_assert_external_trainer_in_run(identity, args=args, trainer_id=trainer_id)
|
||||
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
||||
|
||||
from miles.backends.megatron_utils.megatron_config import MegatronConfig, compute_trainer_args, resolve_megatron_config
|
||||
from miles.backends.sglang_utils.sglang_config import resolve_sglang_config
|
||||
from miles.ray.placement_group import create_training_model, wait_external_trainers
|
||||
from miles.ray.placement_group import create_trainer_handles, create_training_model, wait_external_trainers
|
||||
from miles.ray.specs.train import compute_trainer_configs
|
||||
from miles.utils.arguments import validate_async_off_policy_correction
|
||||
from miles.utils.multi_policy.checkpoint_state import MultiPolicyCheckpointState
|
||||
@@ -22,14 +22,18 @@ class TrainerInfo:
|
||||
|
||||
|
||||
async def create_trainers(args, *, rollout_executor: BaseWorkerHandle) -> dict[str, TrainerInfo]:
|
||||
await wait_external_trainers(args)
|
||||
trainer_configs = compute_trainer_configs(args)
|
||||
handles = create_trainer_handles(args, trainer_configs=trainer_configs)
|
||||
await wait_external_trainers(args, handles=handles)
|
||||
|
||||
trainers: dict[str, TrainerInfo] = {}
|
||||
for trainer_config in compute_trainer_configs(args):
|
||||
for trainer_config in trainer_configs:
|
||||
model_id = trainer_config.model_id
|
||||
assert model_id is not None, f"{trainer_config} carries no policy model id"
|
||||
created = await create_training_model(
|
||||
compute_trainer_args(args, trainer_config), trainer_id=trainer_config.trainer_id
|
||||
compute_trainer_args(args, trainer_config),
|
||||
handle=handles[trainer_config.trainer_id],
|
||||
trainer_id=trainer_config.trainer_id,
|
||||
)
|
||||
assert model_id not in trainers, f"{trainer_config} shares its model id with an already created trainer"
|
||||
trainers[model_id] = TrainerInfo(
|
||||
|
||||
@@ -412,6 +412,19 @@ class TestCreateTrainingModels:
|
||||
|
||||
assert requested == ["alpha-actor"]
|
||||
|
||||
async def test_an_external_trainer_is_identified_and_driven_through_one_handle(self, tmp_path, monkeypatch):
|
||||
"""A second handle would identify one connection and drive another, so the check would guard nothing."""
|
||||
handles = self._patched(monkeypatch, [])
|
||||
monkeypatch.setattr(placement_group_module, "wait_static_addrs_ready", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(placement_group_module, "_assert_external_trainer_in_run", lambda identity, **kwargs: None)
|
||||
args = self._args(tmp_path, trainer_controller_addrs=["alpha-actor=10.0.0.5:1234"])
|
||||
|
||||
await create_training_models(args, self._rollout_executor())
|
||||
|
||||
[handle] = handles
|
||||
called = [name for name, _args, _kwargs in handle.mock_calls]
|
||||
assert called.index("get_deployment_identity") < called.index("is_initialized") < called.index("init")
|
||||
|
||||
async def test_a_run_without_a_megatron_config_still_addresses_the_actor_and_critic_pools(self, monkeypatch):
|
||||
"""Every existing single policy deployment names its two pools 'actor' and 'critic'."""
|
||||
requested: list[str] = []
|
||||
@@ -456,69 +469,65 @@ class TestCreateTrainingModels:
|
||||
|
||||
class TestCreateTrainingModel:
|
||||
@staticmethod
|
||||
def _patch_handle(monkeypatch, *, restored: list[int]) -> None:
|
||||
def _create_handle(*, capability, trainer_id: str):
|
||||
handle = MagicMock()
|
||||
handle.init = AsyncMock(return_value=restored)
|
||||
return handle
|
||||
def _handle(*, restored: list[int]) -> MagicMock:
|
||||
handle = MagicMock()
|
||||
handle.init = AsyncMock(return_value=restored)
|
||||
return handle
|
||||
|
||||
monkeypatch.setattr(placement_group_module, "create_trainer_controller_handle", _create_handle)
|
||||
monkeypatch.setattr(placement_group_module, "get_backend_capability", lambda args: object())
|
||||
|
||||
async def test_a_trainer_whose_cells_restored_different_rollouts_is_refused(self, monkeypatch):
|
||||
async def test_a_trainer_whose_cells_restored_different_rollouts_is_refused(self):
|
||||
"""Cells of one trainer hold one model, so disagreeing positions mean a corrupted checkpoint set."""
|
||||
self._patch_handle(monkeypatch, restored=[5, 4])
|
||||
|
||||
with pytest.raises(AssertionError, match=r"trainer 'alpha-actor' restored \[5, 4\]"):
|
||||
await create_training_model(Namespace(start_rollout_id=None), trainer_id="alpha-actor")
|
||||
await create_training_model(
|
||||
Namespace(start_rollout_id=None), handle=self._handle(restored=[5, 4]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
async def test_a_trainer_starts_where_its_cells_restored(self, monkeypatch):
|
||||
async def test_a_trainer_starts_where_its_cells_restored(self):
|
||||
"""The restored position is what makes a resume continue instead of retraining old rounds."""
|
||||
self._patch_handle(monkeypatch, restored=[3, 3])
|
||||
|
||||
info = await create_training_model(Namespace(start_rollout_id=None), trainer_id="alpha-actor")
|
||||
info = await create_training_model(
|
||||
Namespace(start_rollout_id=None), handle=self._handle(restored=[3, 3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert info.start_rollout_id == 3
|
||||
|
||||
async def test_an_explicit_start_rollout_id_wins_over_the_restored_one(self, monkeypatch):
|
||||
async def test_an_explicit_start_rollout_id_wins_over_the_restored_one(self):
|
||||
"""--start-rollout-id is the manual override for replaying or skipping rounds."""
|
||||
self._patch_handle(monkeypatch, restored=[3])
|
||||
|
||||
info = await create_training_model(Namespace(start_rollout_id=9), trainer_id="alpha-actor")
|
||||
info = await create_training_model(
|
||||
Namespace(start_rollout_id=9), handle=self._handle(restored=[3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert info.start_rollout_id == 9
|
||||
|
||||
async def test_a_trainer_told_to_start_elsewhere_than_it_restored_says_so(self, monkeypatch, caplog):
|
||||
async def test_a_trainer_told_to_start_elsewhere_than_it_restored_says_so(self, caplog):
|
||||
"""A trainer that silently starts somewhere other than where it restored gives the operator nothing to read."""
|
||||
self._patch_handle(monkeypatch, restored=[3])
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="miles.ray.placement_group"):
|
||||
await create_training_model(Namespace(start_rollout_id=9), trainer_id="alpha-actor")
|
||||
await create_training_model(
|
||||
Namespace(start_rollout_id=9), handle=self._handle(restored=[3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert "alpha-actor" in caplog.text and "--start-rollout-id 9" in caplog.text
|
||||
|
||||
async def test_a_trainer_told_to_start_where_it_restored_says_nothing(self, monkeypatch, caplog):
|
||||
async def test_a_trainer_told_to_start_where_it_restored_says_nothing(self, caplog):
|
||||
"""Logging every trainer that was told where it already stands is noise on every ordinary launch."""
|
||||
self._patch_handle(monkeypatch, restored=[3])
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="miles.ray.placement_group"):
|
||||
await create_training_model(Namespace(start_rollout_id=3), trainer_id="alpha-actor")
|
||||
await create_training_model(
|
||||
Namespace(start_rollout_id=3), handle=self._handle(restored=[3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert "--start-rollout-id" not in caplog.text
|
||||
|
||||
async def test_a_trainer_left_to_its_restored_position_says_nothing(self, monkeypatch, caplog):
|
||||
async def test_a_trainer_left_to_its_restored_position_says_nothing(self, caplog):
|
||||
"""The ordinary resume names no rollout at all, and it must not be reported as an override."""
|
||||
self._patch_handle(monkeypatch, restored=[3])
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="miles.ray.placement_group"):
|
||||
await create_training_model(Namespace(start_rollout_id=None), trainer_id="alpha-actor")
|
||||
await create_training_model(
|
||||
Namespace(start_rollout_id=None), handle=self._handle(restored=[3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert "--start-rollout-id" not in caplog.text
|
||||
|
||||
async def test_the_restored_position_is_kept_beside_the_overridden_start(self, monkeypatch):
|
||||
async def test_the_restored_position_is_kept_beside_the_overridden_start(self):
|
||||
"""Cross trainer checks compare where checkpoints actually were, which an override must not rewrite."""
|
||||
self._patch_handle(monkeypatch, restored=[3])
|
||||
|
||||
info = await create_training_model(Namespace(start_rollout_id=9), trainer_id="alpha-actor")
|
||||
info = await create_training_model(
|
||||
Namespace(start_rollout_id=9), handle=self._handle(restored=[3]), trainer_id="alpha-actor"
|
||||
)
|
||||
|
||||
assert info.restored_rollout_id == 3
|
||||
|
||||
@@ -147,13 +147,17 @@ class TestCreatePolicyTrainers:
|
||||
def _stub_create_training_model(monkeypatch, start_rollout_ids: dict[str, int]) -> list[dict]:
|
||||
created: list[dict] = []
|
||||
|
||||
async def _create(trainer_args, *, trainer_id):
|
||||
handle = AsyncMock()
|
||||
async def _create(trainer_args, *, handle, trainer_id):
|
||||
handle.get_train_parallel_config = AsyncMock(return_value=f"parallel-config-of-{trainer_id}")
|
||||
created.append(dict(trainer_id=trainer_id, args=trainer_args, handle=handle))
|
||||
return SimpleNamespace(handle=handle, start_rollout_id=start_rollout_ids[trainer_args.trainer_model_id])
|
||||
|
||||
monkeypatch.setattr(multi_policy_utils, "create_training_model", _create)
|
||||
monkeypatch.setattr(
|
||||
multi_policy_utils,
|
||||
"create_trainer_handles",
|
||||
lambda args, *, trainer_configs: {config.trainer_id: AsyncMock() for config in trainer_configs},
|
||||
)
|
||||
return created
|
||||
|
||||
async def test_every_policy_gets_a_trainer_keyed_by_its_model_id(self, monkeypatch):
|
||||
|
||||
Reference in New Issue
Block a user