[feat] (1/n) fault tolerance and slime rollout structure (#723)

This commit is contained in:
Ethan (Yusheng) Su
2026-03-16 23:01:53 -07:00
committed by GitHub
parent 64353ccb24
commit 39393552cb
17 changed files with 1594 additions and 238 deletions
+3 -3
View File
@@ -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 }}
+9 -4
View File
@@ -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
+10 -5
View File
@@ -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 = []
+15
View File
@@ -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)
+44 -6
View File
@@ -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
View File
@@ -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):
+31 -2
View File
@@ -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
+8
View File
@@ -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 (
+9 -10
View File
@@ -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()
+140
View File
@@ -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"])