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: 2 additions & 3 deletions miles/backends/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
from contextlib import ExitStack
from typing import TYPE_CHECKING

import ray
import torch
import torch.distributed as dist
from tqdm import tqdm
Expand All @@ -22,7 +21,7 @@
from miles.backends.training_utils.loss import compute_advantages_and_returns, get_log_probs_and_entropy, loss_function
from miles.backends.training_utils.parallel import get_parallel_state, set_parallel_state
from miles.ray.train_actor import TrainRayActor
from miles.utils import train_dump_utils, train_metric_utils
from miles.utils import async_utils, train_dump_utils, train_metric_utils
from miles.utils.context_utils import with_defer
from miles.utils.distributed_utils import get_gloo_group
from miles.utils.flops_utils import flops_args_from_hf_config, fwd_tflops_per_gpu
Expand Down Expand Up @@ -632,7 +631,7 @@ def update_weights(self, info: "EnginesAndLock") -> None: # type: ignore[overri

if self.args.ci_test and len(rollout_engines) > 0:
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
engine_version = async_utils.run(engine.get_weight_version())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
Expand Down
87 changes: 51 additions & 36 deletions miles/backends/fsdp_utils/update_weight_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@
import socket
from argparse import Namespace
from collections.abc import Sequence
from typing import TYPE_CHECKING

import ray
import torch
import torch.distributed as dist
from ray.actor import ActorHandle
from torch.distributed.tensor import DTensor

try:
Expand All @@ -17,8 +17,14 @@

from sglang.srt.utils import MultiprocessingSerializer

from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient
from miles.utils import async_utils
from miles.utils.distributed_utils import get_gloo_group, init_process_group

if TYPE_CHECKING:
from ray.actor import ActorHandle


try:
from sglang.srt.weight_sync.tensor_bucket import FlattenedTensorBucket # type: ignore[import]
except ImportError:
Expand Down Expand Up @@ -60,8 +66,8 @@ def __init__(self, args: Namespace, model: torch.nn.Module) -> None:
@abc.abstractmethod
def connect_rollout_engines(
self,
rollout_engines: Sequence[ActorHandle],
rollout_engine_lock: ActorHandle | None,
rollout_engines: Sequence[SGLangApiClient],
rollout_engine_lock: "ActorHandle | None",
engine_gpu_counts: Sequence[int] | None = None,
engine_gpu_offsets: Sequence[int] | None = None,
) -> None:
Expand All @@ -71,10 +77,13 @@ def update_weights(self) -> None:
self.weight_version += 1

if dist.get_rank() == 0:
futures = [engine.pause_generation.remote() for engine in self.rollout_engines]
futures.extend([engine.flush_cache.remote() for engine in self.rollout_engines])
ray.get(futures)
ray.get([engine.begin_weight_update.remote() for engine in self.rollout_engines])
async_utils.wait_futures(
[async_utils.submit(client.pause_generation()) for client in self.rollout_engines]
)
async_utils.wait_futures([async_utils.submit(client.flush_cache()) for client in self.rollout_engines])

@guapisolo guapisolo Aug 24, 2026

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[non-blocking] potential perf issue? My codex said.
original cost = max_i(pause_i + flush_i)
new cost = max_i(pause_i) + max_i(flush_i)
Seems trivial. non-blocking

async_utils.wait_futures(
[async_utils.submit(client.begin_weight_update()) for client in self.rollout_engines]
)
dist.barrier(group=get_gloo_group())

bucket = []
Expand Down Expand Up @@ -104,8 +113,12 @@ def update_weights(self) -> None:

dist.barrier(group=get_gloo_group())
if dist.get_rank() == 0:
ray.get([engine.end_weight_update.remote() for engine in self.rollout_engines])
ray.get([engine.continue_generation.remote() for engine in self.rollout_engines])
async_utils.wait_futures(
[async_utils.submit(client.end_weight_update()) for client in self.rollout_engines]
)
async_utils.wait_futures(
[async_utils.submit(client.continue_generation()) for client in self.rollout_engines]
)
dist.barrier(group=get_gloo_group())

def wait_and_update_bucket_weights(self, bucket):
Expand All @@ -129,8 +142,8 @@ class UpdateWeightFromTensor(UpdateWeight):

