Files
myLightningOPD/slime/utils/health_monitor.py
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

107 lines
4.1 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import threading
import ray
logger = logging.getLogger(__name__)
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."
logger.info("Starting RolloutHealthMonitor...")
self._stop_event = threading.Event()
self._thread = threading.Thread(
target=self._health_monitor_loop,
name="RolloutHealthMonitor",
daemon=True,
)
self._thread.start()
logger.info("RolloutHealthMonitor started.")
return True
def stop(self) -> None:
if not self._thread:
return
logger.info("Stopping RolloutHealthMonitor...")
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)
else:
logger.info("RolloutHealthMonitor stopped.")
self._thread = None
self._stop_event = None
def _health_monitor_loop(self) -> None:
assert self._stop_event is not None
logger.info(f"Health monitor loop started. Waiting for first wait: {self._check_first_wait}s")
# TODO: need to be waiting for the large moe to be ready. this is hacky.
if self._stop_event.wait(self._check_first_wait):
logger.info("Health monitor stopped during 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:
logger.info(f"Skipping health check for engine {rollout_engine_id} (None)")
return
try:
ray.get(engine.health_generate.remote(timeout=self._check_timeout))
except Exception as e:
logger.error(
f"Health check failed for rollout engine {rollout_engine_id} (ray timeout or error). Killing actor. Exception: {e}"
)
self._kill_engine(rollout_engine_id=rollout_engine_id)
def _kill_engine(self, rollout_engine_id: int):
logger.info(f"Killing engine group {rollout_engine_id}...")
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]
if engine:
logger.info(f"Shutting down and killing engine at index {i}")
try:
ray.get(engine.shutdown.remote())
ray.kill(engine)
logger.info(f"Successfully killed engine at index {i}")
except Exception as e:
logger.warning(f"Fail to kill engine at index {i} (e: {e})")
else:
logger.info(f"Engine at index {i} is already None")
self._rollout_manager.all_rollout_engines[i] = None