Keep the first multi-policy failure and stop followers by flag (#3060)

This commit is contained in:
fzyzcjy
2026-09-26 20:38:24 +08:00
committed by GitHub
parent d1286304c8
commit dfcb8e0597
6 changed files with 240 additions and 10 deletions
+25 -6
View File
@@ -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:
+1
View File
@@ -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
+60 -3
View File
@@ -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
@@ -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
+131
View File
@@ -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
+6 -1
View File
@@ -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)