From dfcb8e0597403a1340cc6b72be666be66cf11d66 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sat, 26 Sep 2026 20:38:24 +0800 Subject: [PATCH] Keep the first multi-policy failure and stop followers by flag (#3060) --- miles/utils/async_utils.py | 31 ++++- miles/utils/multi_policy/parker.py | 1 + tests/fast/test_train_multi_policy.py | 63 ++++++++- tests/fast/utils/multi_policy/test_parker.py | 17 +++ tests/fast/utils/test_async_utils.py | 131 +++++++++++++++++++ train_multi_policy.py | 7 +- 6 files changed, 240 insertions(+), 10 deletions(-) diff --git a/miles/utils/async_utils.py b/miles/utils/async_utils.py index 3a05e0e367..254e0bef07 100644 --- a/miles/utils/async_utils.py +++ b/miles/utils/async_utils.py @@ -2,6 +2,7 @@ import asyncio import concurrent.futures import logging import threading +import traceback from collections.abc import Awaitable, Callable, Coroutine, Sequence from typing import Any, TypeVar @@ -85,19 +86,37 @@ def wait_futures(futures: Sequence[concurrent.futures.Future]) -> list[Any]: return results -async def wait_cancelling_pending_on_first_completion(tasks: Sequence[asyncio.Task]) -> None: - _, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) +async def wait_cancelling_pending_on_first_completion( + tasks: Sequence[asyncio.Task], *, on_first_completion: Callable[[], None] | None = None +) -> None: + done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + + if on_first_completion is not None: + on_first_completion() for task in pending: task.cancel() if pending: await asyncio.gather(*pending, return_exceptions=True) - errors = [error for task in tasks if (error := _compute_task_error(task)) is not None] - for error in errors: + task_errors = [(task, error) for task in tasks if (error := _compute_task_error(task)) is not None] + for _, error in task_errors: logger.error("task failed", exc_info=error) - if errors: - raise errors[0] + if task_errors: + primary_index = next((index for index, (task, _) in enumerate(task_errors) if task in done), 0) + primary_error = task_errors[primary_index][1] + for index, (_, error) in enumerate(task_errors): + if index != primary_index: + note = "Additional task failure while cancelling peers:\n" + "".join(traceback.format_exception(error)) + _exception_add_note_or_log(primary_error, note) + raise primary_error + + +def _exception_add_note_or_log(e: BaseException, msg: str) -> None: + if hasattr(e, "add_note"): + e.add_note(msg) + else: + logger.error(msg) def _compute_task_error(task: asyncio.Task) -> BaseException | None: diff --git a/miles/utils/multi_policy/parker.py b/miles/utils/multi_policy/parker.py index 9d2a9e1fc5..5d916db609 100644 --- a/miles/utils/multi_policy/parker.py +++ b/miles/utils/multi_policy/parker.py @@ -26,6 +26,7 @@ class Parker: self._event.set() async def maybe_park_follower(self) -> None: + await asyncio.sleep(0) self._num_ready += 1 await self._event.wait() self._num_ready -= 1 diff --git a/tests/fast/test_train_multi_policy.py b/tests/fast/test_train_multi_policy.py index 3abf189f00..be202d5ef2 100644 --- a/tests/fast/test_train_multi_policy.py +++ b/tests/fast/test_train_multi_policy.py @@ -38,6 +38,8 @@ 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_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() @@ -95,6 +97,22 @@ 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_never_returning(rollout_id: int, rollout_data_ref: Any, **kwargs: Any) -> None: + await asyncio.Event().wait() + + +def _let_follower_yield(handle) -> None: + async def yield_to_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_leader + + class TestInitialWeightPublication: 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.""" @@ -134,7 +152,10 @@ class TestRunPolicies: async def test_a_policy_only_resumes_the_health_probing_of_its_own_engines(self): """Resuming the whole fleet here un-pauses probing of a policy that is mid weight broadcast.""" - context = await _run(_make_args(num_rollout=1)) + trainers = {"a": AsyncMock(), "b": AsyncMock()} + trainers["b"].train = _train_never_returning + + context = await _run(_make_args(num_rollout=1), trainers=trainers) prepared = context["inference_controller"].prepare_rollout.await_args_list assert sorted((call.args[0], call.kwargs["model_id"]) for call in prepared) == [(0, "a"), (0, "b")] @@ -260,11 +281,17 @@ class TestSaving: async def test_a_parked_follower_is_saved_at_the_round_it_reached(self): """A record naming a policy at a rollout it never checkpointed cannot be resumed.""" trainers = {"a": AsyncMock(), "b": AsyncMock()} + saves: list[tuple[int, int]] = [] + + async def _note_follower_position(rollout_id: int, **kwargs: Any) -> None: + saves.append((rollout_id, trainers["b"].train.await_args_list[-1].args[0])) + + trainers["b"].save_model = AsyncMock(side_effect=_note_follower_position) await _run(_make_args(num_rollout=1, save=None, save_interval=1), trainers=trainers) - [saved_at] = [call.args[0] for call in trainers["b"].save_model.await_args_list] - assert saved_at == trainers["b"].train.await_args_list[-1].args[0] + [(saved_at, reached)] = saves + assert saved_at == reached async def test_every_policy_is_on_disk_before_the_record_claims_the_checkpoint_exists(self): """An asynchronous follower save still running would leave the record pointing at files nobody wrote.""" @@ -342,3 +369,33 @@ 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 diff --git a/tests/fast/utils/multi_policy/test_parker.py b/tests/fast/utils/multi_policy/test_parker.py index 9e25713732..be226a5c64 100644 --- a/tests/fast/utils/multi_policy/test_parker.py +++ b/tests/fast/utils/multi_policy/test_parker.py @@ -103,3 +103,20 @@ class TestMaybeParkFollower: await asyncio.wait_for(task, timeout=5) assert finished + + +class TestTheFollowerAlwaysGivesTheLoopATurn: + async def test_an_open_gate_still_suspends_the_follower(self): + """An open gate that returns without suspending lets the follower's unbounded loop starve the loop.""" + parker = Parker(num_followers=1) + other_task_ran = False + + async def other_task() -> None: + nonlocal other_task_ran + other_task_ran = True + + asyncio.create_task(other_task()) + + await parker.maybe_park_follower() + + assert other_task_ran diff --git a/tests/fast/utils/test_async_utils.py b/tests/fast/utils/test_async_utils.py index 4d73537dea..682e0cc695 100644 --- a/tests/fast/utils/test_async_utils.py +++ b/tests/fast/utils/test_async_utils.py @@ -468,6 +468,38 @@ class TestWaitFutures: @pytest.mark.asyncio class TestWaitCancellingPendingOnFirstCompletion: + async def test_callback_runs_before_pending_tasks_are_cancelled(self) -> None: + """The first-completion callback runs once before a pending follower observes cancellation.""" + follower_started = asyncio.Event() + leader_release = asyncio.Event() + events: list[str] = [] + + async def leader() -> None: + await leader_release.wait() + + async def follower() -> None: + follower_started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + events.append("follower cancelled") + raise + + leader_task: asyncio.Task[None] = asyncio.create_task(leader()) + follower_task: asyncio.Task[None] = asyncio.create_task(follower()) + await follower_started.wait() + + def on_first_completion() -> None: + assert follower_task.cancelling() == 0 + events.append("callback") + + leader_release.set() + await wait_cancelling_pending_on_first_completion( + [leader_task, follower_task], on_first_completion=on_first_completion + ) + + assert events == ["callback", "follower cancelled"] + async def test_the_first_task_to_return_cancels_the_rest(self): """The leader owns the length of a multi policy run, and its followers loop until it is over.""" cancelled = False @@ -557,6 +589,26 @@ class TestWaitCancellingPendingOnFirstCompletion: [asyncio.create_task(failing()), asyncio.create_task(slow())] ) + async def test_primary_failure_precedes_a_cleanup_failure_from_an_earlier_task(self): + """Cancellation cleanup is attached without masking the failure that ended the run.""" + + async def cleanup_failure(): + try: + await asyncio.sleep(30) + except asyncio.CancelledError: + raise RuntimeError("cleanup exploded") from None + + async def primary_failure(): + await asyncio.sleep(0.01) + raise ValueError("training failed") + + with pytest.raises(ValueError, match="training failed") as exc_info: + await wait_cancelling_pending_on_first_completion( + [asyncio.create_task(cleanup_failure()), asyncio.create_task(primary_failure())] + ) + + assert any("RuntimeError: cleanup exploded" in note for note in exc_info.value.__notes__) + async def test_the_cancelled_members_are_awaited_before_the_error_is_raised(self): """Raising before the cleanup lands would let a half-cancelled task outlive the caller.""" cleaned_up = False @@ -622,3 +674,82 @@ class TestGatherAndRaiseFirst: await async_utils.gather_and_raise_first([_failing("quiet")]) assert caplog.text == "" + + +class TestGetAsyncLoop: + def test_threads_arriving_together_share_one_loop(self, monkeypatch): + """A second loop would strand whatever already awaits on the first, which no caller can recover from.""" + monkeypatch.setattr(async_utils, "async_loop", None) + built: list[object] = [] + + class SlowToBuild: + def __init__(self): + time.sleep(0.05) + built.append(self) + + monkeypatch.setattr(async_utils, "AsyncLoopThread", SlowToBuild) + + with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool: + loops = [future.result() for future in [pool.submit(async_utils.get_async_loop) for _ in range(8)]] + + assert len(built) == 1 + assert all(loop is built[0] for loop in loops) + + def test_a_later_caller_is_answered_from_the_loop_already_built(self, monkeypatch): + """The guard must not rebuild the loop once one exists, nor pay a lock on every rollout call.""" + monkeypatch.setattr(async_utils, "async_loop", None) + monkeypatch.setattr(async_utils, "AsyncLoopThread", lambda: object()) + + first = async_utils.get_async_loop() + + assert async_utils.get_async_loop() is first + + +class _ErrorWithoutAddNote(Exception): + def __getattribute__(self, name: str): + if name == "add_note": + raise AttributeError(name) + return super().__getattribute__(name) + + +class TestKeepingASecondaryFailureWhereNotesDoNotExist: + def test_a_secondary_failure_is_logged_when_the_primary_cannot_carry_notes(self, caplog): + """On 3.10 the annotation raised AttributeError from the error path and replaced the failure it described.""" + primary = _ErrorWithoutAddNote("training failed") + + with caplog.at_level(logging.ERROR): + async_utils._exception_add_note_or_log(primary, "RuntimeError: cleanup") + + assert "RuntimeError: cleanup" in caplog.text + + def test_a_primary_that_can_carry_notes_is_annotated_rather_than_logged(self, caplog): + """On 3.11 and later the note travels with the exception, which is where a caller reads it.""" + primary = ValueError("training failed") + + with caplog.at_level(logging.ERROR): + async_utils._exception_add_note_or_log(primary, "RuntimeError: cleanup") + + assert any("RuntimeError: cleanup" in note for note in primary.__notes__) + assert caplog.text == "" + + async def test_the_primary_failure_still_reaches_the_caller_unreplaced(self, caplog): + """Annotating the failure must never be able to become the failure the run reports.""" + + async def cleanup_failure(): + try: + await asyncio.sleep(30) + except asyncio.CancelledError: + raise RuntimeError("cleanup exploded") from None + + async def primary_failure(): + await asyncio.sleep(0.01) + raise _ErrorWithoutAddNote("training failed") + + with caplog.at_level(logging.ERROR): + with pytest.raises(_ErrorWithoutAddNote, match="training failed"): + await wait_cancelling_pending_on_first_completion( + [asyncio.create_task(cleanup_failure()), asyncio.create_task(primary_failure())] + ) + + assert "Additional task failure while cancelling peers" in caplog.text + assert "RuntimeError: cleanup exploded" in caplog.text diff --git a/train_multi_policy.py b/train_multi_policy.py index 69743b607c..72547b5570 100644 --- a/train_multi_policy.py +++ b/train_multi_policy.py @@ -70,6 +70,7 @@ async def train_multi_policy(args) -> None: ) parker = Parker(num_followers=len(trainers) - 1) + run_ended = asyncio.Event() rollout_ids: dict[str, int] = {} tasks = [ asyncio.create_task( @@ -81,13 +82,14 @@ async def train_multi_policy(args) -> None: inference_controller=inference_controller, rollout_executor=rollout_executor, parker=parker, + run_ended=run_ended, rollout_ids=rollout_ids, num_rollout_per_epoch=num_rollout_per_epoch, ) ) for trainer in trainers.values() ] - await wait_cancelling_pending_on_first_completion(tasks) + await wait_cancelling_pending_on_first_completion(tasks, on_first_completion=run_ended.set) await rollout_executor.dispose() await inference_controller.dispose() @@ -107,6 +109,7 @@ async def _run_policy( *, trainer: TrainerInfo, is_leader: bool, + run_ended: asyncio.Event, trainers: dict[str, TrainerInfo], inference_controller: BaseWorkerHandle, rollout_executor: BaseWorkerHandle, @@ -120,6 +123,8 @@ async def _run_policy( range(trainer.start_rollout_id, args.num_rollout) if is_leader else itertools.count(trainer.start_rollout_id) ) for rollout_id in rollout_ids_iter: + if run_ended.is_set(): + return rollout_ids[model_id] = rollout_id await inference_controller.prepare_rollout(rollout_id, model_id=model_id) rollout_data_pack = await rollout_executor.get(rollout_id, trainer_model_id=model_id)