def connect_rollout_engines(
self,
rollout_engines: Sequence[ActorHandle],
rollout_engine_lock: ActorHandle | None,
rollout_engines: Sequence[SGLangApiClient],
rollout_engine_lock: "ActorHandle | None",
engine_gpu_counts: Sequence[int] | None = None,
engine_gpu_offsets: Sequence[int] | None = None,
) -> None:
Expand Down Expand Up @@ -197,8 +210,7 @@ def update_bucket_weights(self, named_tensors, weight_version=None) -> None:
"flush_cache": False,
"weight_version": str(weight_version),
}
ref = self._ipc_engine.update_weights_from_tensor.remote(**kwargs)
result = ray.get(ref)
result = async_utils.run(self._ipc_engine.update_weights_from_tensor(**kwargs))
if isinstance(result, dict):
success = result.get("success", True)
error_msg = result.get("error_message") or result.get("message", "unknown error")
Expand All @@ -211,17 +223,16 @@ def update_bucket_weights(self, named_tensors, weight_version=None) -> None:
)

if dist.get_rank() == self._ipc_gather_src:
ref = self._ipc_engine.flush_cache.remote()
ray.get(ref)
async_utils.run(self._ipc_engine.flush_cache())


class UpdateWeightFromDistributed(UpdateWeight):
"""Broadcast weights via a temporary NCCL group to rollout engines."""

def connect_rollout_engines(
self,
rollout_engines: Sequence[ActorHandle],
rollout_engine_lock: ActorHandle | None,
rollout_engines: Sequence[SGLangApiClient],
rollout_engine_lock: "ActorHandle | None",
engine_gpu_counts: Sequence[int] | None = None,
engine_gpu_offsets: Sequence[int] | None = None,
) -> None:
Expand All @@ -240,16 +251,18 @@ def connect_rollout_engines(
# +1 for the trainer's source rank (rank 0); rollout engine ranks start at 1
world_size = self.args.rollout_num_gpus + 1

refs = [
engine.init_weights_update_group.remote(
master_address,
master_port,
i * self.args.rollout_num_gpus_per_engine + 1,
world_size,
self._group_name,
backend="nccl",
futures = [
async_utils.submit(
api_client.init_weights_update_group(
master_address,
master_port,
i * self.args.rollout_num_gpus_per_engine + 1,
world_size,
self._group_name,
backend="nccl",
)
)
for i, engine in enumerate(self.rollout_engines)
for i, api_client in enumerate(self.rollout_engines)
]
self._model_update_groups = init_process_group(
backend="nccl",
Expand All @@ -258,23 +271,25 @@ def connect_rollout_engines(
rank=0,
group_name=self._group_name,
)
ray.get(refs)
async_utils.wait_futures(futures)

def update_bucket_weights(self, named_tensors, weight_version=None) -> None:
"""Send names/dtypes/shapes metadata to engines, then broadcast the tensors (contiguous;
DTensors materialized when world_size == 1)."""
if not self._is_src_rank or not named_tensors:
return

refs = [
engine.update_weights_from_distributed.remote(
names=[name for name, _ in named_tensors],
dtypes=[param.dtype for _, param in named_tensors],
shapes=[param.shape for _, param in named_tensors],
group_name=self._group_name,
weight_version=str(weight_version),
futures = [
async_utils.submit(
client.update_weights_from_distributed(
names=[name for name, _ in named_tensors],
dtypes=[param.dtype for _, param in named_tensors],
shapes=[param.shape for _, param in named_tensors],
group_name=self._group_name,
weight_version=str(weight_version),
)
)
for engine in self.rollout_engines
for client in self.rollout_engines
]

handles = []
Expand All @@ -291,4 +306,4 @@ def update_bucket_weights(self, named_tensors, weight_version=None) -> None:

for handle in handles:
handle.wait()
ray.get(refs)
async_utils.wait_futures(futures)
4 changes: 2 additions & 2 deletions miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
from miles.backends.megatron_utils.rematerialize_utils import build_main_cast_context
from miles.dashboard import hooks as dashboard_hooks
from miles.ray.train_actor import TrainRayActor
from miles.utils import train_dump_utils
from miles.utils import async_utils, train_dump_utils
from miles.utils.argparse_utils import inplace_modify_args
from miles.utils.audit_utils.event_logger.logger import event_logger_context
from miles.utils.audit_utils.witness.allocator import WitnessInfo
Expand Down Expand Up @@ -831,7 +831,7 @@ def update_weights(self, info: "EnginesAndLock") -> None:

if self.args.ci_test and len(rollout_engines) > 0 and not is_lora_enabled(self.args):
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
engine_version = async_utils.run(engine.get_weight_version())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
Expand Down
Loading
Loading