From da3dc08decb23f0a1e9336706b944de956721571 Mon Sep 17 00:00:00 2001 From: Neil Kale <263453039+kalectory@users.noreply.github.com> Date: Tue, 11 Aug 2026 17:54:38 +0000 Subject: [PATCH] Add watchdog for stalled Ray initialization --- skyrl/train/utils/utils.py | 29 +++++++++++++++++++++++- tests/train/utils/test_initialize_ray.py | 29 ++++++++++++++++++++++++ 2 files changed, 57 insertions(+), 1 deletion(-) create mode 100644 tests/train/utils/test_initialize_ray.py diff --git a/skyrl/train/utils/utils.py b/skyrl/train/utils/utils.py index e90c59630e..bf4a4e664d 100644 --- a/skyrl/train/utils/utils.py +++ b/skyrl/train/utils/utils.py @@ -1,3 +1,4 @@ +import faulthandler import functools import ipaddress import logging @@ -5,6 +6,7 @@ import os import socket import sys +import threading import time from copy import deepcopy from datetime import datetime @@ -29,6 +31,23 @@ from skyrl.train.config.config import SkyRLTrainConfig +def _start_ray_init_watchdog(timeout_s: float) -> threading.Event: + """Exit the driver if Ray initialization stops making progress.""" + completed = threading.Event() + if timeout_s <= 0: + return completed + + def watchdog() -> None: + if completed.wait(timeout_s): + return + logger.error(f"ray.init() did not complete within {timeout_s:g}s; terminating the driver") + faulthandler.dump_traceback(file=sys.stderr, all_threads=True) + os._exit(1) + + threading.Thread(target=watchdog, name="skyrl-ray-init-watchdog", daemon=True).start() + return completed + + class Timer: def __init__(self, message, update_dict=None): self.message = message @@ -918,7 +937,15 @@ def initialize_ray(cfg: SkyRLTrainConfig): # log_to_driver=True allows training progress from skyrl_entrypoint to reach stdout. # Infrastructure logs (vLLM, workers) are redirected to log file via os.dup2 in their init. - ray.init(runtime_env={"env_vars": env_vars}, log_to_driver=True) + ray_init_timeout_s = float(os.environ.get("SKYRL_RAY_INIT_TIMEOUT_IN_S", "0")) + logger.info(f"Starting ray.init() (timeout={ray_init_timeout_s:g}s; 0 disables the watchdog)") + started_at = time.monotonic() + ray_init_completed = _start_ray_init_watchdog(ray_init_timeout_s) + try: + ray.init(runtime_env={"env_vars": env_vars}, log_to_driver=True) + finally: + ray_init_completed.set() + logger.info(f"ray.init() completed in {time.monotonic() - started_at:.1f}s") if not verbose_logging: logger.info(f"Infrastructure logs will be written to: {log_file}") diff --git a/tests/train/utils/test_initialize_ray.py b/tests/train/utils/test_initialize_ray.py new file mode 100644 index 0000000000..630fa081b2 --- /dev/null +++ b/tests/train/utils/test_initialize_ray.py @@ -0,0 +1,29 @@ +import threading + +from skyrl.train.utils import utils + + +def test_ray_init_watchdog_exits_after_timeout(monkeypatch): + exited = threading.Event() + exit_codes = [] + + def fake_exit(code): + exit_codes.append(code) + exited.set() + + monkeypatch.setattr(utils.faulthandler, "dump_traceback", lambda **kwargs: None) + monkeypatch.setattr(utils.os, "_exit", fake_exit) + + completed = utils._start_ray_init_watchdog(0.01) + + assert exited.wait(timeout=1) + assert exit_codes == [1] + completed.set() + + +def test_ray_init_watchdog_can_be_disabled(monkeypatch): + monkeypatch.setattr(utils.os, "_exit", lambda code: (_ for _ in ()).throw(AssertionError(code))) + + completed = utils._start_ray_init_watchdog(0) + + assert not completed.is_set()