Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions docs/advanced/fault-tolerance.md
Original file line number Diff line number Diff line change
Expand Up @@ -49,10 +49,10 @@ Each loop iteration does:
## Engine recovery

When `--use-fault-tolerance` is on, `MegatronActor.update_weights` calls
`rollout_manager.recover_updatable_engines` on rank 0 before each weight
`inference_controller.recover_updatable_engines` before each weight
update (`miles/backends/megatron_utils/actor.py`).

`recover_updatable_engines` (`miles/ray/rollout/rollout_manager.py`):
`recover_updatable_engines` (`miles/ray/rollout/inference_controller.py`):

1. Pauses health monitoring.
2. Calls `srv.recover()` on the updatable server.
Expand Down
2 changes: 1 addition & 1 deletion miles/backends/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@
from .update_weight_utils import UpdateWeightFromDistributed, UpdateWeightFromTensor

if TYPE_CHECKING:
from miles.ray.rollout.rollout_manager import EnginesAndLock
from miles.ray.rollout.inference_controller import EnginesAndLock
from miles.utils.audit_utils.witness.allocator import WitnessInfo

logger = logging.getLogger(__name__)
Expand Down
2 changes: 1 addition & 1 deletion miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@
from .replay_utils import register_replay_list_moe

if TYPE_CHECKING:
from miles.ray.rollout.rollout_manager import EnginesAndLock
from miles.ray.rollout.inference_controller import EnginesAndLock

logging.getLogger("megatron").setLevel(logging.WARNING)

Expand Down
4 changes: 2 additions & 2 deletions miles/dashboard/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,7 +322,7 @@ def register_router(args) -> None:


