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
5 changes: 3 additions & 2 deletions miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -789,12 +789,13 @@ def update_weights(self, info: "UpdatableEngines") -> None:
if process_groups_are_temporary:
reload_process_groups()

if has_new_engines or not self.weight_updater.is_rollout_engines_fresh():
if has_new_engines or self.weight_updater.conn_status.needs_reconnect():
self.weight_updater.connect_rollout_engines(
rollout_engines,
engine_gpu_counts=engine_gpu_counts,
engine_gpu_offsets=engine_gpu_offsets,
)
self.weight_updater.conn_status.mark_reconnected()
dist.barrier(group=get_gloo_group())

if self.args.debug_skip_weight_update:
Expand Down Expand Up @@ -910,4 +911,4 @@ def reconfigure_indep_dp(self, indep_dp_info: IndepDPInfo) -> None:
megatron_rank=dist.get_rank(),
megatron_world_size=dist.get_world_size(),
)
self.weight_updater.mark_engine_connection_stale()
self.weight_updater.conn_status.mark_trainer_stale()
14 changes: 14 additions & 0 deletions miles/backends/training_utils/conn_status.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
class ConnStatusManager:
def __init__(self) -> None:
self._initialized: bool = False
self._trainer_stale: bool = False

def needs_reconnect(self) -> bool:
return (not self._initialized) or self._trainer_stale

def mark_trainer_stale(self) -> None:
self._trainer_stale = True

def mark_reconnected(self) -> None:
self._initialized = True
self._trainer_stale = False
7 changes: 0 additions & 7 deletions miles/backends/training_utils/weight_update/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@ class WeightTransferProtocol(ABC):
def __init__(self, args: Namespace) -> None:
self.args = args
self.rollout_engines: Sequence[SGLangApiClient] | None = None
self._connection_stale = False
self.is_sender: bool | None = None
self.group_name = "miles"
self.update_weight_metrics: dict[str, float] = {}
Expand Down Expand Up @@ -63,12 +62,6 @@ def after_base_weights(self) -> None: # noqa: B027 — optional hook
def finalize(self, weight_version: int) -> None: # noqa: B027 — optional hook
"""Hook after all sends (e.g. publish + engine reload)."""

def is_fresh(self) -> bool:
return self.rollout_engines is not None and not self._connection_stale

def mark_stale(self) -> None:
self._connection_stale = True

