mirror of
https://github.com/radixark/miles.git
synced 2026-10-01 23:06:14 +08:00
Keep the first multi-policy failure and stop followers by flag (#3060)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user