mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[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:
@@ -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())
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
):
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user