Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
88 changes: 88 additions & 0 deletions learning_loop_node/tests/unit/test_subprocess.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,16 @@
import asyncio
import multiprocessing
import os
import queue
import signal
import sys
from collections.abc import Callable, Iterator
from typing import Any

import pytest

from ...trainer import subprocess as iterator_module
from ...trainer.exceptions import UnexpectedWorkerExitError
from ...trainer.subprocess import iterator_cpu_bound


Expand Down Expand Up @@ -51,3 +62,80 @@ async def test_leaving_early_does_not_leave_the_process_running():
if item == 2:
break
# reaching here without hanging is the test


@pytest.mark.parametrize(('mode', 'exit_code'), [
('system_exit', 1),
('system_exit_zero', 0),
('native_exit', 7),
('kill', -signal.SIGKILL),
])
async def test_iterator_rejects_exit_without_completion(mode: str, exit_code: int) -> None:
with pytest.raises(UnexpectedWorkerExitError, match=rf'exited with code {exit_code} .*IteratorDone'):
await asyncio.wait_for(_collect_worker(_exit_without_completion, mode), timeout=10)


@pytest.mark.parametrize('error', [SystemExit(1), ValueError('worker failed')])
async def test_iterator_reports_failure_after_progress(error: BaseException) -> None:
seen = []

async def consume() -> None:
async with iterator_cpu_bound(_fail_after_progress, error) as results:
async for item in results:
seen.append(item)

expected_error = UnexpectedWorkerExitError if isinstance(error, SystemExit) else ValueError
with pytest.raises(expected_error):
await asyncio.wait_for(consume(), timeout=10)
assert seen == [0]


async def test_iterator_accepts_completion_arriving_after_queue_timeout(monkeypatch: pytest.MonkeyPatch) -> None:
context = multiprocessing.get_context('spawn')
make_process = context.Process
processes = []

def record_process(*args: Any, **kwargs: Any) -> Any:
process = make_process(*args, **kwargs)
processes.append(process)
return process

to_thread = asyncio.to_thread
first_poll = True

async def delayed_timeout(func: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
nonlocal first_poll
if first_poll:
first_poll = False
await to_thread(processes[0].join, 5)
assert processes[0].exitcode == 0
raise queue.Empty
return await to_thread(func, *args, **kwargs)

monkeypatch.setattr(context, 'Process', record_process)
monkeypatch.setattr(iterator_module.asyncio, 'to_thread', delayed_timeout)

assert await asyncio.wait_for(_collect_worker(counting, 0), timeout=10) == []


async def _collect_worker(it: Callable[..., Iterator[int]], *args: Any) -> list[int]:
async with iterator_cpu_bound(it, *args) as results:
return [item async for item in results]


def _exit_without_completion(mode: str) -> Iterator[int]:
yield from ()
if mode == 'system_exit':
sys.exit(1)
if mode == 'system_exit_zero':
sys.exit(0)
if mode == 'native_exit':
os._exit(7)
if mode == 'kill':
os.kill(os.getpid(), signal.SIGKILL)
raise AssertionError(f'Unexpected exit mode: {mode}')


def _fail_after_progress(error: BaseException) -> Iterator[int]:
yield 0
raise error
4 changes: 4 additions & 0 deletions learning_loop_node/trainer/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,10 @@ class InsufficientMemoryError(RuntimeError):
"""Raised when not even the smallest unit of work fits in memory."""


class UnexpectedWorkerExitError(RuntimeError):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A RuntimeError is what _perform_state treats as retryable: it resets the state to TrainModelDownloaded and starts _train again. The non-finite loss from the description fails the same way every time, so the node would keep retraining and never reach ReadyForCleanup. A worker killed by the OOM killer should be retried though. Maybe CriticalError when the worker exited on its own, and a retry only when a signal killed it?

"""Raised when an iterator worker exits without reporting completion."""


class NodeNeedsRestartError(Exception):
"""
NodeNeedsRestartError is raised when the node needs to be restarted.
Expand Down
14 changes: 11 additions & 3 deletions learning_loop_node/trainer/subprocess.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
from multiprocessing.queues import Queue as MPQueue
from typing import Any, ParamSpec, TypeVar

from .exceptions import UnexpectedWorkerExitError

logger = logging.getLogger(__name__)

T = TypeVar('T')
Expand Down Expand Up @@ -55,9 +57,15 @@ async def _iterator_cpu_bound_inner(
try:
item = await asyncio.to_thread(state_queue.get, True, 0.5)
except queue.Empty:
if not process.is_alive():
break
continue
if process.is_alive():
continue
# Completion may have arrived between the timeout and the exit check.
try:
item = await asyncio.to_thread(state_queue.get_nowait)
except queue.Empty as e:
raise UnexpectedWorkerExitError(
f'{process.name} exited with code {process.exitcode} without sending IteratorDone'
) from e
Comment on lines +66 to +68

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

_iterator_wrapper catches Exception, but sys.exit(1) raises SystemExit. So the worker's reason never reaches the queue and only the exit code is left. process.name is this same literal for detection too, so a failed training and a failed detection produce the same message in the loop. Could _iterator_wrapper catch SystemExit as well and pass the reason along? This error would then be left for the cases where there is nothing to report anyway: os._exit, SIGKILL, the OOM killer. The module docstring would need the new case too.

match item:
case IteratorDone():
break
Expand Down
Loading