Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
73 commits
Select commit Hold shift + click to select a range
ecb80b5
fix: rename dp_src_rank -> dp_cp_src_rank
fzyzcjy Mar 30, 2026
f447fd9
feat: add global ParallelState + compute VPP from args + move init ea…
fzyzcjy Mar 30, 2026
eb72a7d
fmt: auto-format after refactor
fzyzcjy Mar 30, 2026
b1f7b92
fix: move local import to top-level, remove docstring, reorder functions
fzyzcjy Mar 30, 2026
c45000d
fix: simplify _compute_vpp_fields, always use pipeline_model_parallel…
fzyzcjy Mar 30, 2026
88e5828
fix: merge adjacent f-strings in assert message
fzyzcjy Mar 31, 2026
d402488
fix: apply black formatting to assert statement in parallel.py
fzyzcjy Apr 1, 2026
8cd7719
refactor: migrate mpu.get_data_parallel_* to get_parallel_state() + b…
fzyzcjy Mar 30, 2026
04baa3c
fmt: auto-format after migration
fzyzcjy Mar 30, 2026
b509b61
refactor: migrate mpu.get_data_parallel_* to get_parallel_state() + b…
fzyzcjy Mar 30, 2026
e5c8705
fmt: auto-format after migration
fzyzcjy Mar 30, 2026
96faeef
refactor: rename dp_* -> intra_dp_* on ParallelState fields
fzyzcjy Mar 30, 2026
615f9d5
fmt: auto-format after rename
fzyzcjy Mar 30, 2026
5cfd41c
refactor: rename dp_* -> intra_dp_* on ParallelState fields
fzyzcjy Mar 30, 2026
038c16d
fmt: auto-format after rename
fzyzcjy Mar 30, 2026
c32d955
refactor: introduce GroupInfo dataclass for ParallelState
fzyzcjy Mar 30, 2026
3815263
fix: revert incorrect renames of self.tp_rank and self.cp_group (not …
fzyzcjy Mar 30, 2026
4e2547f
fmt
fzyzcjy Mar 30, 2026
3b3238c
more
fzyzcjy Mar 30, 2026
2eebace
fix
fzyzcjy Mar 30, 2026
5de0831
temp revert
fzyzcjy Mar 30, 2026
19de975
Revert "temp revert"
fzyzcjy Mar 30, 2026
819a81b
feat: add GroupInfo __post_init__ validation for rank/size against pr…
fzyzcjy Mar 30, 2026
a2938b5
refactor: remove src_rank from GroupInfo, use dst=0 in gather_object
fzyzcjy Mar 30, 2026
fda4443
Revert "refactor: remove src_rank from GroupInfo, use dst=0 in gather…
fzyzcjy Mar 30, 2026
ea7a8dd
Revert "Revert "refactor: remove src_rank from GroupInfo, use dst=0 i…
fzyzcjy Mar 30, 2026
eb58bb9
refactor: remove src_rank from GroupInfo, derive from group at call site
fzyzcjy Mar 30, 2026
c79b720
refactor
fzyzcjy Mar 30, 2026
f30b819
revert to /4
fzyzcjy Mar 30, 2026
1251132
refactor
fzyzcjy Mar 30, 2026
c3294a4
Revert "refactor"
fzyzcjy Mar 30, 2026
8b1bc82
Revert "revert to /4"
fzyzcjy Mar 30, 2026
5f44dae
cherry pick ci update
fzyzcjy Mar 31, 2026
15f87fe
[async-refactor] refactor: convert RayTrainGroup methods to async
fzyzcjy Mar 31, 2026
581fec8
[async-refactor] refactor: convert create_training_models to async
fzyzcjy Mar 31, 2026
4598caa
[async-refactor] refactor: convert train loops to async
fzyzcjy Mar 31, 2026
1c1ac4e
[async-refactor] fix: update stale comments after async conversion
fzyzcjy Apr 1, 2026
8d26e22
[async-refactor] simplify: extract _broadcast helper, remove unused i…
fzyzcjy Apr 1, 2026
37ddda4
[async-refactor] docs: add async migration guide for custom train loops
fzyzcjy Apr 1, 2026
b7536ba
[async-refactor] feat: crash on unawaited coroutines via warning-to-e…
fzyzcjy Apr 1, 2026
3e4a7af
[async-refactor] feat: extract configure_raise_unawaited_coroutine wi…
fzyzcjy Apr 1, 2026
1aaa7cd
[async-refactor] docs: simplify migration guide
fzyzcjy Apr 1, 2026
eb4e9ec
[async-refactor] fix: add sleep(0) after create_task to ensure critic…
fzyzcjy Apr 1, 2026
0cebcd5
[async-refactor] refactor: extract dispatch_async_task, rename test file
fzyzcjy Apr 1, 2026
b68920b
[async-refactor] refactor: rename dispatch_async_task to eager_create…
fzyzcjy Apr 1, 2026
9137281
fmt
fzyzcjy Apr 1, 2026
6d04c9b
more
fzyzcjy Apr 1, 2026
82097f1
[async-refactor] feat: catch destroyed-pending-task warning, rename t…
fzyzcjy Apr 1, 2026
321dbe6
[async-refactor] docs: add motivation and simplify migration guide
fzyzcjy Apr 1, 2026
3c53023
more
fzyzcjy Apr 1, 2026
6faacfd
more
fzyzcjy Apr 1, 2026
a9818b8
[async-refactor] docs: fix step 3 description in migration guide
fzyzcjy Apr 1, 2026
97caf56
more
fzyzcjy Apr 1, 2026
b5b1c9c
[async-refactor] docs: merge Rule A and B into single step
fzyzcjy Apr 1, 2026
1b961ca
[async-refactor] docs: mention adding await on previously sync group …
fzyzcjy Apr 1, 2026
22f38d9
[async-refactor] docs: remove import line from migration example
fzyzcjy Apr 1, 2026
5deb30a
[async-refactor] docs: side-by-side layout for dispatch handles example
fzyzcjy Apr 1, 2026
519784a
more
fzyzcjy Apr 1, 2026
8259d46
fmt
fzyzcjy Apr 1, 2026
22856de
[async-refactor] fix: restore original docstrings and comments to min…
fzyzcjy Apr 1, 2026
3664205
[async-refactor] fix: trim eager_create_task docstring
fzyzcjy Apr 1, 2026
530844a
[async-refactor] test: parametrize eager vs plain create_task to show…
fzyzcjy Apr 1, 2026
e10b0ac
[async-refactor] fix: parametrize by string to avoid async wrapper, f…
fzyzcjy Apr 1, 2026
419b97d
[async-refactor] fix: test warning filters directly instead of relyin…
fzyzcjy Apr 1, 2026
8c51f7a
Revert "[async-refactor] fix: test warning filters directly instead o…
fzyzcjy Apr 1, 2026
1bd553e
[async-refactor] fix: use sys.unraisablehook + os._exit for real cras…
fzyzcjy Apr 1, 2026
841658e
fmt
fzyzcjy Apr 1, 2026
a6906bf
fix: add pytest-asyncio to CI and fix async test issues
fzyzcjy Apr 1, 2026
dc74b1d
feat: add pytest-asyncio to test extras in setup.py
fzyzcjy Apr 1, 2026
ec4bb42
fix: move pytest-asyncio to default dependencies in requirements.txt
fzyzcjy Apr 1, 2026
35eb213
fix: remove unnecessary asyncio marker registration
fzyzcjy Apr 1, 2026
ecf27cb
merge: merge main into feat/refactor_dp/7
fzyzcjy Apr 3, 2026
3f17886
merge: merge main into feat/refactor_dp/7
fzyzcjy Apr 3, 2026
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
50 changes: 50 additions & 0 deletions docs/en/developer_guide/migration.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
# Migration Guide

## Train Loop: Sync → Async

### What is changed

The train loop (`train.py`, `train_async.py`) and `RayTrainGroup` now use Python async/await instead of sync `ray.get()`.

### Why it is changed

Python async is more expressive than sync code with `ray.get`. As a concrete example, in fault tolerance, we need to capture ray actor results and do retries when calling `actor_model.train`, while still allowing it to be overlapped freely with `critic_model.train`. This is hard to achieve without Python async.

### How to mechanically migrate

**1. Make the train function async:**

```python
# Before # After
def train(args): async def train(args):
... ...

if __name__ == "__main__": if __name__ == "__main__":
train(parse_args()) asyncio.run(train(parse_args()))
```

**2. `ray.get(x)` → `await x`, drop the `async_` prefix, and add `await` on group methods that previously had none:**

```python
ray.get(group.async_init(...)) → await group.init(...)
ray.get(group.async_train(...)) → await group.train(...)
group.save_model(...) → await group.save_model(...)
group.update_weights() → await group.update_weights()
ray.get(rollout_manager.generate.remote(id)) → await rollout_manager.generate.remote(id)
# Same pattern for offload, onload, clear_memory, connect, set_rollout_manager
```

**3. Dispatch handles:** replace `handle = group.async_fn(...)` with `task = await eager_create_task(group.fn(...))`.

```python
# Before # After
handle = critic.async_train(...) task = await eager_create_task(critic.train(...))
ray.get(actor.async_train(...)) await actor.train(...)
ray.get(handle) await task
```

**4. `create_training_models` is now async:**

```python
actor, critic = await create_training_models(args, pgs, rollout_manager)
```
51 changes: 27 additions & 24 deletions miles/ray/actor_group.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
import os

import ray
Expand All @@ -10,7 +11,6 @@
class RayTrainGroup:
"""
A group of ray actors
Functions start with 'async' should return list of object refs

Args:
args (Namespace): Arguments for the actor group.
Expand Down Expand Up @@ -105,40 +105,43 @@ def _allocate_gpus_for_actor(self, pg, num_gpus_per_actor):

return actor_handles

def async_init(self):
async def init(self):
"""
Allocate GPU resourced and initialize model, optimizer, local ckpt, etc.
"""
return [actor.init.remote(self.args, self.role, with_ref=self.with_ref) for actor in self._actor_handles]
return await self._broadcast("init", self.args, self.role, with_ref=self.with_ref)

def async_train(self, rollout_id, rollout_data_ref):
async def train(self, rollout_id, rollout_data_ref):
"""Do one rollout training"""
return [actor.train.remote(rollout_id, rollout_data_ref) for actor in self._actor_handles]
await self._broadcast("train", rollout_id, rollout_data_ref)

def save_model(self, rollout_id, force_sync=False):
async def save_model(self, rollout_id, force_sync=False):
"""Save actor model"""
return ray.get([actor.save_model.remote(rollout_id, force_sync=force_sync) for actor in self._actor_handles])
await self._broadcast("save_model", rollout_id, force_sync=force_sync)

def update_weights(self):
async def update_weights(self):
"""Broadcast weights from rank 0 to all other ranks."""
return ray.get([actor.update_weights.remote() for actor in self._actor_handles])
await self._broadcast("update_weights")

def onload(self):
return ray.get([actor.wake_up.remote() for actor in self._actor_handles])
async def onload(self):
await self._broadcast("wake_up")

def offload(self):
return ray.get([actor.sleep.remote() for actor in self._actor_handles])
async def offload(self):
await self._broadcast("sleep")

def clear_memory(self):
return ray.get([actor.clear_memory.remote() for actor in self._actor_handles])
async def clear_memory(self):
await self._broadcast("clear_memory")

def connect(self, critic_group):
return ray.get(
[
actor.connect_actor_critic.remote(critic)
for actor, critic in zip(self._actor_handles, critic_group._actor_handles, strict=False)
]
)
async def connect(self, critic_group):
refs = [
actor.connect_actor_critic.remote(critic)
for actor, critic in zip(self._actor_handles, critic_group._actor_handles, strict=False)
]
await asyncio.gather(*refs)

def set_rollout_manager(self, rollout_manager):
return ray.get([actor.set_rollout_manager.remote(rollout_manager) for actor in self._actor_handles])
async def set_rollout_manager(self, rollout_manager):
await self._broadcast("set_rollout_manager", rollout_manager)

async def _broadcast(self, method_name: str, *args, **kwargs) -> list:
refs = [getattr(actor, method_name).remote(*args, **kwargs) for actor in self._actor_handles]
return await asyncio.gather(*refs)
16 changes: 9 additions & 7 deletions miles/ray/placement_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
from ray.util.placement_group import placement_group
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy

from miles.utils.async_utils import eager_create_task

from ..utils.ray_utils import compute_ray_pin_head_options
from .actor_group import RayTrainGroup
from .rollout import RolloutManager
Expand Down Expand Up @@ -132,7 +134,7 @@ def allocate_train_group(args, num_nodes, num_gpus_per_node, pg, role: str, with
)


def create_training_models(args, pgs, rollout_manager):
async def create_training_models(args, pgs, rollout_manager):
actor_model = allocate_train_group(
args=args,
num_nodes=args.actor_num_nodes,
Expand All @@ -150,23 +152,23 @@ def create_training_models(args, pgs, rollout_manager):
role="critic",
with_ref=False,
)
critic_init_handle = critic_model.async_init()
critic_init_task = await eager_create_task(critic_model.init())
else:
critic_model = None

start_rollout_ids = ray.get(actor_model.async_init())
start_rollout_ids = await actor_model.init()

assert len(set(start_rollout_ids)) == 1
if args.start_rollout_id is None:
args.start_rollout_id = start_rollout_ids[0]

if args.use_critic:
ray.get(critic_init_handle)
actor_model.connect(critic_model)
await critic_init_task
await actor_model.connect(critic_model)

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

return actor_model, critic_model

Expand Down
91 changes: 91 additions & 0 deletions tests/fast/utils/test_async_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
"""Tests for eager_create_task — contrast with plain asyncio.create_task."""

import asyncio

import pytest

from miles.utils.async_utils import eager_create_task


@pytest.mark.asyncio
@pytest.mark.parametrize("create_mode", ["eager", "plain"])
class TestCreateTaskComparison:
async def test_returns_asyncio_task(self, create_mode):
async def coro():
return 42

if create_mode == "eager":
task = await eager_create_task(coro())
else:
task = asyncio.create_task(coro())

assert isinstance(task, asyncio.Task)
assert await task == 42

async def test_started_before_next_line(self, create_mode):
"""eager starts immediately; plain does not."""
started = False

async def coro():
nonlocal started
started = True
await asyncio.sleep(10)

if create_mode == "eager":
task = await eager_create_task(coro())
assert started, "eager_create_task should have started the task"
else:
task = asyncio.create_task(coro())
assert not started, "plain create_task should NOT have started the task yet"

task.cancel()
with pytest.raises(asyncio.CancelledError):
await task

async def test_dispatch_order(self, create_mode):
"""eager preserves critic-before-actor dispatch order; plain reverses it."""
order: list[str] = []

async def critic():
order.append("critic")
await asyncio.sleep(0.1)

async def actor():
order.append("actor")
await asyncio.sleep(0.1)

if create_mode == "eager":
critic_task = await eager_create_task(critic())
else:
critic_task = asyncio.create_task(critic())

await actor()
await critic_task

if create_mode == "eager":
assert order == ["critic", "actor"]
else:
assert order == ["actor", "critic"]

async def test_exception_propagates(self, create_mode):
async def failing():
raise ValueError("boom")

if create_mode == "eager":
task = await eager_create_task(failing())
else:
task = asyncio.create_task(failing())

with pytest.raises(ValueError, match="boom"):
await task

async def test_result_available(self, create_mode):
async def compute():
return {"key": "value"}

if create_mode == "eager":
task = await eager_create_task(compute())
else:
task = asyncio.create_task(compute())

assert await task == {"key": "value"}
Loading
Loading