mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[feat] (1/n) fault tolerance and slime rollout structure (#723)
This commit is contained in:
@@ -77,7 +77,7 @@ jobs:
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
info: [{"num_gpus": 0, "test_file": "fast"}]
|
||||
info: [{"num_gpus": 0, "test_file": "fast"}, {"num_gpus": 0, "test_file": "utils/test_sglang_config.py"}]
|
||||
defaults:
|
||||
run:
|
||||
working-directory: ${{ github.workspace }}
|
||||
@@ -401,7 +401,7 @@ jobs:
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
info: [{"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_async_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen3_0.6B_fsdp_colocated_2xGPU.py"}]
|
||||
info: [{"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_async_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen3_0.6B_fsdp_colocated_2xGPU.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config_mixed_offload.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config_mixed_offload_ft.py"}]
|
||||
defaults:
|
||||
run:
|
||||
working-directory: ${{ github.workspace }}
|
||||
@@ -1328,7 +1328,7 @@ jobs:
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
info: [{"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_4B_fsdp_true_on_policy.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_vl_4B_fsdp.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_0.6B_fsdp_distributed.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_0.6B_megatron_fsdp_align.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_quick_start_glm4_9B.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B.py", "use_deepep": "1", "use_fp8_rollout": "1"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B_r3.py", "use_deepep": "1", "use_fp8_rollout": "1"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B_r3.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_4B_ppo.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_moonlight_16B_A3B.py"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_moonlight_16B_A3B_r3.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_mimo_7B_mtp_only_grad.py"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_glm47_flash_r3_mtp.py"}, {"num_gpus": 8, "test_file": "e2e/lora/test_lora_qwen2.5_0.5B.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_async_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen3_0.6B_fsdp_colocated_2xGPU.py"}, {"num_gpus": 8, "test_file": "e2e/precision/test_qwen3_0.6B_parallel_check.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_qwen3_4B_ckpt.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_qwen3_4B_ckpt.py --async-save"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_glm47_flash_ckpt.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_glm47_flash_ckpt.py --async-save"}, {"num_gpus": 8, "test_file": "e2e/long/test_qwen2.5_0.5B_gsm8k.py"}, {"num_gpus": 8, "test_file": "e2e/long/test_qwen2.5_0.5B_gsm8k_async.py"}]
|
||||
info: [{"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_4B_fsdp_true_on_policy.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_vl_4B_fsdp.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_0.6B_fsdp_distributed.py"}, {"num_gpus": 8, "test_file": "e2e/fsdp/test_qwen3_0.6B_megatron_fsdp_align.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_quick_start_glm4_9B.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B.py", "use_deepep": "1", "use_fp8_rollout": "1"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B_r3.py", "use_deepep": "1", "use_fp8_rollout": "1"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_30B_A3B_r3.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_qwen3_4B_ppo.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_moonlight_16B_A3B.py"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_moonlight_16B_A3B_r3.py"}, {"num_gpus": 8, "test_file": "e2e/megatron/test_mimo_7B_mtp_only_grad.py"}, {"enable_eval": "0", "num_gpus": 8, "test_file": "e2e/megatron/test_glm47_flash_r3_mtp.py"}, {"num_gpus": 8, "test_file": "e2e/lora/test_lora_qwen2.5_0.5B.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_async_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen2.5_0.5B_gsm8k_short.py"}, {"num_gpus": 8, "test_file": "e2e/short/test_qwen3_0.6B_fsdp_colocated_2xGPU.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config_mixed_offload.py"}, {"num_gpus": 8, "test_file": "e2e/sglang_config/test_sglang_config_mixed_offload_ft.py"}, {"num_gpus": 8, "test_file": "e2e/precision/test_qwen3_0.6B_parallel_check.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_qwen3_4B_ckpt.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_qwen3_4B_ckpt.py --async-save"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_glm47_flash_ckpt.py"}, {"num_gpus": 8, "test_file": "e2e/ckpt/test_glm47_flash_ckpt.py --async-save"}, {"num_gpus": 8, "test_file": "e2e/long/test_qwen2.5_0.5B_gsm8k.py"}, {"num_gpus": 8, "test_file": "e2e/long/test_qwen2.5_0.5B_gsm8k_async.py"}]
|
||||
defaults:
|
||||
run:
|
||||
working-directory: ${{ github.workspace }}
|
||||
|
||||
@@ -549,14 +549,19 @@ class FSDPTrainRayActor(TrainRayActor):
|
||||
if self.args.debug_train_only or self.args.debug_rollout_only:
|
||||
return
|
||||
|
||||
rollout_engines, rollout_engine_lock, num_new_engines = ray.get(
|
||||
self.rollout_manager.get_rollout_engines_and_lock.remote()
|
||||
rollout_engines, rollout_engine_lock, num_new_engines, engine_gpu_counts, engine_gpu_offsets = ray.get(
|
||||
self.rollout_manager.get_updatable_engines_and_lock.remote()
|
||||
)
|
||||
if num_new_engines > 0:
|
||||
self.weight_updater.connect_rollout_engines(rollout_engines, rollout_engine_lock)
|
||||
self.weight_updater.connect_rollout_engines(
|
||||
rollout_engines,
|
||||
rollout_engine_lock,
|
||||
engine_gpu_counts=engine_gpu_counts,
|
||||
engine_gpu_offsets=engine_gpu_offsets,
|
||||
)
|
||||
dist.barrier(group=get_gloo_group())
|
||||
if dist.get_rank() == 0:
|
||||
ray.get(self.rollout_manager.clear_num_new_engines.remote())
|
||||
ray.get(self.rollout_manager.clear_updatable_num_new_engines.remote())
|
||||
|
||||
self.weight_updater.update_weights()
|
||||
|
||||
|
||||
@@ -40,6 +40,8 @@ class UpdateWeight(abc.ABC):
|
||||
self,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
rollout_engine_lock: ActorHandle | None,
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
engine_gpu_offsets: Sequence[int] | None = None,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
@@ -92,6 +94,8 @@ class UpdateWeightFromTensor(UpdateWeight):
|
||||
self,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
rollout_engine_lock: ActorHandle | None,
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
engine_gpu_offsets: Sequence[int] | None = None,
|
||||
) -> None:
|
||||
"""Attach rollout engines and create per-engine IPC (Gloo) groups.
|
||||
|
||||
@@ -186,6 +190,8 @@ class UpdateWeightFromDistributed(UpdateWeight):
|
||||
self,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
rollout_engine_lock: ActorHandle | None,
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
engine_gpu_offsets: Sequence[int] | None = None,
|
||||
) -> None:
|
||||
"""On rank 0, initialize a temporary NCCL group for parameter broadcast."""
|
||||
self.rollout_engines = rollout_engines
|
||||
|
||||
@@ -480,21 +480,26 @@ class MegatronTrainRayActor(TrainRayActor):
|
||||
|
||||
if self.args.use_fault_tolerance:
|
||||
if dist.get_rank() == 0:
|
||||
ray.get(self.rollout_manager.recover_rollout_engines.remote())
|
||||
ray.get(self.rollout_manager.recover_updatable_engines.remote())
|
||||
dist.barrier(group=get_gloo_group())
|
||||
|
||||
rollout_engines, rollout_engine_lock, num_new_engines = ray.get(
|
||||
self.rollout_manager.get_rollout_engines_and_lock.remote()
|
||||
rollout_engines, rollout_engine_lock, num_new_engines, engine_gpu_counts, engine_gpu_offsets = ray.get(
|
||||
self.rollout_manager.get_updatable_engines_and_lock.remote()
|
||||
)
|
||||
|
||||
if self.args.offload_train:
|
||||
reload_process_groups()
|
||||
|
||||
if num_new_engines > 0:
|
||||
self.weight_updater.connect_rollout_engines(rollout_engines, rollout_engine_lock)
|
||||
self.weight_updater.connect_rollout_engines(
|
||||
rollout_engines,
|
||||
rollout_engine_lock,
|
||||
engine_gpu_counts=engine_gpu_counts,
|
||||
engine_gpu_offsets=engine_gpu_offsets,
|
||||
)
|
||||
dist.barrier(group=get_gloo_group())
|
||||
if dist.get_rank() == 0:
|
||||
ray.get(self.rollout_manager.clear_num_new_engines.remote())
|
||||
ray.get(self.rollout_manager.clear_updatable_num_new_engines.remote())
|
||||
|
||||
if self.args.offload_train and is_lora_enabled(self.args):
|
||||
# For LoRA, we must resume() to restore GPU memory backing for adapter
|
||||
|
||||
@@ -44,13 +44,18 @@ class UpdateWeightFromDistributed:
|
||||
self._model_update_groups = None
|
||||
|
||||
def connect_rollout_engines(
|
||||
self, rollout_engines: Sequence[ActorHandle], rollout_engine_lock: ActorHandle
|
||||
self,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
rollout_engine_lock: ActorHandle,
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
engine_gpu_offsets: Sequence[int] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Create NCCL "miles-pp_{pp_rank}" if PP source (DP=TP=0). Lock prevents concurrent broadcasts.
|
||||
"""
|
||||
self.rollout_engines = rollout_engines
|
||||
self.rollout_engine_lock = rollout_engine_lock
|
||||
self._engine_gpu_counts = engine_gpu_counts
|
||||
|
||||
# For TP:
|
||||
# 1. AllGather parameters to rank 0
|
||||
@@ -246,28 +251,40 @@ class UpdateWeightFromDistributed:
|
||||
|
||||
|
||||
def connect_rollout_engines_from_distributed(
|
||||
args: Namespace, group_name: str, rollout_engines: Sequence[ActorHandle]
|
||||
args: Namespace,
|
||||
group_name: str,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
) -> dist.ProcessGroup:
|
||||
"""
|
||||
Create NCCL group: training rank 0 + all engine GPUs. Blocks until joined.
|
||||
|
||||
``engine_gpu_counts`` gives the number of GPUs per engine. When engines
|
||||
have heterogeneous TP sizes (e.g. prefill TP=2, decode TP=4), each engine
|
||||
occupies a different number of ranks in the NCCL group.
|
||||
"""
|
||||
if engine_gpu_counts is None:
|
||||
engine_gpu_counts = [args.rollout_num_gpus_per_engine] * len(rollout_engines)
|
||||
master_address = ray._private.services.get_node_ip_address()
|
||||
with socket.socket() as sock:
|
||||
sock.bind(("", 0))
|
||||
master_port = sock.getsockname()[1]
|
||||
world_size = len(rollout_engines) * args.rollout_num_gpus_per_engine + 1
|
||||
world_size = sum(engine_gpu_counts) + 1
|
||||
|
||||
refs = [
|
||||
engine.init_weights_update_group.remote(
|
||||
master_address,
|
||||
master_port,
|
||||
i * args.rollout_num_gpus_per_engine + 1,
|
||||
world_size,
|
||||
group_name,
|
||||
backend="nccl",
|
||||
refs = []
|
||||
rank_cursor = 1
|
||||
for i, engine in enumerate(rollout_engines):
|
||||
refs.append(
|
||||
engine.init_weights_update_group.remote(
|
||||
master_address,
|
||||
master_port,
|
||||
rank_cursor,
|
||||
world_size,
|
||||
group_name,
|
||||
backend="nccl",
|
||||
)
|
||||
)
|
||||
for i, engine in enumerate(rollout_engines)
|
||||
]
|
||||
rank_cursor += engine_gpu_counts[i]
|
||||
model_update_groups = init_process_group(
|
||||
backend="nccl",
|
||||
init_method=f"tcp://{master_address}:{master_port}",
|
||||
|
||||
@@ -76,21 +76,42 @@ class UpdateWeightFromTensor:
|
||||
self._model_update_groups = None
|
||||
|
||||
def connect_rollout_engines(
|
||||
self, rollout_engines: Sequence[ActorHandle], rollout_engine_lock: ActorHandle
|
||||
self,
|
||||
rollout_engines: Sequence[ActorHandle],
|
||||
rollout_engine_lock: ActorHandle,
|
||||
engine_gpu_counts: Sequence[int] | None = None,
|
||||
engine_gpu_offsets: Sequence[int] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Split colocated/distributed engines. Global source rank (DP=TP=PP=0) creates NCCL
|
||||
for distributed. Map ranks to colocated IPC engines.
|
||||
"""
|
||||
self.rollout_engines = rollout_engines
|
||||
colocate_engine_nums = (
|
||||
self.args.actor_num_nodes * self.args.actor_num_gpus_per_node // self.args.rollout_num_gpus_per_engine
|
||||
)
|
||||
|
||||
if engine_gpu_counts is None:
|
||||
engine_gpu_counts = [self.args.rollout_num_gpus_per_engine] * len(rollout_engines)
|
||||
if engine_gpu_offsets is None:
|
||||
# Fallback: assume engines are densely packed (no placeholder gaps).
|
||||
engine_gpu_offsets = []
|
||||
offset = 0
|
||||
for c in engine_gpu_counts:
|
||||
engine_gpu_offsets.append(offset)
|
||||
offset += c
|
||||
|
||||
# Compute colocated engine count: engines whose GPUs fall within actor GPU range.
|
||||
total_actor_gpus = self.args.actor_num_nodes * self.args.actor_num_gpus_per_node
|
||||
colocate_engine_nums = 0
|
||||
for gpu_offset, gpu_count in zip(engine_gpu_offsets, engine_gpu_counts, strict=True):
|
||||
if gpu_offset + gpu_count > total_actor_gpus:
|
||||
break
|
||||
colocate_engine_nums += 1
|
||||
|
||||
self.use_distribute = len(rollout_engines) > colocate_engine_nums
|
||||
|
||||
if self.use_distribute:
|
||||
self.rollout_engines = rollout_engines[:colocate_engine_nums]
|
||||
self.distributed_rollout_engines = rollout_engines[colocate_engine_nums:]
|
||||
distributed_gpu_counts = engine_gpu_counts[colocate_engine_nums:]
|
||||
self._is_distributed_src_rank = (
|
||||
mpu.get_data_parallel_rank(with_context_parallel=True) == 0
|
||||
and mpu.get_tensor_model_parallel_rank() == 0
|
||||
@@ -102,16 +123,47 @@ class UpdateWeightFromTensor:
|
||||
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
|
||||
self.args,
|
||||
self._group_name,
|
||||
self.distributed_rollout_engines,
|
||||
engine_gpu_counts=distributed_gpu_counts,
|
||||
)
|
||||
|
||||
# Here we assume the gpu id of rollout engines and train actors are the same.
|
||||
colocate_gpu_offsets = engine_gpu_offsets[:colocate_engine_nums]
|
||||
colocate_gpu_counts = engine_gpu_counts[:colocate_engine_nums]
|
||||
|
||||
# Determine whether this rank is covered by any colocated engine.
|
||||
all_colocated_ranks = set()
|
||||
for offset, count in zip(colocate_gpu_offsets, colocate_gpu_counts, strict=True):
|
||||
all_colocated_ranks.update(range(offset, offset + count))
|
||||
rank_has_engine = dist.get_rank() in all_colocated_ranks
|
||||
|
||||
# Create IPC Gloo gather groups matching actual engine layout.
|
||||
# Re-create on first call or when engine layout changes (placeholder ranks
|
||||
# that had a group from __init__ but no actual engine need to be reset).
|
||||
if rank_has_engine:
|
||||
if self._ipc_gather_group is None:
|
||||
for i in range(colocate_engine_nums):
|
||||
group_ranks = list(
|
||||
range(colocate_gpu_offsets[i], colocate_gpu_offsets[i] + colocate_gpu_counts[i])
|
||||
)
|
||||
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 = colocate_gpu_offsets[i]
|
||||
else:
|
||||
# Ranks not covered by any engine (e.g. placeholder GPU slots)
|
||||
self._ipc_gather_group = None
|
||||
self._ipc_gather_src = None
|
||||
|
||||
# Map training ranks to colocated engine actors.
|
||||
self._ipc_engine = None
|
||||
for i, engine in enumerate(self.rollout_engines):
|
||||
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))
|
||||
if dist.get_rank() in group_ranks:
|
||||
start = colocate_gpu_offsets[i]
|
||||
end = start + colocate_gpu_counts[i]
|
||||
if start <= dist.get_rank() < end:
|
||||
self._ipc_engine = engine
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -234,6 +286,11 @@ def _send_to_colocated_engine(
|
||||
lora_name: str | None = None,
|
||||
lora_loaded: bool = False,
|
||||
) -> tuple[list[ObjectRef], Any]:
|
||||
# Placeholder ranks (GPU slots reserved but no engine) have no gather group.
|
||||
# gather_object is only collective among group members, so we skip entirely.
|
||||
if ipc_gather_group is None:
|
||||
return [], None
|
||||
|
||||
is_lora = lora_config is not None
|
||||
long_live_tensors = []
|
||||
|
||||
|
||||
@@ -39,6 +39,7 @@ def add_sglang_arguments(parser):
|
||||
|
||||
skipped_args = [
|
||||
"model_path",
|
||||
"config",
|
||||
"trust_remote_code",
|
||||
"random_seed",
|
||||
# memory
|
||||
@@ -108,6 +109,20 @@ def add_sglang_arguments(parser):
|
||||
ServerArgs.add_cli_args(parser)
|
||||
parser.add_argument = old_add_argument
|
||||
|
||||
parser.add_argument(
|
||||
"--sglang-config",
|
||||
type=str,
|
||||
default=None,
|
||||
help=(
|
||||
"Path to a YAML config for SGLang engine deployment. "
|
||||
"Defines server_groups with worker_type (regular/prefill/decode/placeholder), "
|
||||
"num_gpus per group, and optional per-group 'overrides' dict of "
|
||||
"ServerArgs field names that override the base --sglang-* CLI args. "
|
||||
"Placeholder groups reserve GPU slots without creating engines. "
|
||||
"Mutually exclusive with --prefill-num-servers."
|
||||
),
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
"""Configuration dataclasses for SGLang engine deployment."""
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
|
||||
import yaml
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ServerGroupConfig:
|
||||
"""Configuration for a single server group.
|
||||
|
||||
Attributes:
|
||||
worker_type: One of "regular", "prefill", "decode", or "placeholder".
|
||||
"placeholder" reserves GPU slots without creating engines.
|
||||
num_gpus: Total number of GPUs for this group.
|
||||
num_gpus_per_engine: GPUs per engine for this group. Overrides the
|
||||
model-level or global ``--rollout-num-gpus-per-engine``.
|
||||
overrides: Optional dict of SGLang ``ServerArgs`` field overrides.
|
||||
These are applied on top of the base CLI ``--sglang-*``
|
||||
arguments in ``_compute_server_args``.
|
||||
"""
|
||||
|
||||
worker_type: str
|
||||
num_gpus: int
|
||||
num_gpus_per_engine: int | None = None
|
||||
overrides: dict = dataclasses.field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
valid_types = {"regular", "prefill", "decode", "placeholder"}
|
||||
assert (
|
||||
self.worker_type in valid_types
|
||||
), f"Invalid worker_type '{self.worker_type}', must be one of {valid_types}"
|
||||
assert self.num_gpus > 0, f"num_gpus must be > 0, got {self.num_gpus}"
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ModelConfig:
|
||||
"""Configuration for a single model deployment.
|
||||
|
||||
Attributes:
|
||||
name: Unique name for this model (e.g. "actor", "reward").
|
||||
model_path: HF checkpoint path. Falls back to ``args.hf_checkpoint``.
|
||||
num_gpus_per_engine: Default GPUs per engine for all groups in this
|
||||
model. Individual groups can override.
|
||||
server_groups: Server group configurations for this model.
|
||||
update_weights: Whether this model receives weight updates from
|
||||
training. Set to ``False`` for frozen models
|
||||
(reference, reward, etc.). When ``None`` (default),
|
||||
automatically inferred in ``resolve()``: ``True`` if
|
||||
model_path matches ``args.hf_checkpoint``, ``False``
|
||||
otherwise.
|
||||
"""
|
||||
|
||||
name: str
|
||||
model_path: str | None = None
|
||||
num_gpus_per_engine: int | None = None
|
||||
server_groups: list[ServerGroupConfig] = dataclasses.field(default_factory=list)
|
||||
update_weights: bool | None = None
|
||||
|
||||
def resolve(self, args) -> None:
|
||||
"""Resolve per-group defaults from model-level then args-level values."""
|
||||
default_gpus_per_engine = self.num_gpus_per_engine or args.rollout_num_gpus_per_engine
|
||||
default_model_path = self.model_path or args.hf_checkpoint
|
||||
for g in self.server_groups:
|
||||
if g.num_gpus_per_engine is None:
|
||||
g.num_gpus_per_engine = default_gpus_per_engine
|
||||
if "model_path" not in g.overrides:
|
||||
g.overrides["model_path"] = default_model_path
|
||||
|
||||
if self.server_groups:
|
||||
model_paths = {g.overrides["model_path"] for g in self.server_groups}
|
||||
assert len(model_paths) == 1, (
|
||||
f"Model '{self.name}' has server groups with different model_path values: "
|
||||
f"{model_paths}. All server groups within a model must use the same model_path."
|
||||
)
|
||||
effective_model_path = model_paths.pop()
|
||||
else:
|
||||
effective_model_path = default_model_path
|
||||
|
||||
if self.update_weights is None:
|
||||
if effective_model_path != args.hf_checkpoint:
|
||||
logger.warning(
|
||||
f"Model '{self.name}' uses model_path='{effective_model_path}' which differs "
|
||||
f"from hf_checkpoint='{args.hf_checkpoint}'. Defaulting update_weights to False. "
|
||||
f"Set update_weights explicitly in the config to suppress this warning."
|
||||
)
|
||||
self.update_weights = False
|
||||
else:
|
||||
self.update_weights = True
|
||||
|
||||
@property
|
||||
def has_pd_disaggregation(self) -> bool:
|
||||
return any(g.worker_type in ("prefill", "decode") for g in self.server_groups)
|
||||
|
||||
@property
|
||||
def total_num_gpus(self) -> int:
|
||||
return sum(g.num_gpus for g in self.server_groups)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class SglangConfig:
|
||||
"""Configuration for SGLang engine deployment.
|
||||
|
||||
Loaded from ``--sglang-config`` YAML file.
|
||||
|
||||
**Config format**::
|
||||
|
||||
sglang:
|
||||
- name: actor
|
||||
model_path: /path/to/actor
|
||||
update_weights: true # receives training weight updates (default)
|
||||
num_gpus_per_engine: 2
|
||||
server_groups:
|
||||
- worker_type: prefill
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 2
|
||||
- worker_type: decode
|
||||
num_gpus: 8
|
||||
num_gpus_per_engine: 4
|
||||
- name: ref
|
||||
model_path: /path/to/ref
|
||||
update_weights: false # frozen, no weight updates
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
|
||||
Each model gets its own router. ``placeholder`` groups reserve GPU
|
||||
slots without creating engines. ``overrides`` are ``ServerArgs``
|
||||
field names applied on top of the base ``--sglang-*`` CLI args.
|
||||
|
||||
Set ``update_weights: false`` for frozen models (reference, reward,
|
||||
etc.) that should not receive weight updates from training.
|
||||
|
||||
.. note::
|
||||
|
||||
``engine_groups`` is accepted as a backward-compatible alias for
|
||||
``server_groups`` in the YAML config.
|
||||
"""
|
||||
|
||||
models: list[ModelConfig]
|
||||
|
||||
@staticmethod
|
||||
def from_yaml(path: str) -> "SglangConfig":
|
||||
with open(path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
assert "sglang" in data, (
|
||||
f"sglang config must have a 'sglang' key, got {list(data.keys())}. "
|
||||
f"Wrap your server_groups inside a model entry under 'sglang'."
|
||||
)
|
||||
models = []
|
||||
for m in data["sglang"]:
|
||||
raw_groups = m.get("server_groups") or m.get("engine_groups") or []
|
||||
groups = [ServerGroupConfig(**g) for g in raw_groups]
|
||||
models.append(
|
||||
ModelConfig(
|
||||
name=m["name"],
|
||||
model_path=m.get("model_path"),
|
||||
num_gpus_per_engine=m.get("num_gpus_per_engine"),
|
||||
server_groups=groups,
|
||||
update_weights=m.get("update_weights"),
|
||||
)
|
||||
)
|
||||
return SglangConfig(models=models)
|
||||
|
||||
@staticmethod
|
||||
def from_prefill_num_servers(args) -> "SglangConfig":
|
||||
"""Build a config equivalent to the legacy --prefill-num-servers flag."""
|
||||
total_gpus = args.rollout_num_gpus
|
||||
prefill_gpus = args.prefill_num_servers * args.rollout_num_gpus_per_engine
|
||||
decode_gpus = total_gpus - prefill_gpus
|
||||
assert decode_gpus > 0, f"No decode GPUs: total {total_gpus}, prefill {prefill_gpus}"
|
||||
return SglangConfig(
|
||||
models=[
|
||||
ModelConfig(
|
||||
name="default",
|
||||
server_groups=[
|
||||
ServerGroupConfig(worker_type="prefill", num_gpus=prefill_gpus),
|
||||
ServerGroupConfig(worker_type="decode", num_gpus=decode_gpus),
|
||||
],
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
@property
|
||||
def has_pd_disaggregation(self) -> bool:
|
||||
return any(m.has_pd_disaggregation for m in self.models)
|
||||
|
||||
@property
|
||||
def total_num_gpus(self) -> int:
|
||||
return sum(m.total_num_gpus for m in self.models)
|
||||
@@ -108,15 +108,34 @@ def _wait_server_healthy(base_url, api_key, is_process_alive):
|
||||
|
||||
|
||||
class SGLangEngine(RayActor):
|
||||
def __init__(self, args, rank: int, worker_type: str = "regular", base_gpu_id: int | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
rank: int,
|
||||
worker_type: str = "regular",
|
||||
base_gpu_id: int | None = None,
|
||||
sglang_overrides: dict | None = None,
|
||||
num_gpus_per_engine: int | None = None,
|
||||
):
|
||||
self.args = args
|
||||
self.rank = rank
|
||||
self.worker_type = worker_type
|
||||
self.base_gpu_id = base_gpu_id
|
||||
self.sglang_overrides = sglang_overrides or {}
|
||||
self.num_gpus_per_engine = num_gpus_per_engine
|
||||
|
||||
def init(self, dist_init_addr, port, nccl_port, host=None, disaggregation_bootstrap_port=None):
|
||||
self.router_ip = self.args.sglang_router_ip
|
||||
self.router_port = self.args.sglang_router_port
|
||||
def init(
|
||||
self,
|
||||
dist_init_addr,
|
||||
port,
|
||||
nccl_port,
|
||||
host=None,
|
||||
disaggregation_bootstrap_port=None,
|
||||
router_ip=None,
|
||||
router_port=None,
|
||||
):
|
||||
self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip
|
||||
self.router_port = router_port if router_port is not None else self.args.sglang_router_port
|
||||
|
||||
host = host or get_host_info()[1]
|
||||
|
||||
@@ -144,6 +163,8 @@ class SGLangEngine(RayActor):
|
||||
self.worker_type,
|
||||
disaggregation_bootstrap_port,
|
||||
base_gpu_id=self.base_gpu_id,
|
||||
sglang_overrides=self.sglang_overrides,
|
||||
num_gpus_per_engine=self.num_gpus_per_engine,
|
||||
)
|
||||
|
||||
self.node_rank = server_args_dict["node_rank"]
|
||||
@@ -390,6 +411,17 @@ class SGLangEngine(RayActor):
|
||||
def check_weights(self, action: str):
|
||||
return self._make_request("weights_checker", {"action": action})
|
||||
|
||||
def update_weights_from_disk(self, model_path: str, load_format: str | None = None):
|
||||
"""Reload weights from *model_path* without restarting the engine.
|
||||
|
||||
Used for non-updatable (frozen) models that overlap with megatron:
|
||||
after offload, weights are restored from disk instead of CPU cache.
|
||||
"""
|
||||
payload = {"model_path": model_path}
|
||||
if load_format is not None:
|
||||
payload["load_format"] = load_format
|
||||
return self._make_request("update_weights_from_disk", payload)
|
||||
|
||||
def init_weights_update_group(self, master_address, master_port, rank_offset, world_size, group_name, backend):
|
||||
return self._make_request(
|
||||
"init_weights_update_group",
|
||||
@@ -517,8 +549,11 @@ def _compute_server_args(
|
||||
worker_type: str = "regular",
|
||||
disaggregation_bootstrap_port: int | None = None,
|
||||
base_gpu_id: int | None = None,
|
||||
sglang_overrides: dict | None = None,
|
||||
num_gpus_per_engine: int | None = None,
|
||||
):
|
||||
nnodes = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node)
|
||||
_gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine
|
||||
nnodes = max(1, _gpus_per_engine // args.num_gpus_per_node)
|
||||
node_rank = rank % nnodes
|
||||
base = base_gpu_id if base_gpu_id is not None else get_base_gpu_id(args, rank)
|
||||
base = _to_local_gpu_id(base)
|
||||
@@ -538,7 +573,7 @@ def _compute_server_args(
|
||||
"gpu_id_step": 1,
|
||||
"base_gpu_id": base,
|
||||
# parallel
|
||||
"tp_size": args.rollout_num_gpus_per_engine,
|
||||
"tp_size": _gpus_per_engine,
|
||||
"dp_size": args.sglang_dp_size,
|
||||
"pp_size": args.sglang_pp_size,
|
||||
"ep_size": args.sglang_ep_size,
|
||||
@@ -548,6 +583,9 @@ def _compute_server_args(
|
||||
"enable_draft_weights_cpu_backup": True,
|
||||
}
|
||||
|
||||
if sglang_overrides:
|
||||
kwargs.update(sglang_overrides)
|
||||
|
||||
if worker_type == "prefill":
|
||||
kwargs["disaggregation_mode"] = "prefill"
|
||||
kwargs["load_balance_method"] = "round_robin"
|
||||
|
||||
+563
-185
@@ -1,6 +1,8 @@
|
||||
import dataclasses
|
||||
import itertools
|
||||
import logging
|
||||
import multiprocessing
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from pathlib import Path
|
||||
@@ -12,6 +14,7 @@ import torch
|
||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH, GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_WEIGHTS
|
||||
|
||||
from miles.backends.sglang_utils.sglang_config import ModelConfig, ServerGroupConfig, SglangConfig
|
||||
from miles.backends.sglang_utils.sglang_engine import SGLangEngine
|
||||
from miles.rollout.base_types import (
|
||||
RolloutFnConstructorInput,
|
||||
@@ -43,6 +46,282 @@ logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ServerGroup / RolloutServer abstractions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ServerGroup:
|
||||
"""A group of homogeneous SGLang engines with the same configuration.
|
||||
|
||||
All engines in a group share the same tp_size / nodes_per_engine / pg.
|
||||
A RolloutServer may contain multiple ServerGroups (e.g. prefill vs decode
|
||||
in PD disaggregation).
|
||||
"""
|
||||
|
||||
args: Any
|
||||
pg: Any # (placement_group, reordered_bundle_indices, reordered_gpu_ids)
|
||||
all_engines: list
|
||||
num_gpus_per_engine: int
|
||||
num_new_engines: int
|
||||
worker_type: str = "regular" # "regular", "prefill", or "decode"
|
||||
rank_offset: int = 0
|
||||
gpu_offset: int = 0
|
||||
sglang_overrides: dict = dataclasses.field(default_factory=dict)
|
||||
needs_offload: bool = False
|
||||
model_path: str | None = None
|
||||
router_ip: str | None = None
|
||||
router_port: int | None = None
|
||||
|
||||
@property
|
||||
def nodes_per_engine(self):
|
||||
return max(1, self.num_gpus_per_engine // self.args.num_gpus_per_node)
|
||||
|
||||
@property
|
||||
def engines(self):
|
||||
"""Node-0 engines only (for multi-node serving)."""
|
||||
return self.all_engines[:: self.nodes_per_engine]
|
||||
|
||||
def start_engines(self, port_cursors: dict[int, int] | None = None) -> tuple[list, dict[int, int]]:
|
||||
"""Create Ray actors, allocate ports, and fire ``engine.init()`` without waiting.
|
||||
|
||||
Returns ``(init_handles, port_cursors)`` where *init_handles* is a list
|
||||
of Ray ObjectRefs and *port_cursors* maps node index -> next free port.
|
||||
"""
|
||||
if port_cursors is None:
|
||||
port_cursors = {}
|
||||
if self.args.debug_train_only or self.worker_type == "placeholder":
|
||||
self.num_new_engines = 0
|
||||
return [], port_cursors
|
||||
|
||||
num_gpu_per_engine = min(self.num_gpus_per_engine, self.args.num_gpus_per_node)
|
||||
|
||||
pg, reordered_bundle_indices, reordered_gpu_ids = self.pg
|
||||
|
||||
RolloutRayActor = ray.remote(SGLangEngine)
|
||||
|
||||
rollout_engines = []
|
||||
for i in range(len(self.all_engines)):
|
||||
if self.all_engines[i] is not None:
|
||||
continue
|
||||
|
||||
global_rank = self.rank_offset + i
|
||||
num_gpus = 0.2
|
||||
num_cpus = num_gpus
|
||||
|
||||
gpu_index = self.gpu_offset + i * num_gpu_per_engine
|
||||
base_gpu_id = int(reordered_gpu_ids[gpu_index])
|
||||
|
||||
scheduling_strategy = PlacementGroupSchedulingStrategy(
|
||||
placement_group=pg,
|
||||
placement_group_capture_child_tasks=True,
|
||||
placement_group_bundle_index=reordered_bundle_indices[gpu_index],
|
||||
)
|
||||
|
||||
env_vars = {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST} | {
|
||||
key: os.environ.get(key, default_val)
|
||||
for key, default_val in {
|
||||
"SGLANG_JIT_DEEPGEMM_PRECOMPILE": "false",
|
||||
"SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
|
||||
"SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
|
||||
"SGLANG_MEMORY_SAVER_CUDA_GRAPH": "true",
|
||||
"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_FALLBACK_VARIANT": "true",
|
||||
"SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION": "false",
|
||||
"SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "false",
|
||||
}.items()
|
||||
}
|
||||
|
||||
rollout_engine = RolloutRayActor.options(
|
||||
num_cpus=num_cpus,
|
||||
num_gpus=num_gpus,
|
||||
scheduling_strategy=scheduling_strategy,
|
||||
runtime_env={
|
||||
"env_vars": env_vars,
|
||||
},
|
||||
).remote(
|
||||
self.args,
|
||||
rank=global_rank,
|
||||
worker_type=self.worker_type,
|
||||
base_gpu_id=base_gpu_id,
|
||||
sglang_overrides=self.sglang_overrides,
|
||||
num_gpus_per_engine=self.num_gpus_per_engine,
|
||||
)
|
||||
|
||||
rollout_engines.append((global_rank, rollout_engine))
|
||||
self.all_engines[i] = rollout_engine
|
||||
|
||||
self.num_new_engines = len(rollout_engines)
|
||||
|
||||
if self.num_new_engines == 0:
|
||||
return [], port_cursors
|
||||
|
||||
if self.args.rollout_external:
|
||||
addr_and_ports = _allocate_rollout_engine_addr_and_ports_external(
|
||||
args=self.args, rollout_engines=rollout_engines
|
||||
)
|
||||
else:
|
||||
base_port = max(port_cursors.values()) if port_cursors else 15000
|
||||
addr_and_ports, port_cursors = _allocate_rollout_engine_addr_and_ports_normal(
|
||||
args=self.args,
|
||||
rollout_engines=rollout_engines,
|
||||
worker_type=self.worker_type,
|
||||
num_gpus_per_engine=self.num_gpus_per_engine,
|
||||
rank_offset=self.rank_offset,
|
||||
base_port=base_port,
|
||||
)
|
||||
|
||||
init_handles = [
|
||||
engine.init.remote(
|
||||
**(addr_and_ports[rank]),
|
||||
router_ip=self.router_ip,
|
||||
router_port=self.router_port,
|
||||
)
|
||||
for rank, engine in rollout_engines
|
||||
]
|
||||
return init_handles, port_cursors
|
||||
|
||||
def offload(self):
|
||||
if not self.needs_offload:
|
||||
return []
|
||||
return [engine.release_memory_occupation.remote() for engine in self.engines if engine is not None]
|
||||
|
||||
def onload(self, tags: list[str] | None = None):
|
||||
if not self.needs_offload:
|
||||
return []
|
||||
return [engine.resume_memory_occupation.remote(tags=tags) for engine in self.engines if engine is not None]
|
||||
|
||||
def onload_weights_from_disk(self):
|
||||
"""Reload weights from ``model_path`` for non-updatable groups."""
|
||||
if not self.needs_offload or not self.model_path:
|
||||
return []
|
||||
return [
|
||||
engine.update_weights_from_disk.remote(self.model_path) for engine in self.engines if engine is not None
|
||||
]
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class RolloutServer:
|
||||
"""A model served behind a shared router, with one or more server groups.
|
||||
|
||||
Each RolloutServer represents one model deployed behind a single router.
|
||||
"""
|
||||
|
||||
server_groups: list[ServerGroup]
|
||||
router_ip: str | None = None
|
||||
router_port: int | None = None
|
||||
model_name: str = "default"
|
||||
update_weights: bool = True
|
||||
|
||||
@property
|
||||
def engines(self):
|
||||
"""All node-0 engines across all groups."""
|
||||
return [e for g in self.server_groups for e in g.engines]
|
||||
|
||||
@property
|
||||
def all_engines(self):
|
||||
return [e for g in self.server_groups for e in g.all_engines]
|
||||
|
||||
@property
|
||||
def num_new_engines(self):
|
||||
return sum(g.num_new_engines for g in self.server_groups)
|
||||
|
||||
@num_new_engines.setter
|
||||
def num_new_engines(self, value):
|
||||
for g in self.server_groups:
|
||||
g.num_new_engines = value
|
||||
|
||||
@property
|
||||
def engine_gpu_counts(self) -> list[int]:
|
||||
"""Per-engine GPU count for all node-0 engines, parallel to ``engines``."""
|
||||
return [g.num_gpus_per_engine for g in self.server_groups for _ in g.engines]
|
||||
|
||||
@property
|
||||
def engine_gpu_offsets(self) -> list[int]:
|
||||
offsets = []
|
||||
for g in self.server_groups:
|
||||
for j in range(len(g.engines)):
|
||||
offsets.append(g.gpu_offset + j * g.num_gpus_per_engine)
|
||||
return offsets
|
||||
|
||||
@property
|
||||
def nodes_per_engine(self):
|
||||
values = {g.nodes_per_engine for g in self.server_groups}
|
||||
if len(values) != 1:
|
||||
raise ValueError(f"Heterogeneous nodes_per_engine across groups: {values}")
|
||||
return values.pop()
|
||||
|
||||
def recover(self):
|
||||
"""Recover dead engines across all active groups, overlapping init."""
|
||||
dead_per_group = [[i for i, engine in enumerate(g.all_engines) if engine is None] for g in self.server_groups]
|
||||
|
||||
all_handles = []
|
||||
port_cursors: dict[int, int] = {}
|
||||
for g in self.server_groups:
|
||||
handles, port_cursors = g.start_engines(port_cursors)
|
||||
all_handles.extend(handles)
|
||||
if all_handles:
|
||||
ray.get(all_handles)
|
||||
|
||||
release_handles = []
|
||||
updatable_new_engines = []
|
||||
non_updatable_groups_engines: list[tuple[str, list]] = []
|
||||
for g, dead_indices in zip(self.server_groups, dead_per_group, strict=True):
|
||||
logger.info(f"Recovered {g.num_new_engines} dead rollout engines (worker_type={g.worker_type})")
|
||||
assert g.num_new_engines == len(dead_indices), "num_new_engines does not match dead_indices length"
|
||||
if g.needs_offload and dead_indices:
|
||||
new_engines = [g.all_engines[i] for i in dead_indices]
|
||||
release_handles.extend(engine.release_memory_occupation.remote() for engine in new_engines)
|
||||
if self.update_weights:
|
||||
updatable_new_engines.extend(new_engines)
|
||||
elif g.model_path:
|
||||
non_updatable_groups_engines.append((g.model_path, new_engines))
|
||||
|
||||
if release_handles:
|
||||
ray.get(release_handles)
|
||||
all_resume_engines = updatable_new_engines[:]
|
||||
for _model_path, engines in non_updatable_groups_engines:
|
||||
all_resume_engines.extend(engines)
|
||||
if all_resume_engines:
|
||||
ray.get(
|
||||
[
|
||||
engine.resume_memory_occupation.remote(tags=[GPU_MEMORY_TYPE_WEIGHTS])
|
||||
for engine in all_resume_engines
|
||||
]
|
||||
)
|
||||
|
||||
def offload(self):
|
||||
handles = []
|
||||
for g in self.server_groups:
|
||||
handles.extend(g.offload())
|
||||
return ray.get(handles) if handles else []
|
||||
|
||||
def onload(self, tags: list[str] | None = None):
|
||||
handles = []
|
||||
for g in self.server_groups:
|
||||
handles.extend(g.onload(tags))
|
||||
return ray.get(handles) if handles else []
|
||||
|
||||
def onload_weights(self):
|
||||
handles = []
|
||||
for g in self.server_groups:
|
||||
if not g.needs_offload:
|
||||
continue
|
||||
handles.extend(g.onload(tags=[GPU_MEMORY_TYPE_WEIGHTS]))
|
||||
return ray.get(handles) if handles else []
|
||||
|
||||
def onload_kv(self):
|
||||
handles = []
|
||||
for g in self.server_groups:
|
||||
handles.extend(g.onload(tags=[GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_CUDA_GRAPH]))
|
||||
return ray.get(handles) if handles else []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RolloutManager
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@ray.remote
|
||||
class RolloutManager:
|
||||
"""The class to run rollout and convert rollout data to training data."""
|
||||
@@ -50,12 +329,10 @@ class RolloutManager:
|
||||
def __init__(self, args, pg):
|
||||
configure_logger()
|
||||
|
||||
self.args = args
|
||||
self.pg = pg
|
||||
_start_router(args)
|
||||
self.args = args
|
||||
# TODO make args immutable
|
||||
init_tracking(args, primary=False, router_addr=f"http://{args.sglang_router_ip}:{args.sglang_router_port}")
|
||||
init_http_client(args)
|
||||
|
||||
data_source_cls = load_function(self.args.data_source_path)
|
||||
self.data_source = data_source_cls(args)
|
||||
@@ -80,21 +357,21 @@ class RolloutManager:
|
||||
logger.info(f"import {self.args.eval_function_path} as eval_generate_rollout function.")
|
||||
|
||||
if self.args.debug_train_only:
|
||||
self.all_rollout_engines = []
|
||||
self.servers: dict[str, RolloutServer] = {}
|
||||
else:
|
||||
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)
|
||||
init_http_client(args)
|
||||
self.servers = start_rollout_servers(args, pg)
|
||||
self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote()
|
||||
self.rollout_id = -1
|
||||
|
||||
self._metric_checker = MetricChecker.maybe_create(args)
|
||||
self._health_monitor = None
|
||||
if self.args.use_fault_tolerance:
|
||||
self._health_monitor = RolloutHealthMonitor(self, args)
|
||||
self._health_monitor.start() # Start the monitor thread (in paused state)
|
||||
self._health_monitors = []
|
||||
if not self.args.debug_train_only and self.args.use_fault_tolerance:
|
||||
for srv in self.servers.values():
|
||||
for group in srv.server_groups:
|
||||
monitor = RolloutHealthMonitor(group, args)
|
||||
monitor.start()
|
||||
self._health_monitors.append(monitor)
|
||||
self._ci_fault_injection_pending = self.args.ci_test # Flag for CI fault injection
|
||||
|
||||
def _try_ci_fault_injection(self):
|
||||
@@ -105,11 +382,11 @@ class RolloutManager:
|
||||
# Only inject fault once
|
||||
self._ci_fault_injection_pending = False
|
||||
|
||||
if self.all_rollout_engines and self.all_rollout_engines[0]:
|
||||
if self.server and self.server.server_groups[0].all_engines and self.server.server_groups[0].all_engines[0]:
|
||||
logger.info("CI Fault Injection: Simulating crash on engine 0 during generate")
|
||||
try:
|
||||
# This will cause the ray actor to exit
|
||||
self.all_rollout_engines[0].simulate_crash.remote()
|
||||
self.server.server_groups[0].all_engines[0].simulate_crash.remote()
|
||||
# Wait for health monitor to detect the crash and mark engine as None
|
||||
# health_check_interval + health_check_timeout + buffer
|
||||
wait_time = self.args.rollout_health_check_interval + self.args.rollout_health_check_timeout + 5
|
||||
@@ -121,17 +398,35 @@ class RolloutManager:
|
||||
def dispose(self):
|
||||
if self._metric_checker is not None:
|
||||
self._metric_checker.dispose()
|
||||
if self._health_monitor is not None:
|
||||
self._health_monitor.stop()
|
||||
for monitor in self._health_monitors:
|
||||
monitor.stop()
|
||||
|
||||
@property
|
||||
def server(self) -> RolloutServer | None:
|
||||
"""Default server (first model). For backward compatibility."""
|
||||
if not self.servers:
|
||||
return None
|
||||
return next(iter(self.servers.values()))
|
||||
|
||||
def _get_updatable_server(self) -> RolloutServer | None:
|
||||
for srv in self.servers.values():
|
||||
if srv.update_weights:
|
||||
return srv
|
||||
return None
|
||||
|
||||
# TODO maybe rename "rollout_engines" and "all_rollout_engines" later
|
||||
@property
|
||||
def rollout_engines(self):
|
||||
# when doing multi-node serving, we will only send request to node-0 for each engine.
|
||||
return self.all_rollout_engines[:: self.nodes_per_engine]
|
||||
"""All node-0 engines across all servers / models."""
|
||||
return [e for srv in self.servers.values() for e in srv.engines]
|
||||
|
||||
def get_rollout_engines_and_lock(self):
|
||||
return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines
|
||||
def get_updatable_engines_and_lock(self):
|
||||
"""Return engines eligible for weight updates."""
|
||||
srv = self._get_updatable_server()
|
||||
engines = srv.engines if srv else []
|
||||
gpu_counts = srv.engine_gpu_counts if srv else []
|
||||
gpu_offsets = srv.engine_gpu_offsets if srv else []
|
||||
num_new = srv.num_new_engines if srv else 0
|
||||
return engines, self.rollout_engine_lock, num_new, gpu_counts, gpu_offsets
|
||||
|
||||
def get_num_rollout_per_epoch(self):
|
||||
assert self.args.rollout_global_dataset
|
||||
@@ -175,57 +470,64 @@ class RolloutManager:
|
||||
|
||||
def offload(self, tags: list[str] | None = None):
|
||||
self.health_monitoring_pause()
|
||||
return ray.get(
|
||||
[
|
||||
if tags is not None:
|
||||
handles = [
|
||||
engine.release_memory_occupation.remote(tags=tags)
|
||||
for engine in self.rollout_engines
|
||||
if engine is not None
|
||||
]
|
||||
)
|
||||
return ray.get(handles) if handles else []
|
||||
for srv in self.servers.values():
|
||||
srv.offload()
|
||||
|
||||
def onload(self, tags: list[str] | None = None):
|
||||
return ray.get(
|
||||
[
|
||||
engine.resume_memory_occupation.remote(tags=tags)
|
||||
for engine in self.rollout_engines
|
||||
if engine is not None
|
||||
]
|
||||
)
|
||||
for srv in self.servers.values():
|
||||
srv.onload(tags)
|
||||
|
||||
def health_monitoring_pause(self):
|
||||
if self.args.use_fault_tolerance and self._health_monitor is not None:
|
||||
self._health_monitor.pause()
|
||||
def health_monitoring_pause(self) -> None:
|
||||
for monitor in self._health_monitors:
|
||||
monitor.pause()
|
||||
|
||||
def health_monitoring_resume(self):
|
||||
if self.args.use_fault_tolerance and self._health_monitor is not None:
|
||||
self._health_monitor.resume()
|
||||
def health_monitoring_resume(self) -> None:
|
||||
for monitor in self._health_monitors:
|
||||
monitor.resume()
|
||||
|
||||
def onload_weights(self):
|
||||
self.onload(tags=[GPU_MEMORY_TYPE_WEIGHTS])
|
||||
for srv in self.servers.values():
|
||||
srv.onload_weights()
|
||||
|
||||
def onload_kv(self):
|
||||
self.onload(tags=[GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_CUDA_GRAPH])
|
||||
for srv in self.servers.values():
|
||||
srv.onload_kv()
|
||||
|
||||
def recover_rollout_engines(self):
|
||||
"""Restart any dead rollout engines and update num_new_engines for update_weights detection."""
|
||||
def recover_updatable_engines(self):
|
||||
"""Restart any dead rollout engines and update num_new_engines for update_weights detection.
|
||||
|
||||
Recovers the updatable model (the one that receives weight
|
||||
updates from training).
|
||||
"""
|
||||
self.health_monitoring_pause()
|
||||
if self.rollout_id == -1:
|
||||
return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines
|
||||
srv = self._get_updatable_server()
|
||||
if self.rollout_id == -1 or srv is None:
|
||||
engines = srv.engines if srv else []
|
||||
gpu_counts = srv.engine_gpu_counts if srv else []
|
||||
gpu_offsets = srv.engine_gpu_offsets if srv else []
|
||||
return engines, self.rollout_engine_lock, (srv.num_new_engines if srv else 0), gpu_counts, gpu_offsets
|
||||
|
||||
dead_indices = [i for i, engine in enumerate(self.all_rollout_engines) if engine is None]
|
||||
self.num_new_engines = init_rollout_engines(self.args, self.pg, self.all_rollout_engines)
|
||||
logger.info(f"Recovered {self.num_new_engines} dead rollout engines")
|
||||
assert self.num_new_engines == len(dead_indices), "num_new_engines does not match dead_indices length"
|
||||
if self.args.offload_rollout and dead_indices:
|
||||
new_engines = [self.all_rollout_engines[i] for i in dead_indices]
|
||||
ray.get([engine.release_memory_occupation.remote() for engine in new_engines])
|
||||
ray.get([engine.resume_memory_occupation.remote(tags=[GPU_MEMORY_TYPE_WEIGHTS]) for engine in new_engines])
|
||||
srv.recover()
|
||||
return (
|
||||
srv.engines,
|
||||
self.rollout_engine_lock,
|
||||
srv.num_new_engines,
|
||||
srv.engine_gpu_counts,
|
||||
srv.engine_gpu_offsets,
|
||||
)
|
||||
|
||||
return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines
|
||||
|
||||
def clear_num_new_engines(self):
|
||||
def clear_updatable_num_new_engines(self):
|
||||
# when fault tolerance is not enabled, we need to manually clear num_new_engines after update_weights
|
||||
self.num_new_engines = 0
|
||||
srv = self._get_updatable_server()
|
||||
if srv:
|
||||
srv.num_new_engines = 0
|
||||
|
||||
def check_weights(self, action: str):
|
||||
return ray.get([engine.check_weights.remote(action=action) for engine in self.rollout_engines])
|
||||
@@ -409,7 +711,7 @@ class RolloutManager:
|
||||
if samples[0].train_metadata is not None:
|
||||
train_data["metadata"] = [sample.train_metadata for sample in samples]
|
||||
|
||||
if samples[0].multimodal_train_inputs is not None:
|
||||
if any(sample.multimodal_train_inputs is not None for sample in samples):
|
||||
train_data["multimodal_train_inputs"] = [sample.multimodal_train_inputs for sample in samples]
|
||||
|
||||
if "teacher_log_probs" in samples[0].__dict__:
|
||||
@@ -474,138 +776,63 @@ class RolloutManager:
|
||||
return rollout_data_refs
|
||||
|
||||
|
||||
def init_rollout_engines(args, pg, all_rollout_engines):
|
||||
if args.debug_train_only:
|
||||
return 0
|
||||
|
||||
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
|
||||
if args.prefill_num_servers is not None:
|
||||
prefill_num_servers = args.prefill_num_servers * args.rollout_num_gpus_per_engine // num_gpu_per_engine
|
||||
assert (
|
||||
num_engines > prefill_num_servers
|
||||
), f"num_engines {num_engines} should be larger than prefill_num_servers {prefill_num_servers}"
|
||||
|
||||
pg, reordered_bundle_indices, reordered_gpu_ids = pg
|
||||
|
||||
RolloutRayActor = ray.remote(SGLangEngine)
|
||||
|
||||
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
|
||||
|
||||
# Get the base GPU ID from placement group
|
||||
base_gpu_id = int(reordered_gpu_ids[i * num_gpu_per_engine])
|
||||
|
||||
scheduling_strategy = PlacementGroupSchedulingStrategy(
|
||||
placement_group=pg,
|
||||
placement_group_capture_child_tasks=True,
|
||||
placement_group_bundle_index=reordered_bundle_indices[i * num_gpu_per_engine],
|
||||
)
|
||||
|
||||
env_vars = {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST} | {
|
||||
"SGL_JIT_DEEPGEMM_PRECOMPILE": "false",
|
||||
"SGLANG_JIT_DEEPGEMM_PRECOMPILE": "false",
|
||||
"SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
|
||||
"SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true",
|
||||
"SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK": "false",
|
||||
"SGLANG_MEMORY_SAVER_CUDA_GRAPH": "true",
|
||||
"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_FALLBACK_VARIANT": "true",
|
||||
"SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION": "false",
|
||||
"SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "false",
|
||||
}
|
||||
|
||||
worker_type = "regular"
|
||||
if args.prefill_num_servers is not None:
|
||||
if i < prefill_num_servers:
|
||||
worker_type = "prefill"
|
||||
else:
|
||||
worker_type = "decode"
|
||||
|
||||
rollout_engine = RolloutRayActor.options(
|
||||
num_cpus=num_cpus,
|
||||
num_gpus=num_gpus,
|
||||
scheduling_strategy=scheduling_strategy,
|
||||
runtime_env={
|
||||
"env_vars": env_vars,
|
||||
},
|
||||
).remote(args, rank=i, worker_type=worker_type, base_gpu_id=base_gpu_id)
|
||||
|
||||
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
|
||||
|
||||
if args.rollout_external:
|
||||
addr_and_ports = _allocate_rollout_engine_addr_and_ports_external(args=args, rollout_engines=rollout_engines)
|
||||
else:
|
||||
addr_and_ports = _allocate_rollout_engine_addr_and_ports_normal(
|
||||
args=args, num_engines=num_engines, rollout_engines=rollout_engines
|
||||
)
|
||||
|
||||
# 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(**(addr_and_ports[rank])) for rank, engine in rollout_engines]
|
||||
ray.get(init_handles)
|
||||
|
||||
return num_new_engines
|
||||
# ---------------------------------------------------------------------------
|
||||
# Port allocation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _allocate_rollout_engine_addr_and_ports_external(args, rollout_engines):
|
||||
addr_and_ports = []
|
||||
addr_and_ports = {}
|
||||
for rank, _ in rollout_engines:
|
||||
addr = args.rollout_external_engine_addrs[rank]
|
||||
[host, port] = addr.split(":")
|
||||
addr_and_ports.append(
|
||||
dict(
|
||||
dist_init_addr=addr,
|
||||
nccl_port=None,
|
||||
host=host,
|
||||
port=int(port),
|
||||
)
|
||||
addr_and_ports[rank] = dict(
|
||||
dist_init_addr=addr,
|
||||
nccl_port=None,
|
||||
host=host,
|
||||
port=int(port),
|
||||
)
|
||||
return addr_and_ports
|
||||
|
||||
|
||||
def _allocate_rollout_engine_addr_and_ports_normal(*, args, num_engines, rollout_engines):
|
||||
def _allocate_rollout_engine_addr_and_ports_normal(
|
||||
*,
|
||||
args,
|
||||
rollout_engines,
|
||||
worker_type="regular",
|
||||
num_gpus_per_engine=None,
|
||||
rank_offset=0,
|
||||
base_port=15000,
|
||||
):
|
||||
# get ports
|
||||
# there are 4 ports we need to allocate
|
||||
# 1. server port
|
||||
# 2. nccl port
|
||||
# 3. dist_init_addr port
|
||||
# 4. other ports for dp_attention, which is of size 4 + dp_size
|
||||
num_engines_per_node = max(
|
||||
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)]
|
||||
_gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine
|
||||
num_engines_per_node = max(1, args.num_gpus_per_node // _gpus_per_engine)
|
||||
addr_and_ports: dict[int, dict] = {}
|
||||
|
||||
# Calculate prefill limit to identify prefill engines
|
||||
prefill_limit = 0
|
||||
if args.prefill_num_servers is not None:
|
||||
num_gpu_per_engine = min(args.rollout_num_gpus_per_engine, args.num_gpus_per_node)
|
||||
prefill_limit = args.prefill_num_servers * args.rollout_num_gpus_per_engine // num_gpu_per_engine
|
||||
# Track per-node port cursors so that different server groups (called
|
||||
# sequentially) never race for the same ports on a given node.
|
||||
node_port_cursor: dict[int, int] = {}
|
||||
|
||||
visited_nodes = set()
|
||||
for rank, engine in rollout_engines:
|
||||
if rank // num_engines_per_node in visited_nodes:
|
||||
local_rank = rank - rank_offset
|
||||
node_index = local_rank // num_engines_per_node
|
||||
if node_index in visited_nodes:
|
||||
continue
|
||||
visited_nodes.add(rank // num_engines_per_node)
|
||||
visited_nodes.add(node_index)
|
||||
# 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)
|
||||
num_engines_on_this_node = num_engines_per_node - (local_rank % num_engines_per_node)
|
||||
|
||||
def get_addr_and_ports(engine):
|
||||
def get_addr_and_ports(engine, node_idx):
|
||||
# use small ports to prevent ephemeral port between 32768 and 65536.
|
||||
# also, ray uses port 10002-19999, thus we avoid near-10002 to avoid racing condition
|
||||
start_port = 15000
|
||||
start_port = node_port_cursor.get(node_idx, base_port)
|
||||
|
||||
def port(consecutive=1):
|
||||
nonlocal start_port
|
||||
@@ -616,6 +843,7 @@ def _allocate_rollout_engine_addr_and_ports_normal(*, args, num_engines, rollout
|
||||
)
|
||||
)
|
||||
start_port = port + consecutive
|
||||
node_port_cursor[node_idx] = start_port
|
||||
return port
|
||||
|
||||
def addr():
|
||||
@@ -624,23 +852,24 @@ def _allocate_rollout_engine_addr_and_ports_normal(*, args, num_engines, rollout
|
||||
|
||||
return addr, port
|
||||
|
||||
get_addr, get_port = get_addr_and_ports(engine)
|
||||
get_addr, get_port = get_addr_and_ports(engine, node_index)
|
||||
|
||||
for i in range(num_engines_on_this_node):
|
||||
current_rank = rank + i
|
||||
addr_and_ports.setdefault(current_rank, {})
|
||||
addr_and_ports[current_rank]["host"] = get_addr()
|
||||
addr_and_ports[current_rank]["port"] = get_port()
|
||||
addr_and_ports[current_rank]["nccl_port"] = get_port()
|
||||
|
||||
if args.prefill_num_servers is not None and current_rank < prefill_limit:
|
||||
if worker_type == "prefill":
|
||||
addr_and_ports[current_rank]["disaggregation_bootstrap_port"] = get_port()
|
||||
|
||||
if args.rollout_num_gpus_per_engine > args.num_gpus_per_node:
|
||||
num_node_per_engine = args.rollout_num_gpus_per_engine // args.num_gpus_per_node
|
||||
if rank % num_node_per_engine == 0:
|
||||
# this is the first node in the engine, we need to allocate the dist_init_addr port
|
||||
if _gpus_per_engine > args.num_gpus_per_node:
|
||||
num_node_per_engine = _gpus_per_engine // args.num_gpus_per_node
|
||||
if local_rank % num_node_per_engine == 0:
|
||||
dist_init_addr = f"{get_addr()}:{get_port(30 + args.sglang_dp_size)}"
|
||||
for i in range(num_node_per_engine):
|
||||
addr_and_ports.setdefault(rank + i, {})
|
||||
addr_and_ports[rank + i]["dist_init_addr"] = dist_init_addr
|
||||
else:
|
||||
for i in range(num_engines_on_this_node):
|
||||
@@ -651,23 +880,40 @@ def _allocate_rollout_engine_addr_and_ports_normal(*, args, num_engines, rollout
|
||||
assert key in addr_and_ports[i], f"Engine {i} {key} is not set."
|
||||
logger.info(f"Ports for engine {i}: {addr_and_ports[i]}")
|
||||
|
||||
return addr_and_ports
|
||||
return addr_and_ports, node_port_cursor
|
||||
|
||||
|
||||
def _start_router(args):
|
||||
"""start sgl router and miles router"""
|
||||
if args.sglang_router_ip is not None:
|
||||
return
|
||||
# ---------------------------------------------------------------------------
|
||||
# Router + server bootstrap
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
args.sglang_router_ip = _wrap_ipv6(get_host_info()[1])
|
||||
if args.sglang_router_port is None:
|
||||
args.sglang_router_port = find_available_port(random.randint(3000, 4000))
|
||||
|
||||
def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool = False) -> tuple[str, int]:
|
||||
"""Start sgl router or miles router and return (router_ip, router_port).
|
||||
|
||||
If ``args.sglang_router_ip`` is already set and ``force_new`` is False,
|
||||
skip launching and return the existing values.
|
||||
"""
|
||||
if not force_new and args.sglang_router_ip is not None:
|
||||
return args.sglang_router_ip, args.sglang_router_port
|
||||
|
||||
router_ip = _wrap_ipv6(get_host_info()[1])
|
||||
if force_new:
|
||||
router_port = find_available_port(random.randint(3000, 4000))
|
||||
else:
|
||||
router_port = args.sglang_router_port
|
||||
if router_port is None:
|
||||
router_port = find_available_port(random.randint(3000, 4000))
|
||||
|
||||
if args.use_miles_router:
|
||||
assert args.prefill_num_servers is None, "miles router does not support prefill_num_servers."
|
||||
import copy
|
||||
|
||||
assert not has_pd_disaggregation, "miles router does not support PD disaggregation."
|
||||
from miles.router.router import run_router
|
||||
|
||||
router_args = args
|
||||
router_args = copy.copy(args)
|
||||
router_args.sglang_router_ip = router_ip
|
||||
router_args.sglang_router_port = router_port
|
||||
|
||||
else:
|
||||
from sglang_router.launch_router import RouterArgs
|
||||
@@ -675,13 +921,13 @@ def _start_router(args):
|
||||
from miles.utils.http_utils import run_router
|
||||
|
||||
router_args = RouterArgs.from_cli_args(args, use_router_prefix=True)
|
||||
router_args.host = args.sglang_router_ip
|
||||
router_args.port = args.sglang_router_port
|
||||
router_args.host = router_ip
|
||||
router_args.port = router_port
|
||||
router_args.prometheus_port = find_available_port(random.randint(4000, 5000))
|
||||
router_args.log_level = "warn"
|
||||
router_args.request_timeout_secs = args.sglang_router_request_timeout_secs
|
||||
|
||||
if args.prefill_num_servers is not None:
|
||||
if has_pd_disaggregation:
|
||||
router_args.pd_disaggregation = True
|
||||
|
||||
logger.info(f"Launch router with args: {router_args}")
|
||||
@@ -690,12 +936,144 @@ def _start_router(args):
|
||||
target=run_router,
|
||||
args=(router_args,),
|
||||
)
|
||||
process.daemon = True # Set the process as a daemon
|
||||
process.daemon = True
|
||||
process.start()
|
||||
# Wait 3 seconds
|
||||
time.sleep(3)
|
||||
assert process.is_alive()
|
||||
logger.info(f"Router launched at {args.sglang_router_ip}:{args.sglang_router_port}")
|
||||
logger.info(f"Router launched at {router_ip}:{router_port}")
|
||||
return router_ip, router_port
|
||||
|
||||
|
||||
def _compute_rollout_offset(args) -> int:
|
||||
"""Offset (in PG bundle slots) where rollout GPUs start."""
|
||||
if args.debug_train_only or args.debug_rollout_only or args.colocate:
|
||||
return 0
|
||||
if getattr(args, "critic_train_only", False):
|
||||
return args.critic_num_nodes * args.critic_num_gpus_per_node
|
||||
offset = args.actor_num_nodes * args.actor_num_gpus_per_node
|
||||
if getattr(args, "use_critic", False):
|
||||
offset += args.critic_num_nodes * args.critic_num_gpus_per_node
|
||||
return offset
|
||||
|
||||
|
||||
def _compute_megatron_num_gpus(args) -> int:
|
||||
"""Total number of megatron (actor + critic) GPU slots in the placement group."""
|
||||
if getattr(args, "debug_rollout_only", False):
|
||||
return 0
|
||||
if getattr(args, "critic_train_only", False):
|
||||
return args.critic_num_nodes * args.critic_num_gpus_per_node
|
||||
num = args.actor_num_nodes * args.actor_num_gpus_per_node
|
||||
if getattr(args, "use_critic", False):
|
||||
num += args.critic_num_nodes * args.critic_num_gpus_per_node
|
||||
return num
|
||||
|
||||
|
||||
def start_rollout_servers(args, pg) -> dict[str, RolloutServer]:
|
||||
"""Start rollout servers: one per model, each with its own router.
|
||||
|
||||
Returns a dict mapping model name -> ``RolloutServer``.
|
||||
"""
|
||||
config = _resolve_sglang_config(args)
|
||||
|
||||
servers: dict[str, RolloutServer] = {}
|
||||
gpu_offset = 0
|
||||
engine_offset = 0
|
||||
|
||||
rollout_pg_offset = _compute_rollout_offset(args)
|
||||
megatron_num_gpus = _compute_megatron_num_gpus(args)
|
||||
|
||||
for model_idx, model_cfg in enumerate(config.models):
|
||||
model_cfg.resolve(args)
|
||||
|
||||
has_pd = model_cfg.has_pd_disaggregation
|
||||
router_ip, router_port = _start_router(args, has_pd_disaggregation=has_pd, force_new=(model_idx > 0))
|
||||
|
||||
if model_idx == 0:
|
||||
args.sglang_router_ip = router_ip
|
||||
args.sglang_router_port = router_port
|
||||
|
||||
server_groups: list[ServerGroup] = []
|
||||
all_init_handles: list = []
|
||||
port_cursors: dict[int, int] = {}
|
||||
|
||||
for group_cfg in model_cfg.server_groups:
|
||||
gpus_per_engine = group_cfg.num_gpus_per_engine
|
||||
num_gpu_per_engine_local = min(gpus_per_engine, args.num_gpus_per_node)
|
||||
num_engines = group_cfg.num_gpus // num_gpu_per_engine_local
|
||||
|
||||
group_abs_start = rollout_pg_offset + gpu_offset
|
||||
needs_offload = args.offload_rollout and group_abs_start < megatron_num_gpus
|
||||
overrides = dict(group_cfg.overrides)
|
||||
if args.offload_rollout and not needs_offload:
|
||||
overrides.setdefault("enable_memory_saver", False)
|
||||
logger.info(
|
||||
f"Engine group '{group_cfg.worker_type}' gpu_offset={gpu_offset} "
|
||||
f"(abs={group_abs_start}): needs_offload={needs_offload}"
|
||||
)
|
||||
|
||||
group = ServerGroup(
|
||||
args=args,
|
||||
pg=pg,
|
||||
all_engines=[None] * num_engines if group_cfg.worker_type != "placeholder" else [],
|
||||
num_gpus_per_engine=gpus_per_engine,
|
||||
num_new_engines=0,
|
||||
worker_type=group_cfg.worker_type,
|
||||
rank_offset=engine_offset,
|
||||
gpu_offset=gpu_offset,
|
||||
sglang_overrides=overrides,
|
||||
needs_offload=needs_offload,
|
||||
model_path=overrides.get("model_path", args.hf_checkpoint),
|
||||
router_ip=router_ip,
|
||||
router_port=router_port,
|
||||
)
|
||||
handles, port_cursors = group.start_engines(port_cursors)
|
||||
all_init_handles.extend(handles)
|
||||
server_groups.append(group)
|
||||
|
||||
engine_offset += num_engines
|
||||
gpu_offset += group_cfg.num_gpus
|
||||
|
||||
if all_init_handles:
|
||||
ray.get(all_init_handles)
|
||||
|
||||
servers[model_cfg.name] = RolloutServer(
|
||||
server_groups=server_groups,
|
||||
router_ip=router_ip,
|
||||
router_port=router_port,
|
||||
model_name=model_cfg.name,
|
||||
update_weights=model_cfg.update_weights,
|
||||
)
|
||||
|
||||
args.sglang_model_routers = {name: (srv.router_ip, srv.router_port) for name, srv in servers.items()}
|
||||
|
||||
return servers
|
||||
|
||||
|
||||
def _resolve_sglang_config(args) -> SglangConfig:
|
||||
"""Build a SglangConfig from args, choosing the right source."""
|
||||
if getattr(args, "sglang_config", None) is not None:
|
||||
config = SglangConfig.from_yaml(args.sglang_config)
|
||||
expected = args.rollout_num_gpus
|
||||
actual = config.total_num_gpus
|
||||
assert actual == expected, f"sglang_config total GPUs ({actual}) != rollout_num_gpus ({expected})"
|
||||
return config
|
||||
|
||||
if args.prefill_num_servers is not None:
|
||||
return SglangConfig.from_prefill_num_servers(args)
|
||||
|
||||
return SglangConfig(
|
||||
models=[
|
||||
ModelConfig(
|
||||
name="default",
|
||||
server_groups=[ServerGroupConfig(worker_type="regular", num_gpus=args.rollout_num_gpus)],
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Logging / metrics helpers (unchanged)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _log_eval_rollout_data(rollout_id, args, data, extra_metrics: dict[str, Any] | None = None):
|
||||
|
||||
@@ -2,6 +2,7 @@ import asyncio
|
||||
import copy
|
||||
import inspect
|
||||
import logging
|
||||
|
||||
from argparse import Namespace
|
||||
from collections.abc import Callable
|
||||
from contextlib import contextmanager
|
||||
@@ -26,11 +27,30 @@ from miles.utils.types import Sample
|
||||
|
||||
from .rm_hub import async_rm, batched_async_rm
|
||||
|
||||
__all__ = ["generate_rollout"]
|
||||
__all__ = ["generate_rollout", "get_model_url"]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_model_url(args: Namespace, model_name: str, endpoint: str = "/generate") -> str:
|
||||
"""Return the router URL for a named model.
|
||||
|
||||
Use this in custom rollout functions to route requests to a specific
|
||||
model when multiple models are deployed via ``--sglang-config``::
|
||||
|
||||
url = get_model_url(args, "ref", "/generate")
|
||||
resp = await post(url, json=payload)
|
||||
|
||||
Falls back to the default router if *model_name* is not found or
|
||||
``sglang_model_routers`` is not set.
|
||||
"""
|
||||
routers = getattr(args, "sglang_model_routers", None)
|
||||
if routers and model_name in routers:
|
||||
ip, port = routers[model_name]
|
||||
return f"http://{ip}:{port}{endpoint}"
|
||||
return f"http://{args.sglang_router_ip}:{args.sglang_router_port}{endpoint}"
|
||||
|
||||
|
||||
class GenerateState(metaclass=SingletonMeta):
|
||||
"""
|
||||
The global state for the generation process.
|
||||
@@ -175,6 +195,11 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A
|
||||
sample.response_length += len(new_response_tokens)
|
||||
sample.response += output["text"]
|
||||
|
||||
# When partial rollout and masking off policy is enabled, update the loss mask
|
||||
if sample.loss_mask is not None:
|
||||
assert args.partial_rollout and args.mask_offpolicy_in_partial_rollout
|
||||
sample.loss_mask += [1] * len(new_response_tokens)
|
||||
|
||||
if sample.rollout_log_probs is None:
|
||||
sample.rollout_log_probs = []
|
||||
sample.rollout_log_probs += new_response_log_probs
|
||||
@@ -303,7 +328,11 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]:
|
||||
urls = [worker["url"] for worker in response["workers"]]
|
||||
|
||||
logger.info(f"Abort request for {urls}")
|
||||
await asyncio.gather(*[post(f"{url}/abort_request", {"abort_all": True}) for url in urls])
|
||||
abort_tasks = [post(f"{url}/abort_request", {"abort_all": True}) for url in urls]
|
||||
abort_results = await asyncio.gather(*abort_tasks, return_exceptions=True)
|
||||
for url, result in zip(urls, abort_results, strict=False):
|
||||
if isinstance(result, Exception):
|
||||
logger.warning(f"Failed to abort worker at {url}: {result}")
|
||||
|
||||
# make sure all the pending tasks are finished
|
||||
count = 0
|
||||
|
||||
@@ -1872,6 +1872,14 @@ def miles_validate_args(args):
|
||||
args.prefill_num_servers is not None and args.rollout_external
|
||||
), "prefill_num_servers cannot be set when rollout_external is set."
|
||||
|
||||
assert not (
|
||||
getattr(args, "sglang_config", None) is not None and args.rollout_external
|
||||
), "sglang_config cannot be set when rollout_external is set."
|
||||
|
||||
assert not (
|
||||
getattr(args, "sglang_config", None) is not None and getattr(args, "prefill_num_servers", None) is not None
|
||||
), "sglang_config and prefill_num_servers are mutually exclusive. Use server_groups in the YAML config instead."
|
||||
|
||||
if args.qkv_format == "bshd":
|
||||
assert args.train_backend == "megatron", "bshd format is only supported for megatron backend."
|
||||
assert (
|
||||
|
||||
@@ -20,9 +20,8 @@ class RolloutHealthMonitor:
|
||||
- stop(): Stop the monitor thread completely (called during dispose)
|
||||
"""
|
||||
|
||||
def __init__(self, rollout_manager, args):
|
||||
# TODO may remove this dependency after refactoring
|
||||
self._rollout_manager = rollout_manager
|
||||
def __init__(self, server_group, args):
|
||||
self._server_group = server_group
|
||||
|
||||
self._thread = None
|
||||
self._stop_event = None
|
||||
@@ -39,7 +38,7 @@ class RolloutHealthMonitor:
|
||||
Returns:
|
||||
True if the monitor was started, False if there are no engines to monitor.
|
||||
"""
|
||||
if not self._rollout_manager.all_rollout_engines:
|
||||
if not self._server_group.all_engines:
|
||||
return False
|
||||
|
||||
if self._thread is not None:
|
||||
@@ -136,7 +135,7 @@ class RolloutHealthMonitor:
|
||||
break
|
||||
|
||||
def _run_health_checks(self) -> None:
|
||||
for rollout_engine_id, engine in enumerate(self._rollout_manager.rollout_engines):
|
||||
for rollout_engine_id, engine in enumerate(self._server_group.engines):
|
||||
if self._stop_event is not None and self._stop_event.is_set():
|
||||
break
|
||||
if self._pause_event is not None and self._pause_event.is_set():
|
||||
@@ -159,12 +158,12 @@ class RolloutHealthMonitor:
|
||||
logger.debug(f"Health check passed for rollout engine {rollout_engine_id}")
|
||||
|
||||
def _kill_engine(self, rollout_engine_id: int):
|
||||
logger.info(f"Killing engine group {rollout_engine_id}...")
|
||||
logger.info(f"Killing server 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,
|
||||
rollout_engine_id * self._server_group.nodes_per_engine,
|
||||
(rollout_engine_id + 1) * self._server_group.nodes_per_engine,
|
||||
):
|
||||
engine = self._rollout_manager.all_rollout_engines[i]
|
||||
engine = self._server_group.all_engines[i]
|
||||
if engine:
|
||||
logger.info(f"Shutting down and killing engine at index {i}")
|
||||
try:
|
||||
@@ -175,4 +174,4 @@ class RolloutHealthMonitor:
|
||||
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
|
||||
self._server_group.all_engines[i] = None
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import miles.utils.external_utils.command_utils as U
|
||||
|
||||
TIGHT_DEVICE_MEMORY = U.get_bool_env_var("MILES_TEST_TIGHT_DEVICE_MEMORY", "1")
|
||||
|
||||
MODEL_NAME = "Qwen2.5-0.5B-Instruct"
|
||||
MODEL_TYPE = "qwen2.5-0.5B"
|
||||
NUM_GPUS = 8
|
||||
|
||||
# Inline sglang config: same model, 2 engine groups with different sizes.
|
||||
# Group 1: 4 GPUs, 1 GPU/engine (tp=1) -> 4 engines
|
||||
# Group 2: 4 GPUs, 1 GPU/engine (tp=1) -> 4 engines
|
||||
# Tests that ServerGroup/RolloutServer correctly manages multiple groups
|
||||
# behind a single router, with separate port cursors per group.
|
||||
SGLANG_CONFIG_YAML = """\
|
||||
sglang:
|
||||
- name: default
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
"""
|
||||
|
||||
|
||||
def prepare():
|
||||
U.exec_command("mkdir -p /root/models /root/datasets")
|
||||
U.exec_command(f"huggingface-cli download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}")
|
||||
U.hf_download_dataset("zhuzilin/gsm8k")
|
||||
|
||||
|
||||
def execute():
|
||||
config_file = tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", prefix="sglang_config_", delete=False)
|
||||
config_file.write(SGLANG_CONFIG_YAML)
|
||||
config_file.flush()
|
||||
config_path = config_file.name
|
||||
|
||||
ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/models/{MODEL_NAME}/ "
|
||||
|
||||
rollout_args = (
|
||||
"--prompt-data /root/datasets/gsm8k/train.parquet "
|
||||
"--input-key messages "
|
||||
"--label-key label "
|
||||
"--apply-chat-template "
|
||||
"--rollout-shuffle "
|
||||
"--rm-type math "
|
||||
"--num-rollout 3 "
|
||||
"--rollout-batch-size 8 "
|
||||
"--n-samples-per-prompt 4 "
|
||||
"--rollout-max-response-len 1024 "
|
||||
"--rollout-temperature 0.8 "
|
||||
"--over-sampling-batch-size 16 "
|
||||
"--dynamic-sampling-filter-path miles.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std "
|
||||
"--global-batch-size 32 "
|
||||
)
|
||||
|
||||
eval_args = (
|
||||
"--eval-interval 20 "
|
||||
"--eval-prompt-data gsm8k /root/datasets/gsm8k/test.parquet "
|
||||
"--n-samples-per-eval-prompt 1 "
|
||||
"--eval-max-response-len 1024 "
|
||||
"--eval-top-k 1 "
|
||||
)
|
||||
|
||||
perf_args = (
|
||||
"--tensor-model-parallel-size 1 "
|
||||
"--sequence-parallel "
|
||||
"--pipeline-model-parallel-size 1 "
|
||||
"--context-parallel-size 1 "
|
||||
"--expert-model-parallel-size 1 "
|
||||
"--expert-tensor-parallel-size 1 "
|
||||
"--use-dynamic-batch-size "
|
||||
"--max-tokens-per-gpu 9216 "
|
||||
)
|
||||
|
||||
grpo_args = (
|
||||
"--advantage-estimator grpo "
|
||||
"--use-kl-loss "
|
||||
"--kl-loss-coef 0.00 "
|
||||
"--kl-loss-type low_var_kl "
|
||||
"--entropy-coef 0.00 "
|
||||
"--eps-clip 0.2 "
|
||||
"--eps-clip-high 0.28 "
|
||||
)
|
||||
|
||||
optimizer_args = (
|
||||
"--optimizer adam "
|
||||
"--lr 1e-6 "
|
||||
"--lr-decay-style constant "
|
||||
"--weight-decay 0.1 "
|
||||
"--adam-beta1 0.9 "
|
||||
"--adam-beta2 0.98 "
|
||||
)
|
||||
|
||||
sglang_args = (
|
||||
"--rollout-num-gpus-per-engine 1 "
|
||||
f"--sglang-mem-fraction-static {0.5 if TIGHT_DEVICE_MEMORY else 0.6} "
|
||||
"--sglang-enable-metrics "
|
||||
"--sglang-cuda-graph-max-bs 32 "
|
||||
f"--sglang-config {config_path} "
|
||||
)
|
||||
|
||||
ci_args = "--ci-test "
|
||||
|
||||
misc_args = (
|
||||
"--attention-dropout 0.0 "
|
||||
"--hidden-dropout 0.0 "
|
||||
"--accumulate-allreduce-grads-in-fp32 "
|
||||
"--attention-softmax-in-fp32 "
|
||||
"--attention-backend flash "
|
||||
"--actor-num-nodes 1 "
|
||||
"--actor-num-gpus-per-node 8 "
|
||||
"--colocate "
|
||||
"--megatron-to-hf-mode bridge "
|
||||
)
|
||||
|
||||
train_args = (
|
||||
f"{ckpt_args} "
|
||||
f"{rollout_args} "
|
||||
f"{optimizer_args} "
|
||||
f"{grpo_args} "
|
||||
f"{U.get_default_wandb_args(__file__)} "
|
||||
f"{perf_args} "
|
||||
f"{eval_args} "
|
||||
f"{sglang_args} "
|
||||
f"{ci_args} "
|
||||
f"{misc_args} "
|
||||
)
|
||||
|
||||
U.execute_train(
|
||||
train_args=train_args,
|
||||
num_gpus_per_node=NUM_GPUS,
|
||||
megatron_model_type=MODEL_TYPE,
|
||||
extra_env_vars={"MILES_EXPERIMENTAL_ROLLOUT_REFACTOR": "1"},
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
prepare()
|
||||
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
|
||||
os.environ.pop(proxy_var, None)
|
||||
execute()
|
||||
@@ -0,0 +1,156 @@
|
||||
"""E2E test: mixed offload with updatable + frozen models.
|
||||
|
||||
Deploys two models via --sglang-config in colocate mode:
|
||||
- "actor": update_weights=true, 4 GPUs -> overlaps with megatron, gets offloaded
|
||||
and weights updated from training.
|
||||
- "ref": update_weights=false, 4 GPUs -> overlaps with megatron, gets offloaded
|
||||
and weights restored from disk (update_weights_from_disk).
|
||||
|
||||
Key coverage:
|
||||
- Per-group needs_offload (both overlap with megatron in colocate mode)
|
||||
- update_weights_from_disk for frozen model
|
||||
- Selective flush_cache (only for offloaded / updatable engines)
|
||||
- Offload/onload cycle completes without crash
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import miles.utils.external_utils.command_utils as U
|
||||
|
||||
TIGHT_DEVICE_MEMORY = U.get_bool_env_var("MILES_TEST_TIGHT_DEVICE_MEMORY", "1")
|
||||
|
||||
MODEL_NAME = "Qwen2.5-0.5B-Instruct"
|
||||
MODEL_TYPE = "qwen2.5-0.5B"
|
||||
NUM_GPUS = 8
|
||||
|
||||
# Two models on 8 GPUs (colocate): actor gets weight updates, ref is frozen.
|
||||
SGLANG_CONFIG_YAML = """\
|
||||
sglang:
|
||||
- name: actor
|
||||
update_weights: true
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
- name: ref
|
||||
update_weights: false
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
"""
|
||||
|
||||
|
||||
def prepare():
|
||||
U.exec_command("mkdir -p /root/models /root/datasets")
|
||||
U.exec_command(f"huggingface-cli download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}")
|
||||
U.hf_download_dataset("zhuzilin/gsm8k")
|
||||
|
||||
|
||||
def execute():
|
||||
config_file = tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", prefix="sglang_mixed_offload_", delete=False)
|
||||
config_file.write(SGLANG_CONFIG_YAML)
|
||||
config_file.flush()
|
||||
config_path = config_file.name
|
||||
|
||||
ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/models/{MODEL_NAME}/ "
|
||||
|
||||
rollout_args = (
|
||||
"--prompt-data /root/datasets/gsm8k/train.parquet "
|
||||
"--input-key messages "
|
||||
"--label-key label "
|
||||
"--apply-chat-template "
|
||||
"--rollout-shuffle "
|
||||
"--rm-type math "
|
||||
"--num-rollout 3 "
|
||||
"--rollout-batch-size 8 "
|
||||
"--n-samples-per-prompt 4 "
|
||||
"--rollout-max-response-len 512 "
|
||||
"--rollout-temperature 0.8 "
|
||||
"--global-batch-size 32 "
|
||||
)
|
||||
|
||||
eval_args = (
|
||||
"--eval-interval 20 "
|
||||
"--eval-prompt-data gsm8k /root/datasets/gsm8k/test.parquet "
|
||||
"--n-samples-per-eval-prompt 1 "
|
||||
"--eval-max-response-len 512 "
|
||||
"--eval-top-k 1 "
|
||||
)
|
||||
|
||||
perf_args = (
|
||||
"--tensor-model-parallel-size 1 "
|
||||
"--sequence-parallel "
|
||||
"--pipeline-model-parallel-size 1 "
|
||||
"--context-parallel-size 1 "
|
||||
"--expert-model-parallel-size 1 "
|
||||
"--expert-tensor-parallel-size 1 "
|
||||
"--use-dynamic-batch-size "
|
||||
"--max-tokens-per-gpu 4096 "
|
||||
)
|
||||
|
||||
grpo_args = (
|
||||
"--advantage-estimator grpo "
|
||||
"--use-kl-loss "
|
||||
"--kl-loss-coef 0.00 "
|
||||
"--kl-loss-type low_var_kl "
|
||||
"--entropy-coef 0.00 "
|
||||
"--eps-clip 0.2 "
|
||||
)
|
||||
|
||||
optimizer_args = (
|
||||
"--optimizer adam "
|
||||
"--lr 1e-6 "
|
||||
"--lr-decay-style constant "
|
||||
"--weight-decay 0.1 "
|
||||
"--adam-beta1 0.9 "
|
||||
"--adam-beta2 0.98 "
|
||||
)
|
||||
|
||||
sglang_args = (
|
||||
"--rollout-num-gpus-per-engine 1 "
|
||||
f"--sglang-mem-fraction-static {0.6 if TIGHT_DEVICE_MEMORY else 0.7} "
|
||||
"--sglang-cuda-graph-max-bs 32 "
|
||||
f"--sglang-config {config_path} "
|
||||
)
|
||||
|
||||
ci_args = "--ci-test "
|
||||
|
||||
misc_args = (
|
||||
"--attention-dropout 0.0 "
|
||||
"--hidden-dropout 0.0 "
|
||||
"--accumulate-allreduce-grads-in-fp32 "
|
||||
"--attention-softmax-in-fp32 "
|
||||
"--attention-backend flash "
|
||||
"--actor-num-nodes 1 "
|
||||
"--actor-num-gpus-per-node 8 "
|
||||
"--colocate "
|
||||
"--megatron-to-hf-mode bridge "
|
||||
)
|
||||
|
||||
train_args = (
|
||||
f"{ckpt_args} "
|
||||
f"{rollout_args} "
|
||||
f"{optimizer_args} "
|
||||
f"{grpo_args} "
|
||||
f"{U.get_default_wandb_args(__file__)} "
|
||||
f"{perf_args} "
|
||||
f"{eval_args} "
|
||||
f"{sglang_args} "
|
||||
f"{ci_args} "
|
||||
f"{misc_args} "
|
||||
)
|
||||
|
||||
U.execute_train(
|
||||
train_args=train_args,
|
||||
num_gpus_per_node=NUM_GPUS,
|
||||
megatron_model_type=MODEL_TYPE,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
prepare()
|
||||
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
|
||||
os.environ.pop(proxy_var, None)
|
||||
execute()
|
||||
@@ -0,0 +1,163 @@
|
||||
"""E2E test: mixed offload with fault tolerance.
|
||||
|
||||
Same two-model layout as test_sglang_config_mixed_offload.py but with
|
||||
fault tolerance enabled. --ci-test triggers a simulated engine crash
|
||||
on the updatable (actor) server, testing:
|
||||
- Health monitor detects crash and marks engine as None
|
||||
- RolloutServer.recover() restarts the dead engine
|
||||
- Updatable engines: offload -> resume_memory_occupation -> update_weights
|
||||
- Non-updatable engines: offload -> update_weights_from_disk
|
||||
- Training continues after recovery
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import miles.utils.external_utils.command_utils as U
|
||||
|
||||
TIGHT_DEVICE_MEMORY = U.get_bool_env_var("MILES_TEST_TIGHT_DEVICE_MEMORY", "1")
|
||||
|
||||
MODEL_NAME = "Qwen2.5-0.5B-Instruct"
|
||||
MODEL_TYPE = "qwen2.5-0.5B"
|
||||
NUM_GPUS = 8
|
||||
|
||||
# Two models on 8 GPUs (colocate): actor gets weight updates, ref is frozen.
|
||||
SGLANG_CONFIG_YAML = """\
|
||||
sglang:
|
||||
- name: actor
|
||||
update_weights: true
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
- name: ref
|
||||
update_weights: false
|
||||
server_groups:
|
||||
- worker_type: regular
|
||||
num_gpus: 4
|
||||
num_gpus_per_engine: 1
|
||||
"""
|
||||
|
||||
|
||||
def prepare():
|
||||
U.exec_command("mkdir -p /root/models /root/datasets")
|
||||
U.exec_command(f"huggingface-cli download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}")
|
||||
U.hf_download_dataset("zhuzilin/gsm8k")
|
||||
|
||||
|
||||
def execute():
|
||||
config_file = tempfile.NamedTemporaryFile(
|
||||
mode="w", suffix=".yaml", prefix="sglang_mixed_offload_ft_", delete=False
|
||||
)
|
||||
config_file.write(SGLANG_CONFIG_YAML)
|
||||
config_file.flush()
|
||||
config_path = config_file.name
|
||||
|
||||
ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/models/{MODEL_NAME}/ "
|
||||
|
||||
rollout_args = (
|
||||
"--prompt-data /root/datasets/gsm8k/train.parquet "
|
||||
"--input-key messages "
|
||||
"--label-key label "
|
||||
"--apply-chat-template "
|
||||
"--rollout-shuffle "
|
||||
"--rm-type math "
|
||||
"--num-rollout 3 "
|
||||
"--rollout-batch-size 8 "
|
||||
"--n-samples-per-prompt 4 "
|
||||
"--rollout-max-response-len 512 "
|
||||
"--rollout-temperature 0.8 "
|
||||
"--global-batch-size 32 "
|
||||
)
|
||||
|
||||
eval_args = (
|
||||
"--eval-interval 20 "
|
||||
"--eval-prompt-data gsm8k /root/datasets/gsm8k/test.parquet "
|
||||
"--n-samples-per-eval-prompt 1 "
|
||||
"--eval-max-response-len 512 "
|
||||
"--eval-top-k 1 "
|
||||
)
|
||||
|
||||
perf_args = (
|
||||
"--tensor-model-parallel-size 1 "
|
||||
"--sequence-parallel "
|
||||
"--pipeline-model-parallel-size 1 "
|
||||
"--context-parallel-size 1 "
|
||||
"--expert-model-parallel-size 1 "
|
||||
"--expert-tensor-parallel-size 1 "
|
||||
"--use-dynamic-batch-size "
|
||||
"--max-tokens-per-gpu 4096 "
|
||||
)
|
||||
|
||||
grpo_args = (
|
||||
"--advantage-estimator grpo "
|
||||
"--use-kl-loss "
|
||||
"--kl-loss-coef 0.00 "
|
||||
"--kl-loss-type low_var_kl "
|
||||
"--entropy-coef 0.00 "
|
||||
"--eps-clip 0.2 "
|
||||
)
|
||||
|
||||
optimizer_args = (
|
||||
"--optimizer adam "
|
||||
"--lr 1e-6 "
|
||||
"--lr-decay-style constant "
|
||||
"--weight-decay 0.1 "
|
||||
"--adam-beta1 0.9 "
|
||||
"--adam-beta2 0.98 "
|
||||
)
|
||||
|
||||
sglang_args = (
|
||||
"--rollout-num-gpus-per-engine 1 "
|
||||
f"--sglang-mem-fraction-static {0.6 if TIGHT_DEVICE_MEMORY else 0.7} "
|
||||
"--sglang-cuda-graph-max-bs 32 "
|
||||
f"--sglang-config {config_path} "
|
||||
)
|
||||
|
||||
ci_args = "--ci-test "
|
||||
|
||||
fault_tolerance_args = (
|
||||
"--use-fault-tolerance "
|
||||
"--rollout-health-check-interval 5 "
|
||||
"--rollout-health-check-timeout 10 "
|
||||
"--rollout-health-check-first-wait 0 "
|
||||
)
|
||||
|
||||
misc_args = (
|
||||
"--attention-dropout 0.0 "
|
||||
"--hidden-dropout 0.0 "
|
||||
"--accumulate-allreduce-grads-in-fp32 "
|
||||
"--attention-softmax-in-fp32 "
|
||||
"--attention-backend flash "
|
||||
"--actor-num-nodes 1 "
|
||||
"--actor-num-gpus-per-node 8 "
|
||||
"--colocate "
|
||||
"--megatron-to-hf-mode bridge "
|
||||
)
|
||||
|
||||
train_args = (
|
||||
f"{ckpt_args} "
|
||||
f"{rollout_args} "
|
||||
f"{optimizer_args} "
|
||||
f"{grpo_args} "
|
||||
f"{U.get_default_wandb_args(__file__)} "
|
||||
f"{perf_args} "
|
||||
f"{eval_args} "
|
||||
f"{sglang_args} "
|
||||
f"{ci_args} "
|
||||
f"{fault_tolerance_args} "
|
||||
f"{misc_args} "
|
||||
)
|
||||
|
||||
U.execute_train(
|
||||
train_args=train_args,
|
||||
num_gpus_per_node=NUM_GPUS,
|
||||
megatron_model_type=MODEL_TYPE,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
prepare()
|
||||
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
|
||||
os.environ.pop(proxy_var, None)
|
||||
execute()
|
||||
@@ -0,0 +1,140 @@
|
||||
"""Unit tests for SglangConfig multi-model parsing with update_weights."""
|
||||
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
|
||||
def _write_yaml(data: dict) -> str:
|
||||
f = tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False)
|
||||
yaml.dump(data, f)
|
||||
f.flush()
|
||||
return f.name
|
||||
|
||||
|
||||
class TestSglangConfigUpdateWeights:
|
||||
def test_update_weights_default_true(self):
|
||||
"""Models without explicit update_weights should resolve to True when model_path matches hf_checkpoint."""
|
||||
from argparse import Namespace
|
||||
|
||||
from miles.backends.sglang_utils.sglang_config import SglangConfig
|
||||
|
||||
path = _write_yaml(
|
||||
{
|
||||
"sglang": [
|
||||
{
|
||||
"name": "actor",
|
||||
"engine_groups": [{"worker_type": "regular", "num_gpus": 4}],
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
config = SglangConfig.from_yaml(path)
|
||||
assert len(config.models) == 1
|
||||
# Before resolve, update_weights is None (not yet inferred)
|
||||
assert config.models[0].update_weights is None
|
||||
# After resolve with matching hf_checkpoint, defaults to True
|
||||
args = Namespace(hf_checkpoint="/path/to/model", rollout_num_gpus_per_engine=1)
|
||||
config.models[0].resolve(args)
|
||||
assert config.models[0].update_weights is True
|
||||
|
||||
def test_update_weights_explicit_false(self):
|
||||
"""Models with update_weights: false should be parsed correctly."""
|
||||
from miles.backends.sglang_utils.sglang_config import SglangConfig
|
||||
|
||||
path = _write_yaml(
|
||||
{
|
||||
"sglang": [
|
||||
{
|
||||
"name": "actor",
|
||||
"update_weights": True,
|
||||
"engine_groups": [{"worker_type": "regular", "num_gpus": 4}],
|
||||
},
|
||||
{
|
||||
"name": "ref",
|
||||
"update_weights": False,
|
||||
"model_path": "/path/to/ref",
|
||||
"engine_groups": [{"worker_type": "regular", "num_gpus": 2}],
|
||||
},
|
||||
]
|
||||
}
|
||||
)
|
||||
config = SglangConfig.from_yaml(path)
|
||||
assert len(config.models) == 2
|
||||
assert config.models[0].name == "actor"
|
||||
assert config.models[0].update_weights is True
|
||||
assert config.models[1].name == "ref"
|
||||
assert config.models[1].update_weights is False
|
||||
assert config.models[1].model_path == "/path/to/ref"
|
||||
|
||||
def test_multi_model_total_gpus(self):
|
||||
"""total_num_gpus should sum across all models."""
|
||||
from miles.backends.sglang_utils.sglang_config import SglangConfig
|
||||
|
||||
path = _write_yaml(
|
||||
{
|
||||
"sglang": [
|
||||
{
|
||||
"name": "actor",
|
||||
"server_groups": [{"worker_type": "regular", "num_gpus": 8}],
|
||||
},
|
||||
{
|
||||
"name": "ref",
|
||||
"update_weights": False,
|
||||
"server_groups": [{"worker_type": "regular", "num_gpus": 4}],
|
||||
},
|
||||
]
|
||||
}
|
||||
)
|
||||
config = SglangConfig.from_yaml(path)
|
||||
assert config.total_num_gpus == 12
|
||||
|
||||
|
||||
class TestGetModelUrl:
|
||||
def test_get_model_url_basic(self):
|
||||
"""get_model_url should return the correct URL for a named model."""
|
||||
from argparse import Namespace
|
||||
|
||||
from miles.rollout.sglang_rollout import get_model_url
|
||||
|
||||
args = Namespace(
|
||||
sglang_router_ip="10.0.0.1",
|
||||
sglang_router_port=3000,
|
||||
sglang_model_routers={
|
||||
"actor": ("10.0.0.1", 3000),
|
||||
"ref": ("10.0.0.1", 3001),
|
||||
},
|
||||
)
|
||||
assert get_model_url(args, "actor") == "http://10.0.0.1:3000/generate"
|
||||
assert get_model_url(args, "ref") == "http://10.0.0.1:3001/generate"
|
||||
assert get_model_url(args, "ref", "/v1/chat/completions") == "http://10.0.0.1:3001/v1/chat/completions"
|
||||
|
||||
def test_get_model_url_fallback(self):
|
||||
"""get_model_url should fall back to default router if model not found."""
|
||||
from argparse import Namespace
|
||||
|
||||
from miles.rollout.sglang_rollout import get_model_url
|
||||
|
||||
args = Namespace(
|
||||
sglang_router_ip="10.0.0.1",
|
||||
sglang_router_port=3000,
|
||||
sglang_model_routers={"actor": ("10.0.0.1", 3000)},
|
||||
)
|
||||
assert get_model_url(args, "unknown") == "http://10.0.0.1:3000/generate"
|
||||
|
||||
def test_get_model_url_no_routers(self):
|
||||
"""get_model_url should work when sglang_model_routers is not set."""
|
||||
from argparse import Namespace
|
||||
|
||||
from miles.rollout.sglang_rollout import get_model_url
|
||||
|
||||
args = Namespace(
|
||||
sglang_router_ip="10.0.0.1",
|
||||
sglang_router_port=3000,
|
||||
)
|
||||
assert get_model_url(args, "anything") == "http://10.0.0.1:3000/generate"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
Reference in New Issue
Block a user