def register_engines(servers) -> None:
"""Called at the top of every RolloutManager.generate(): pushes an engine
"""Called at the top of every InferenceController.prepare_rollout(): pushes an engine
topology snapshot whenever the set of engine actors changed (startup,
fault-tolerance recovery). Steady state costs one local tuple compare."""
global _engines_fingerprint
Expand All @@ -344,7 +344,7 @@ def register_engines(servers) -> None:


def report_data_buffer(length: int | None) -> None:
"""Called at the top of every ``RolloutManager.generate()`` alongside
"""Called at the top of every ``RolloutExecutor.generate()`` alongside
``register_engines``, with ``getattr(data_source, "get_buffer_length",
lambda: None)()``. A no-op for ``length is None`` — most data sources
(plain ``RolloutDataSource``) never buffer samples across steps."""
Expand Down
2 changes: 1 addition & 1 deletion miles/dashboard/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -151,7 +151,7 @@ def from_dict(cls, data: dict) -> TopologySnapshot:
class DataBufferSample(Record):
"""One report of ``RolloutDataSourceWithBuffer.get_buffer_length()``
(design doc's "show the data status in databuffer" ask). Appended once
per ``RolloutManager.generate()`` call — a no-op for plain
per ``RolloutExecutor.generate()`` call — a no-op for plain
``RolloutDataSource`` runs, which never buffer samples across steps.
Low-rate and unpartitioned like ``TopologySnapshot``: only the latest
value matters for display."""
Expand Down
18 changes: 10 additions & 8 deletions miles/ray/actor_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,8 @@ def __init__(
num_gpus_per_node,
pg: tuple[PlacementGroup, list[int], list[int]],
*,
rollout_manager: object | None,
inference_controller: object | None,
rollout_executor: object | None,
num_gpus_per_actor: float = 1,
role: str,
with_ref: bool,
Expand All @@ -43,7 +44,8 @@ def __init__(
self._num_gpus_per_node = num_gpus_per_node
self.role = role
self.with_ref = with_ref
self._rollout_manager = rollout_manager
self._inference_controller = inference_controller
self._rollout_executor = rollout_executor
self.with_opd_teacher = with_opd_teacher

# Allocate the GPUs for actors w/o instantiating them
Expand Down Expand Up @@ -122,13 +124,13 @@ async def update_weights(self, rollout_id: int | None = None):
return

if self.args.use_fault_tolerance and "rollout" in self.args.ft_components:
await self._rollout_manager.recover_updatable_engines.remote()
await self._inference_controller.recover_updatable_engines()

info = await self._rollout_manager.get_updatable_engines_and_lock.remote()
await self._rollout_manager.health_monitoring_pause.remote()
info = await self._inference_controller.get_updatable_engines_and_lock()
await self._inference_controller.health_monitoring_pause()

await self._broadcast("update_weights", info=info)
await self._rollout_manager.clear_updatable_has_new_engines.remote()
await self._inference_controller.clear_updatable_has_new_engines()

async def reconcile_adapters(self) -> None:
"""Multi-LoRA: reconcile loaded adapters with the controller's active set
Expand All @@ -144,8 +146,8 @@ async def offload(self):
async def clear_memory(self):
await self._broadcast("clear_memory")

async def set_rollout_manager(self):
await self._broadcast("set_rollout_manager", self._rollout_manager)
async def set_rollout_executor(self):
await self._broadcast("set_rollout_executor", self._rollout_executor)

async def _broadcast(self, method_name: str, *args, **kwargs) -> list:
refs = [getattr(actor, method_name).remote(*args, **kwargs) for actor in self._actor_handles]
Expand Down
63 changes: 44 additions & 19 deletions miles/ray/placement_group.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
import copy
import logging
import socket
from typing import NamedTuple

import ray
from ray.util.placement_group import placement_group
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy

from miles.utils.environ import enable_experimental_ft_trainer
from ..utils.ray_utils import compute_ray_pin_head_options
from .rollout.rollout_manager import RolloutManager
from .rollout.inference_controller import InferenceController
from .rollout.rollout_executor import RolloutExecutor

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -124,7 +126,15 @@ def create_placement_groups(args):


def allocate_train_group(
args, num_nodes, num_gpus_per_node, pg, role: str, with_ref: bool, rollout_manager, with_opd_teacher: bool = False
args,
num_nodes,
num_gpus_per_node,
pg,
role: str,
with_ref: bool,
inference_controller,
rollout_executor,
with_opd_teacher: bool = False,
):
train_group_cls = _select_train_group_class()
return train_group_cls(
Expand All @@ -135,20 +145,22 @@ def allocate_train_group(
num_gpus_per_actor=0.4,
role=role,
with_ref=with_ref,
rollout_manager=rollout_manager,
inference_controller=inference_controller,
rollout_executor=rollout_executor,
with_opd_teacher=with_opd_teacher,
)


async def create_training_models(args, pgs, rollout_manager):
async def create_training_models(args, pgs, inference_controller, rollout_executor):
actor_model = allocate_train_group(
args=args,
num_nodes=args.actor_num_nodes,
num_gpus_per_node=args.actor_num_gpus_per_node,
pg=pgs["actor"],
role="actor",
with_ref=args.kl_coef != 0 or args.use_kl_loss,
rollout_manager=rollout_manager,
inference_controller=inference_controller,
rollout_executor=rollout_executor,
with_opd_teacher=args.use_opd and args.opd_type == "megatron",
)
actor_start_rollout_ids = await actor_model.init()
Expand All @@ -165,7 +177,8 @@ async def create_training_models(args, pgs, rollout_manager):
pg=pgs["critic"],
role="critic",
with_ref=False,
rollout_manager=None,
inference_controller=None,
rollout_executor=None,
)
critic_start_rollout_ids = await critic_model.init()
else:
Expand All @@ -177,36 +190,48 @@ async def create_training_models(args, pgs, rollout_manager):
if args.start_rollout_id is None:
args.start_rollout_id = start_rollout_ids[0]

await actor_model.set_rollout_manager()
await actor_model.set_rollout_executor()
if args.rollout_global_dataset:
await rollout_manager.load.remote(args.start_rollout_id - 1)
await rollout_executor.load.remote(args.start_rollout_id - 1)

return actor_model, critic_model


def create_rollout_manager(args, pg):
rollout_manager = RolloutManager.options(
class RolloutComponents(NamedTuple):
inference_controller: InferenceController
rollout_executor: ray.actor.ActorHandle
num_rollout_per_epoch: int | None


async def create_rollout_components(args, pg) -> RolloutComponents:
inference_controller = InferenceController(args, pg)

rollout_executor = RolloutExecutor.options(
num_cpus=1, num_gpus=0, **(compute_ray_pin_head_options() if args.pin_rollout_manager_to_head else {})
).remote(args, pg)
).remote(args)

# calculate num_rollout from num_epoch
num_rollout_per_epoch = None
if args.num_rollout is None:
num_rollout_per_epoch = ray.get(rollout_manager.get_num_rollout_per_epoch.remote())
num_rollout_per_epoch = ray.get(rollout_executor.get_num_rollout_per_epoch.remote())
args.num_rollout = num_rollout_per_epoch * args.num_epoch
assert args.num_rollout > 0

await rollout_executor.set_eval_fleet.remote(inference_controller.eval_fleet)

if args.check_weight_update_equal:
ray.get(rollout_manager.check_weights.remote(action="snapshot"))
ray.get(
rollout_manager.check_weights.remote(action="reset_tensors", skip_list=args.check_weight_update_skip_list)
)
await inference_controller.check_weights(action="snapshot")
await inference_controller.check_weights(action="reset_tensors", skip_list=args.check_weight_update_skip_list)

if args.offload_rollout:
if args.colocate_memory_peak_device == "gpu":
# keep weight on GPU to reduce peak CPU memory
ray.get(rollout_manager.offload_kv.remote())
await inference_controller.offload_kv()
else:
ray.get(rollout_manager.offload.remote())
await inference_controller.offload()

return rollout_manager, num_rollout_per_epoch
return RolloutComponents(
inference_controller=inference_controller,
rollout_executor=rollout_executor,
num_rollout_per_epoch=num_rollout_per_epoch,
)
14 changes: 6 additions & 8 deletions miles/ray/rollout/inference_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,6 @@
from miles.ray.rollout.server_cell import get_cell_indexer_of_id_map
from miles.ray.utils import Lock
from miles.utils.health_monitor import RolloutHealthMonitor
from miles.utils.http_utils import init_http_client


logger = logging.getLogger(__name__)
Expand All @@ -29,7 +28,6 @@ def __init__(self, args, pg):
else:
self.servers = start_rollout_servers(args, pg)
dashboard_hooks.register_router(args)
init_http_client(args)
start_session_server(args)
self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote()
self.rollout_id = -1
Expand All @@ -56,18 +54,18 @@ async def prepare_rollout(self, rollout_id):
await self._try_ci_fault_injection()
dashboard_hooks.register_engines(self.servers)

def prepare_eval(self):
async def prepare_eval(self):
self._health_monitoring_resume()

def dispose(self):
async def dispose(self):
for monitor in self._health_monitors:
monitor.stop()

# -------------------------- offload/onload -----------------------------

# TODO may parallelly execute offload/onload across services
async def offload(self, tags: list[str] | None = None):
self.health_monitoring_pause()
await self.health_monitoring_pause()
for srv in self.servers.values():
await srv.offload(tags=tags)

Expand Down Expand Up @@ -117,7 +115,7 @@ async def get_updatable_engines_and_lock(self):
engine_gpu_offsets=srv.engine_gpu_offsets,
)

def clear_updatable_has_new_engines(self):
async def clear_updatable_has_new_engines(self):
# when fault tolerance is not enabled, we need to manually clear has_new_engines after update_weights
srv = self._get_updatable_server()
if srv:
Expand All @@ -129,7 +127,7 @@ async def recover_updatable_engines(self) -> None:
Recovers the updatable model (the one that receives weight
updates from training).
"""
self.health_monitoring_pause()
await self.health_monitoring_pause()
srv = self._get_updatable_server()
if self.rollout_id == -1 or srv is None:
return
Expand Down Expand Up @@ -177,7 +175,7 @@ async def check_weights(

# -------------------------- utils -----------------------------

def health_monitoring_pause(self) -> None:
async def health_monitoring_pause(self) -> None:
for monitor in self._health_monitors:
monitor.pause()

Expand Down
11 changes: 6 additions & 5 deletions miles/ray/rollout/rollout_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,9 +25,10 @@
from miles.utils import object_store
from miles.utils.audit_utils.event_analyzer import analyzer as event_analyzer
from miles.utils.audit_utils.event_logger import checkpoint as event_logger_checkpoint
from miles.utils.audit_utils.process_identity import RolloutManagerProcessIdentity
from miles.utils.audit_utils.process_identity import RolloutExecutorProcessIdentity
from miles.utils.environ import use_legacy_rollout_v1
from miles.utils.hf_config import is_complete_hf_export
from miles.utils.http_utils import init_http_client
from miles.utils.logging_utils import configure_logger
from miles.utils.metric_checker import MetricChecker
from miles.utils.misc import load_function
Expand All @@ -47,7 +48,7 @@ class RolloutExecutor:

def __init__(self, args):
event_logger_checkpoint.restore(args)
configure_logger(args, source=RolloutManagerProcessIdentity())
configure_logger(args, source=RolloutExecutorProcessIdentity())

self.args = args
# set by the training actor after each weight update
Expand All @@ -56,6 +57,9 @@ def __init__(self, args):
init_tracking(args, primary=False, router_addr=f"http://{args.sglang_router_ip}:{args.sglang_router_port}")
object_store.init_instance(args, contribute_segment=False)

if not self.args.debug_train_only:
init_http_client(args)

data_source_cls = load_function(self.args.data_source_path)
self.data_source = data_source_cls(args)

Expand Down Expand Up @@ -95,9 +99,6 @@ def __init__(self, args):
# -------------------------- lifecycle -----------------------------
# TODO: may have a `async def init` here later

def get_router_address(self) -> tuple[str, int]:
return self.args.sglang_router_ip, self.args.sglang_router_port

def dispose(self):
if (close := getattr(self.data_source, "close", None)) is not None:
close()
Expand Down
2 changes: 1 addition & 1 deletion miles/ray/rollout/server_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,7 +170,7 @@ def start_engines(
return init_handles, new_engine_indices

# There are two callers, only one of them will exist in a running system
# 1. For new callers (RolloutManager.stop_cell, main thread, async),
# 1. For new callers (InferenceController.stop_cell, main thread, async),
# deliberately make this function non-async here to avoid introducing two states
# like "stopping (but not stopped)" vs "stopped", since single-thread async code will not yield
# without an await point
Expand Down
12 changes: 6 additions & 6 deletions miles/ray/train/cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,15 +36,15 @@ def __init__(
with_opd_teacher: bool = False,
cell_index: int,
actor_factory: ActorFactory,
rollout_manager: object | None,
rollout_executor: object | None,
health_checker: BaseHealthChecker,
) -> None:
self.args = args
self.cell_index = cell_index
self.role = role
self.with_ref = with_ref
self.with_opd_teacher = with_opd_teacher
self.rollout_manager = rollout_manager
self.rollout_executor = rollout_executor
self.actor_factory = actor_factory
self.health_checker = health_checker

Expand Down Expand Up @@ -73,9 +73,9 @@ async def init(
await self.health_checker.start()
return results

async def set_rollout_manager(self):
if (m := self.rollout_manager) is not None:
return await self.execute("set_rollout_manager", m)
async def set_rollout_executor(self):
if (executor := self.rollout_executor) is not None:
return await self.execute("set_rollout_executor", executor)
return []

# ------------------------ API :: cooperatively prepare ------------------------
Expand All @@ -101,7 +101,7 @@ async def prepare_indep_dp_mode_healing(
recv_ckpt_src_rank=recv_ckpt_src_rank,
)

await self.set_rollout_manager()
await self.set_rollout_executor()

# ------------------------ state transition ------------------------

Expand Down
Loading
Loading