[feat] support fault tolerant for rollout engines (#405)

* [feat] support fault tolerant for rollout engines

* support fault tolerant for UpdateWeightFromDistributed

* bugfix

* bugfix
This commit is contained in:
Zilin Zhu
2025-09-30 14:18:13 +08:00
committed by GitHub
parent e9c677bbb3
commit dde524e2c6
9 changed files with 216 additions and 59 deletions
+4 -4
View File
@@ -99,7 +99,6 @@ class FSDPTrainRayActor(TrainRayActor):
self.update_cpu_params_dict(self.weights["actor"])
self.connected = False
self.weight_updator = (
UpdateWeightFromTensor(self.args, self.model)
if self.args.colocate
@@ -405,9 +404,10 @@ class FSDPTrainRayActor(TrainRayActor):
if self.args.debug_train_only or self.args.debug_rollout_only:
return
if not self.connected:
self.connected = True
rollout_engines, rollout_engine_lock = ray.get(self.rollout_manager.get_rollout_engines_and_lock.remote())
rollout_engines, rollout_engine_lock, num_new_engines = ray.get(
self.rollout_manager.get_rollout_engines_and_lock.remote()
)
if num_new_engines > 0:
self.weight_updator.connect_rollout_engines(rollout_engines, rollout_engine_lock)
dist.barrier(group=get_gloo_group())
+4 -3
View File
@@ -401,9 +401,10 @@ class MegatronTrainRayActor(TrainRayActor):
if self.args.offload and hasattr(mpu, "reload_process_groups"):
mpu.reload_process_groups()
if not self.connected:
self.connected = True
rollout_engines, rollout_engine_lock = ray.get(self.rollout_manager.get_rollout_engines_and_lock.remote())
rollout_engines, rollout_engine_lock, num_new_engines = ray.get(
self.rollout_manager.get_rollout_engines_and_lock.remote()
)
if num_new_engines > 0:
self.weight_updator.connect_rollout_engines(rollout_engines, rollout_engine_lock)
dist.barrier(group=get_gloo_group())
@@ -305,6 +305,17 @@ class UpdateWeightFromTensor:
self.param_info_buckets = get_param_info_buckets(self.args, self.model)
self.weight_version = 0
# create the group within megatron.
for start_rank in range(0, dist.get_world_size(), self.args.rollout_num_gpus_per_engine):
end_rank = start_rank + self.args.rollout_num_gpus_per_engine
group_ranks = list(range(start_rank, end_rank))
new_group = dist.new_group(ranks=group_ranks, backend="gloo")
if dist.get_rank() in group_ranks:
self._ipc_gather_group = new_group
self._ipc_gather_src = start_rank
self._model_update_groups = None
def connect_rollout_engines(self, rollout_engines, rollout_engine_lock):
self.rollout_engines = rollout_engines
colocate_engine_nums = (
@@ -322,6 +333,11 @@ class UpdateWeightFromTensor:
)
self._group_name = "slime"
if self._is_distributed_src_rank:
if self._model_update_groups is not None:
disconnect_rollout_engines_from_distributed(
self.args, self._group_name, self._model_update_groups, self.distributed_rollout_engines
)
self._model_update_groups = connect_rollout_engines_from_distributed(
self.args, self._group_name, self.distributed_rollout_engines
)
@@ -331,13 +347,7 @@ class UpdateWeightFromTensor:
start_rank = i * self.args.rollout_num_gpus_per_engine
end_rank = (i + 1) * self.args.rollout_num_gpus_per_engine
group_ranks = list(range(start_rank, end_rank))
new_group = dist.new_group(
ranks=group_ranks,
backend="gloo",
)
if dist.get_rank() in group_ranks:
self._ipc_gather_src = start_rank
self._ipc_gather_group = new_group
self._ipc_engine = engine
@torch.no_grad()
@@ -496,6 +506,7 @@ class UpdateWeightFromDistributed:
self.vocab_size = vocab_size
self.quantization_config = quantization_config
self.weight_version = 0
self._model_update_groups = None
def connect_rollout_engines(self, rollout_engines, rollout_engine_lock):
self.rollout_engines = rollout_engines
@@ -512,6 +523,10 @@ class UpdateWeightFromDistributed:
self._group_name = f"slime-pp_{pp_rank}"
if self._is_pp_src_rank:
if self._model_update_groups is not None:
disconnect_rollout_engines_from_distributed(
self.args, self._group_name, self._model_update_groups, self.rollout_engines
)
self._model_update_groups = connect_rollout_engines_from_distributed(
self.args, self._group_name, rollout_engines
)
@@ -670,6 +685,12 @@ def connect_rollout_engines_from_distributed(args, group_name, rollout_engines):
return model_update_groups
def disconnect_rollout_engines_from_distributed(args, group_name, model_update_groups, rollout_engines):
refs = [engine.destroy_weights_update_group.remote(group_name) for engine in rollout_engines]
dist.destroy_process_group(model_update_groups)
ray.get(refs)
def update_weights_from_distributed(args, group_name, group, weight_version, rollout_engines, converted_named_tensors):
refs = [
engine.update_weights_from_distributed.remote(
+38 -1
View File
@@ -149,6 +149,28 @@ class SGLangEngine(RayActor):
response.raise_for_status()
return response.json()
def health_generate(self, timeout: float = 5.0) -> bool:
"""Run /health_generate on the underlying SGLang HTTP server.
Args:
timeout: Timeout for the health request in seconds.
Returns:
True if the server responds with HTTP 200.
Raises:
requests.RequestException: If the request fails for any reason, including timeout.
"""
if self.node_rank != 0:
return True
response = requests.get(
f"http://{self.server_args.host}:{self.server_args.port}/health_generate",
timeout=timeout,
)
response.raise_for_status()
return True
def update_weights_from_tensor(
self,
serialized_named_tensors: List[str],
@@ -179,7 +201,7 @@ class SGLangEngine(RayActor):
if self.node_rank != 0:
return
# flush cache will not return status_code 200 when there are pending requests
while True:
for _ in range(60):
try:
response = requests.get(f"http://{self.server_args.host}:{self.server_args.port}/flush_cache")
if response.status_code == 200:
@@ -188,7 +210,10 @@ class SGLangEngine(RayActor):
raise e
except Exception as e:
print(f"Error flushing cache: {e}")
time.sleep(1)
continue
else:
raise TimeoutError("Timeout while flushing cache.")
def shutdown(self):
requests.post(
@@ -230,6 +255,18 @@ class SGLangEngine(RayActor):
},
)
def destroy_weights_update_group(self, group_name):
try:
return self._make_request(
"destroy_weights_update_group",
{
"group_name": group_name,
},
)
except:
# catch the case there the engine is just created and does not have the group.
pass
def update_weights_from_distributed(
self, names, dtypes, shapes, group_name, flush_cache=False, weight_version: Optional[str] = None
):
+4 -3
View File
@@ -260,9 +260,10 @@ class XTunerTrainRayActor(TrainRayActor):
if self.args.debug_train_only or self.args.debug_rollout_only:
return
if not self.connected:
self.connected = True
rollout_engines, rollout_engine_lock = ray.get(self.rollout_manager.get_rollout_engines_and_lock.remote())
rollout_engines, rollout_engine_lock, num_new_engines = ray.get(
self.rollout_manager.get_rollout_engines_and_lock.remote()
)
if num_new_engines > 0:
self.weight_updator.connect_rollout_engines(rollout_engines, rollout_engine_lock)
dist.barrier(group=get_gloo_group())
-3
View File
@@ -175,9 +175,6 @@ def create_rollout_manager(args, pg, wandb_run_id):
if args.rollout_global_dataset:
ray.get(rollout_manager.load.remote(args.start_rollout_id - 1))
# TODO: extract this to single function
rollout_engines, rollout_engine_lock = ray.get(rollout_manager.get_rollout_engines_and_lock.remote())
# calculate num_rollout from num_epoch
num_rollout_per_epoch = None
if args.num_rollout is None:
+126 -39
View File
@@ -1,6 +1,7 @@
import logging
import multiprocessing
import random
import threading
import time
from pathlib import Path
from typing import List, Union
@@ -30,6 +31,7 @@ class RolloutManager:
def __init__(self, args, pg, wandb_run_id):
self.args = args
self.pg = pg
_start_router(args)
init_wandb_secondary(args, wandb_run_id)
init_http_client(args)
@@ -44,17 +46,23 @@ class RolloutManager:
print(f"import {self.args.rollout_function_path} as generate_rollout function.")
print(f"import {self.args.eval_function_path} as eval_generate_rollout function.")
self.all_rollout_engines = _create_rollout_engines(args, pg)
nodes_per_engine = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node)
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
num_engines = args.rollout_num_gpus // num_gpu_per_engine
self.all_rollout_engines = [None] * num_engines
self.num_new_engines = init_rollout_engines(args, pg, self.all_rollout_engines)
self.nodes_per_engine = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node)
# when doing multi-node serving, we will only send request to node-0 for each engine.
self.rollout_engines = self.all_rollout_engines[::nodes_per_engine]
self.rollout_engine_lock = Lock.options(
num_cpus=1,
num_gpus=0,
).remote()
self.rollout_engines = self.all_rollout_engines[:: self.nodes_per_engine]
self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote()
# fault tolerance
self._health_monitor_thread = None
self._health_monitor_stop_event = None
self._health_check_interval = getattr(args, "rollout_health_check_interval", 10.0)
self._health_check_timeout = getattr(args, "rollout_health_check_timeout", 5.0)
def get_rollout_engines_and_lock(self):
return self.rollout_engines, self.rollout_engine_lock
return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines
def get_num_rollout_per_epoch(self):
assert self.args.rollout_global_dataset
@@ -62,17 +70,25 @@ class RolloutManager:
def generate(self, rollout_id):
self.rollout_id = rollout_id
monitor_started = self._start_health_monitor()
start_time = time.time()
data = self._get_rollout_data()
self._save_debug_rollout_data(data)
_log_rollout_data(rollout_id, self.args, data, time.time() - start_time)
data = self._convert_samples_to_train_data(data)
return Box(ray.put(data))
try:
data = self._get_rollout_data()
self._save_debug_rollout_data(data)
_log_rollout_data(rollout_id, self.args, data, time.time() - start_time)
data = self._convert_samples_to_train_data(data)
return Box(ray.put(data))
finally:
if monitor_started:
self._stop_health_monitor()
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]
def eval(self, rollout_id):
if self.args.debug_train_only:
# if debug train only, we don't generate evaluation data
return
# TODO: add fault tolerance to eval
data = self.eval_generate_rollout(self.args, rollout_id, self.data_source, evaluation=True)
_log_eval_rollout_data(rollout_id, self.args, data)
@@ -88,6 +104,65 @@ class RolloutManager:
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
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):
if self.args.load_debug_rollout_data:
data = torch.load(
@@ -199,12 +274,13 @@ class RolloutManager:
return train_data
def _create_rollout_engines(args, pg):
def init_rollout_engines(args, pg, all_rollout_engines):
if args.debug_train_only:
return []
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
num_engines = args.rollout_num_gpus // num_gpu_per_engine
assert len(all_rollout_engines) == num_engines
pg, reordered_bundle_indices = pg
@@ -212,6 +288,9 @@ def _create_rollout_engines(args, pg):
rollout_engines = []
for i in range(num_engines):
if all_rollout_engines[i] is not None:
continue
num_gpus = 0.2
num_cpus = num_gpus
@@ -221,20 +300,26 @@ def _create_rollout_engines(args, pg):
placement_group_bundle_index=reordered_bundle_indices[i * num_gpu_per_engine],
)
rollout_engines.append(
RolloutRayActor.options(
num_cpus=num_cpus,
num_gpus=num_gpus,
scheduling_strategy=scheduling_strategy,
runtime_env={
"env_vars": {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST}
| {
"SGL_JIT_DEEPGEMM_PRECOMPILE": "false",
"SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
}
},
).remote(args, rank=i)
)
rollout_engine = RolloutRayActor.options(
num_cpus=num_cpus,
num_gpus=num_gpus,
scheduling_strategy=scheduling_strategy,
runtime_env={
"env_vars": {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST}
| {
"SGL_JIT_DEEPGEMM_PRECOMPILE": "false",
"SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
}
},
).remote(args, rank=i)
rollout_engines.append((i, rollout_engine))
all_rollout_engines[i] = rollout_engine
num_new_engines = len(rollout_engines)
if num_new_engines == 0:
return num_new_engines
# get ports
# there are 4 ports we need to allocate
@@ -246,9 +331,15 @@ def _create_rollout_engines(args, pg):
1, min(args.num_gpus_per_node, args.rollout_num_gpus) // args.rollout_num_gpus_per_engine
)
addr_and_ports = [{} for _ in range(num_engines)]
for rank, engine in enumerate(rollout_engines):
if rank % num_engines_per_node != 0:
visited_nodes = set()
for rank, engine in rollout_engines:
if rank // num_engines_per_node in visited_nodes:
continue
visited_nodes.add(rank // num_engines_per_node)
# TODO: currently when restarting engines, we will set port for all engines on this node starting with this rank.
# e.g. for 8 gpus, if we are restarting engine on gpu 3, we will set port for engine 3,4,5,6,7 on this node.
num_engines_on_this_node = num_engines_per_node - (rank % num_engines_per_node)
def get_addr_and_ports():
# use small ports to prevent ephemeral port between 32768 and 65536.
@@ -273,7 +364,7 @@ def _create_rollout_engines(args, pg):
get_addr, get_port = get_addr_and_ports()
for i in range(num_engines_per_node):
for i in range(num_engines_on_this_node):
addr_and_ports[rank + i]["port"] = get_port()
addr_and_ports[rank + i]["nccl_port"] = get_port()
@@ -285,23 +376,19 @@ def _create_rollout_engines(args, pg):
for i in range(num_node_per_engine):
addr_and_ports[rank + i]["dist_init_addr"] = dist_init_addr
else:
for i in range(num_engines_per_node):
for i in range(num_engines_on_this_node):
addr_and_ports[rank + i]["dist_init_addr"] = f"{get_addr()}:{get_port(6 + args.sglang_dp_size)}"
for i in range(num_engines):
for i, _ in rollout_engines:
for key in ["port", "nccl_port", "dist_init_addr"]:
assert key in addr_and_ports[i], f"Engine {i} {key} is not set."
print(f"Ports for engine {i}: {addr_and_ports[i]}")
# TODO: don't ray.get here to overlap train actor init with rollout engine init.
# somehow if we don't sync here, the --debug-rollout-only mode will crash.
init_handles = [engine.init.remote(**ports) for engine, ports in zip(rollout_engines, addr_and_ports)]
init_handles = [engine.init.remote(**(addr_and_ports[rank])) for rank, engine in rollout_engines]
ray.get(init_handles)
if args.offload:
ray.get([engine.release_memory_occupation.remote() for engine in rollout_engines])
return rollout_engines
return num_new_engines
def _start_router(args):
+12
View File
@@ -216,6 +216,18 @@ def get_slime_extra_args_provider(add_custom_arguments=None):
"This is used to shuffle the prompts and also for the random sampling of the prompts."
),
)
parser.add_argument(
"--rollout-health-check-interval",
type=float,
default=10.0,
help="Interval in seconds between rollout engine /health_generate checks during generate/eval.",
)
parser.add_argument(
"--rollout-health-check-timeout",
type=float,
default=5.0,
help="Timeout in seconds to wait for a rollout engine /health_generate response before killing it.",
)
# sampling
parser.add_argument(
+1
View File
@@ -20,6 +20,7 @@ def train(args):
actor_model.set_rollout_manager(rollout_manager)
if args.offload:
ray.get(rollout_manager.offload.remote())
ray.get(rollout_manager.onload.remote(tags=[GPU_MEMORY_TYPE_WEIGHTS]))
# always update weight first so that sglang has the loaded weights from training.