Skip to content
Merged
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
76 changes: 4 additions & 72 deletions miles/ray/rollout.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import logging
import multiprocessing
import random
import threading
import time
from pathlib import Path
from typing import List, Union
Expand All @@ -13,6 +12,7 @@

from miles.backends.sglang_utils.sglang_engine import SGLangEngine
from miles.ray.rollout_data_source import RolloutDataSourceWithBuffer
from miles.utils.health_monitor import RolloutHealthMonitor
from miles.utils.http_utils import find_available_port, get_host_info, init_http_client
from miles.utils.metric_checker import MetricChecker
from miles.utils.misc import load_function
Expand Down Expand Up @@ -63,13 +63,7 @@ def __init__(self, args, pg, wandb_run_id):
self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote()

self._metric_checker = MetricChecker.maybe_create(args)

# fault tolerance
self._health_monitor_thread = None
self._health_monitor_stop_event = None
self._health_check_interval = args.rollout_health_check_interval
self._health_check_timeout = args.rollout_health_check_timeout
self._health_check_first_wait = args.rollout_health_check_first_wait
self._health_monitor = RolloutHealthMonitor(self, args)

def dispose(self):
if self._metric_checker is not None:
Expand All @@ -83,7 +77,7 @@ def get_num_rollout_per_epoch(self):
return len(self.data_source.dataset) // self.args.rollout_batch_size

def generate(self, rollout_id):
monitor_started = self._start_health_monitor()
monitor_started = self._health_monitor.start()
start_time = time.time()
try:
data = self._get_rollout_data(rollout_id=rollout_id)
Expand All @@ -93,7 +87,7 @@ def generate(self, rollout_id):
return Box(ray.put(data))
finally:
if monitor_started:
self._stop_health_monitor()
self._health_monitor.stop()
self.num_new_engines = init_rollout_engines(self.args, self.pg, self.all_rollout_engines)
self.rollout_engines = self.all_rollout_engines[:: self.nodes_per_engine]

Expand All @@ -119,68 +113,6 @@ def offload(self):
def onload(self, tags: List[str] = None):
return [engine.resume_memory_occupation.remote(tags=tags) for engine in self.rollout_engines]

def _start_health_monitor(self) -> bool:
if not self.rollout_engines:
return False

assert self._health_monitor_thread is None, "Health monitor thread is already running."

self._health_monitor_stop_event = threading.Event()
self._health_monitor_thread = threading.Thread(
target=self._health_monitor_loop,
name="RolloutHealthMonitor",
daemon=True,
)
self._health_monitor_thread.start()
return True

def _stop_health_monitor(self) -> None:
if not self._health_monitor_thread:
return

assert self._health_monitor_stop_event is not None
self._health_monitor_stop_event.set()
timeout = self._health_check_timeout + self._health_check_interval + 5
self._health_monitor_thread.join(timeout=timeout)
if self._health_monitor_thread.is_alive():
logging.warning("Rollout health monitor thread did not terminate within %.1fs", timeout)

self._health_monitor_thread = None
self._health_monitor_stop_event = None

def _health_monitor_loop(self) -> None:
assert self._health_monitor_stop_event is not None
# TODO: need to be waiting for the large moe to be ready. this is hacky.
if self._health_monitor_stop_event.wait(self._health_check_first_wait):
return
while not self._health_monitor_stop_event.is_set():
self._run_health_checks()
if self._health_monitor_stop_event.wait(self._health_check_interval):
break

def _run_health_checks(self) -> None:
for rollout_engine_id, engine in enumerate(self.rollout_engines):
if self._health_monitor_stop_event is not None and self._health_monitor_stop_event.is_set():
break
self._check_engine_health(rollout_engine_id, engine)

def _check_engine_health(self, rollout_engine_id, engine) -> None:
if engine is None:
return

try:
ray.get(engine.health_generate.remote(timeout=self._health_check_timeout))
except Exception as e:
print(f"Health check timed out for rollout engine {rollout_engine_id} (ray timeout). Killing actor.")
for i in range(rollout_engine_id * self.nodes_per_engine, (rollout_engine_id + 1) * self.nodes_per_engine):
engine = self.all_rollout_engines[i]
try:
ray.kill(engine)
except Exception:
pass
self.all_rollout_engines[i] = None
self.rollout_engines[rollout_engine_id] = None

def _get_rollout_data(self, rollout_id):
if self.args.load_debug_rollout_data:
data = torch.load(
Expand Down
84 changes: 84 additions & 0 deletions miles/utils/health_monitor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
import logging
import threading

import ray


class RolloutHealthMonitor:
def __init__(self, rollout_manager, args):
# TODO may remove this dependency after refactoring
self._rollout_manager = rollout_manager

self._thread = None
self._stop_event = None
self._check_interval = args.rollout_health_check_interval
self._check_timeout = args.rollout_health_check_timeout
self._check_first_wait = args.rollout_health_check_first_wait

def start(self) -> bool:
if not self._rollout_manager.rollout_engines:
return False

assert self._thread is None, "Health monitor thread is already running."

self._stop_event = threading.Event()
self._thread = threading.Thread(
target=self._health_monitor_loop,
name="RolloutHealthMonitor",
daemon=True,
)
self._thread.start()
return True

def stop(self) -> None:
if not self._thread:
return

assert self._stop_event is not None
self._stop_event.set()
timeout = self._check_timeout + self._check_interval + 5
self._thread.join(timeout=timeout)
if self._thread.is_alive():
logging.warning("Rollout health monitor thread did not terminate within %.1fs", timeout)

self._thread = None
self._stop_event = None

def _health_monitor_loop(self) -> None:
assert self._stop_event is not None
# TODO: need to be waiting for the large moe to be ready. this is hacky.
if self._stop_event.wait(self._check_first_wait):
return
while not self._stop_event.is_set():
self._run_health_checks()
if self._stop_event.wait(self._check_interval):
break

def _run_health_checks(self) -> None:
for rollout_engine_id, engine in enumerate(self._rollout_manager.rollout_engines):
if self._stop_event is not None and self._stop_event.is_set():
break
self._check_engine_health(rollout_engine_id, engine)

def _check_engine_health(self, rollout_engine_id, engine) -> None:
if engine is None:
return

try:
ray.get(engine.health_generate.remote(timeout=self._check_timeout))
except Exception as e:
print(f"Health check timed out for rollout engine {rollout_engine_id} (ray timeout). Killing actor.")
self._kill_engine(rollout_engine_id=rollout_engine_id)

def _kill_engine(self, rollout_engine_id: int):
for i in range(
rollout_engine_id * self._rollout_manager.nodes_per_engine,
(rollout_engine_id + 1) * self._rollout_manager.nodes_per_engine,
):
engine = self._rollout_manager.all_rollout_engines[i]
try:
ray.kill(engine)
except Exception as e:
print(f"Fail to kill engine and skip (e: {e})")
self._rollout_manager.all_rollout_engines[i] = None
self._rollout_manager.rollout_engines[rollout_engine_id] = None