def pop_metrics(self) -> dict[str, float]:
metrics, self.update_weight_metrics = self.update_weight_metrics, {}
return metrics
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,6 @@ def connect(
Create NCCL "miles-pp_{pp_rank}" if PP source (DP=TP=0). Lock prevents concurrent broadcasts.
"""
self.rollout_engines = rollout_engines
self._connection_stale = False
self._selector = selector
self._engine_gpu_counts = engine_gpu_counts

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,6 @@ def connect(
for distributed. Map ranks to colocated IPC engines.
"""
self.rollout_engines = rollout_engines
self._connection_stale = False
self._selector = selector

if engine_gpu_counts is None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,6 @@ def connect(
# uses isn't needed either — the engine-side apply is serialized by a per-host flock
# behind /pull_weights.
self.rollout_engines = rollout_engines
self._connection_stale = False
self.group_name = "miles-disk-delta"
replica_rank, _ = get_data_replica_rank_and_size(parallel_state, placement)
self.is_sender = replica_rank == 0
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,6 @@ def connect(
weight format conversion before transfer.
"""
self.rollout_engines = rollout_engines
self._connection_stale = False

self.is_sender = self.transfer_plan._gathered_dp_rank < self.transfer_plan._rollout_num_gpus

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@

import ray
import torch
from mooncake.engine import TransferEngine
from sglang.srt.server_args import ServerArgs
from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient
from miles.backends.training_utils.parallel import get_parallel_state
Expand Down Expand Up @@ -202,6 +201,8 @@ def register_cpu_memory(params_dict: dict, transfer_engine) -> dict:


def create_transfer_engine():
from mooncake.engine import TransferEngine

transfer_engine = TransferEngine()
local_ip = ray._private.services.get_node_ip_address()
transfer_engine.initialize(local_ip, "P2PHANDSHAKE", "rdma", "")
Expand Down
8 changes: 2 additions & 6 deletions miles/backends/training_utils/weight_update/updater.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from tqdm import tqdm

from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient
from miles.backends.training_utils.conn_status import ConnStatusManager
from miles.backends.training_utils.parallel import ParallelState
from miles.backends.training_utils.weight_update.protocol import get_weight_transfer_protocol
from miles.backends.training_utils.weight_update.session import (
Expand Down Expand Up @@ -51,6 +52,7 @@ def __init__(
self.args = args
self.parallel_state = parallel_state
self.protocol = get_weight_transfer_protocol(args)
self.conn_status = ConnStatusManager()
assert (
not is_lora or self.protocol.supports_lora
), f"LoRA weight sync is not supported for {args.update_weight_transfer_mode!r} weight transfer."
Expand Down Expand Up @@ -88,12 +90,6 @@ def connect_rollout_engines(
assert self.protocol.is_sender is not None, "connect() must set is_sender"
self._registered_adapters.clear()

def is_rollout_engines_fresh(self) -> bool:
return self.protocol.is_fresh()

def mark_engine_connection_stale(self) -> None:
self.protocol.mark_stale()

def pop_metrics(self) -> dict[str, float]:
"""Return and clear the protocol's metrics; the actor drains them onto the step log."""
return self.protocol.pop_metrics()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@

from miles.utils.replay_base import IndexerReplayManager, RoutingReplayManager

from miles.backends.training_utils.conn_status import ConnStatusManager


@pytest.fixture(scope="module")
def actor_module():
Expand Down Expand Up @@ -186,7 +188,8 @@ def test_update_weights_only_uses_temporary_process_groups_when_asleep(actor_mod
worker._asleep = asleep
worker._heartbeat = Mock()
worker.weight_updater = Mock()
worker.weight_updater.is_rollout_engines_fresh.return_value = True
worker.weight_updater.conn_status = Mock(spec=ConnStatusManager)
worker.weight_updater.conn_status.needs_reconnect.return_value = False
info = Namespace(
engine_gpu_counts=[],
engine_gpu_offsets=[],
Expand Down
48 changes: 48 additions & 0 deletions tests/fast/backends/training_utils/test_conn_status.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
from miles.backends.training_utils.conn_status import ConnStatusManager


def test_fresh_manager_requires_an_initial_connect() -> None:
"""A never-connected manager reports that a reconnect is still required."""
manager = ConnStatusManager()

assert manager.needs_reconnect() is True


def test_reconnect_is_not_repeated_after_a_successful_connect() -> None:
"""Once marked reconnected, the manager stops asking for further reconnects."""
manager = ConnStatusManager()

manager.mark_reconnected()

assert manager.needs_reconnect() is False


def test_marking_the_trainer_stale_forces_another_reconnect() -> None:
"""A stale trainer makes an already-connected manager require a reconnect again."""
manager = ConnStatusManager()
manager.mark_reconnected()

manager.mark_trainer_stale()

assert manager.needs_reconnect() is True


def test_reconnecting_clears_the_stale_trainer_flag() -> None:
"""Reconnecting after a stale trainer clears the flag so later windows skip reconnects."""
manager = ConnStatusManager()
manager.mark_reconnected()
manager.mark_trainer_stale()

manager.mark_reconnected()

assert manager.needs_reconnect() is False
assert manager.needs_reconnect() is False


def test_marking_the_trainer_stale_before_any_connect_still_requires_reconnect() -> None:
"""Staleness on a never-connected manager keeps the reconnect requirement set."""
manager = ConnStatusManager()

manager.mark_trainer_stale()

assert manager.needs_reconnect() is True
184 changes: 184 additions & 0 deletions tests/fast/backends/training_utils/weight_update/test_broadcast.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,184 @@
import asyncio
import threading
from argparse import Namespace
from unittest.mock import MagicMock, patch

import pytest
import torch

from miles.backends.training_utils.weight_update.protocols.broadcast import (
connect_rollout_engines_from_distributed,
disconnect_rollout_engines_from_distributed,
update_weights_from_distributed,
)
from miles.utils import async_utils

_BROADCAST_MODULE = "miles.backends.training_utils.weight_update.protocols.broadcast"

_BASE_NAMED_TENSORS = [
("model.layers.0.mlp.gate_proj.weight", torch.zeros(2, 3)),
("model.layers.0.mlp.up_proj.weight", torch.arange(6, dtype=torch.float32).reshape(2, 3).t()),
]


class _GatedEngine:
def __init__(self, started: threading.Semaphore, release: threading.Event) -> None:
self.calls: list[tuple[str, tuple, dict]] = []
self._started = started
self._release = release

def __getattr__(self, name: str):
async def method(*args, **kwargs):
self.calls.append((name, args, kwargs))
self._started.release()
released = await asyncio.get_running_loop().run_in_executor(None, self._release.wait, 5)
if not released:
raise AssertionError("the caller waited for an engine before reaching the collective")
return {"success": True}

return method


class _AcceptingEngine:
async def init_weights_update_group(self, *args, **kwargs) -> dict:
return {"success": True}


class _RefusingEngine:
async def init_weights_update_group(self, *args, **kwargs) -> None:
raise RuntimeError("engine refused the group")


class _SlowTeardownEngine:
def __init__(self) -> None:
self.calls: list[tuple[str, str]] = []
self.finished = False

async def destroy_weights_update_group(self, group_name: str) -> dict:
self.calls.append(("destroy_weights_update_group", group_name))
await asyncio.sleep(0.05)
self.finished = True
return {"success": True}


class _RecordingHandle:
def __init__(self) -> None:
self.waited = False

def wait(self) -> None:
self.waited = True


class TestConnectRolloutEnginesFromDistributed:
def test_group_init_starts_all_engines_before_local_join_with_heterogeneous_offsets(self) -> None:
"""Every engine is asked to join, at its own rank offset, before rank zero blocks on the handshake."""
started = threading.Semaphore(0)
release = threading.Event()
engines = [_GatedEngine(started, release) for _ in range(3)]
group = MagicMock(name="nccl_group")

def join(**kwargs):
for _ in engines:
assert started.acquire(timeout=30), "an engine had not been asked before the local join"
release.set()
return group

with (
patch(f"{_BROADCAST_MODULE}.ray") as ray_mock,
patch(f"{_BROADCAST_MODULE}.init_process_group", side_effect=join) as init_process_group,
):
ray_mock._private.services.get_node_ip_address.return_value = "10.0.0.1"
result = connect_rollout_engines_from_distributed(
Namespace(),
"miles-pp_0",
engines,
engine_gpu_counts=[2, 4, 1],
)

assert result is group
master_port = engines[0].calls[0][1][1]
assert [engine.calls for engine in engines] == [
[("init_weights_update_group", ("10.0.0.1", master_port, 1, 8, "miles-pp_0"), {"backend": "nccl"})],
[("init_weights_update_group", ("10.0.0.1", master_port, 3, 8, "miles-pp_0"), {"backend": "nccl"})],
[("init_weights_update_group", ("10.0.0.1", master_port, 7, 8, "miles-pp_0"), {"backend": "nccl"})],
]
assert init_process_group.call_args.kwargs == {
"backend": "nccl",
"init_method": f"tcp://10.0.0.1:{master_port}",
"world_size": 8,
"rank": 0,
"group_name": "miles-pp_0",
}

def test_an_engine_that_refuses_the_group_fails_the_connect(self) -> None:
"""The submitted joins are awaited, so a refusing engine surfaces instead of being dropped."""
with (
patch(f"{_BROADCAST_MODULE}.ray") as ray_mock,
patch(f"{_BROADCAST_MODULE}.init_process_group", return_value=MagicMock(name="nccl_group")),
):
ray_mock._private.services.get_node_ip_address.return_value = "10.0.0.1"
with pytest.raises(RuntimeError, match="engine refused the group"):
connect_rollout_engines_from_distributed(
Namespace(rollout_num_gpus_per_engine=2),
"miles-pp_0",
[_AcceptingEngine(), _AcceptingEngine(), _RefusingEngine()],
)


class TestDisconnectRolloutEnginesFromDistributed:
def test_disconnect_awaits_engine_cleanup_when_local_destroy_fails(self) -> None:
"""A failed local teardown still surfaces, and it must not abandon the engine-side teardown."""
engines = [_SlowTeardownEngine(), _SlowTeardownEngine()]

with patch(f"{_BROADCAST_MODULE}.dist") as dist_mock:
dist_mock.destroy_process_group.side_effect = RuntimeError("nccl teardown failed")
with pytest.raises(RuntimeError, match="nccl teardown failed"):
disconnect_rollout_engines_from_distributed(
Namespace(),
"miles-pp_0",
MagicMock(name="nccl_group"),
engines,
)

assert [engine.calls for engine in engines] == [
[("destroy_weights_update_group", "miles-pp_0")],
[("destroy_weights_update_group", "miles-pp_0")],
]
assert [engine.finished for engine in engines] == [True, True]


class TestUpdateWeightsFromDistributed:
def test_base_metadata_is_in_flight_before_broadcast_and_all_handles_finish(self) -> None:
"""No engine is still unasked when the first tensor goes out, every tensor is staged contiguously, and no handle is left unwaited."""
started = threading.Semaphore(0)
release = threading.Event()
engines = [_GatedEngine(started, release) for _ in range(3)]
handles: list[_RecordingHandle] = []
broadcast_tensors: list[torch.Tensor] = []

def broadcast(tensor, src, group=None, async_op=False):
broadcast_tensors.append(tensor)
if len(broadcast_tensors) == 1:
for _ in engines:
assert started.acquire(timeout=30), "an engine had not been asked before the first broadcast"
release.set()
handles.append(_RecordingHandle())
return handles[-1]

with patch(f"{_BROADCAST_MODULE}.dist") as dist_mock:
dist_mock.broadcast.side_effect = broadcast
futures = update_weights_from_distributed(
"miles-pp_0",
MagicMock(name="nccl_group"),
engines,
_BASE_NAMED_TENSORS,
selector="target",
)
async_utils.wait_futures(futures)

assert [len(engine.calls) for engine in engines] == [1, 1, 1]
assert engines[0].calls[0][2]["names"] == [name for name, _ in _BASE_NAMED_TENSORS]
assert engines[0].calls[0][2]["selector"] == "target"
assert [handle.waited for handle in handles] == [True, True]
assert [tensor.is_contiguous() for tensor in broadcast_tensors] == [True, True]
assert torch.equal(broadcast_tensors[1], _BASE_NAMED_TENSORS[1][1])
Loading
Loading