diff --git a/docs/en/developer_guide/migration.md b/docs/en/developer_guide/migration.md new file mode 100644 index 00000000000..6b52832c461 --- /dev/null +++ b/docs/en/developer_guide/migration.md @@ -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) +``` diff --git a/miles/ray/actor_group.py b/miles/ray/actor_group.py index 265814664cf..54228f42282 100644 --- a/miles/ray/actor_group.py +++ b/miles/ray/actor_group.py @@ -1,3 +1,4 @@ +import asyncio import os import ray @@ -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. @@ -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) diff --git a/miles/ray/placement_group.py b/miles/ray/placement_group.py index 443d2f7fbbd..bf8244b90ab 100644 --- a/miles/ray/placement_group.py +++ b/miles/ray/placement_group.py @@ -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 @@ -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, @@ -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 diff --git a/tests/fast/utils/test_async_utils.py b/tests/fast/utils/test_async_utils.py new file mode 100644 index 00000000000..145ac10fbc7 --- /dev/null +++ b/tests/fast/utils/test_async_utils.py @@ -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"} diff --git a/train.py b/train.py index bd22ba56135..9de6c30e0dd 100644 --- a/train.py +++ b/train.py @@ -1,14 +1,16 @@ -import ray +import asyncio + from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH, GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_WEIGHTS from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from miles.utils.arguments import parse_args +from miles.utils.async_utils import eager_create_task from miles.utils.logging_utils import configure_logger from miles.utils.misc import should_run_periodic_action from miles.utils.tracking_utils import init_tracking -def train(args): +async def train(args): configure_logger() # allocate the GPUs pgs = create_placement_groups(args) @@ -19,56 +21,56 @@ def train(args): rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) # create the actor and critic models - actor_model, critic_model = create_training_models(args, pgs, rollout_manager) + actor_model, critic_model = await create_training_models(args, pgs, rollout_manager) if args.offload_rollout: - ray.get(rollout_manager.onload_weights.remote()) + await rollout_manager.onload_weights.remote() # always update weight first so that sglang has the loaded weights from training. - actor_model.update_weights() + await actor_model.update_weights() if args.check_weight_update_equal: - ray.get(rollout_manager.check_weights.remote(action="compare")) + await rollout_manager.check_weights.remote(action="compare") if args.offload_rollout: - ray.get(rollout_manager.onload_kv.remote()) + await rollout_manager.onload_kv.remote() # special case for eval-only if args.num_rollout == 0 and args.eval_interval is not None: - ray.get(rollout_manager.eval.remote(rollout_id=0)) + await rollout_manager.eval.remote(rollout_id=0) - def offload_train(): + async def offload_train(): if args.offload_train: if args.use_critic: - critic_model.offload() + await critic_model.offload() if rollout_id >= args.num_critic_only_steps: - actor_model.offload() + await actor_model.offload() else: - actor_model.offload() + await actor_model.offload() else: - actor_model.clear_memory() + await actor_model.clear_memory() - def save(rollout_id): + async def save(rollout_id): if (not args.use_critic) or (rollout_id >= args.num_critic_only_steps): - actor_model.save_model( + await actor_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, ) if args.use_critic: - critic_model.save_model( + await critic_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, ) if args.rollout_global_dataset: - ray.get(rollout_manager.save.remote(rollout_id)) + await rollout_manager.save.remote(rollout_id) # train loop. # note that for async training, one can change the position of the sync operation(ray.get). for rollout_id in range(args.start_rollout_id, args.num_rollout): if args.eval_interval is not None and rollout_id == 0 and not args.skip_eval_before_train: - ray.get(rollout_manager.eval.remote(rollout_id)) + await rollout_manager.eval.remote(rollout_id) - rollout_data_ref = ray.get(rollout_manager.generate.remote(rollout_id)) + rollout_data_ref = await rollout_manager.generate.remote(rollout_id) if args.offload_rollout: offload_tags = [GPU_MEMORY_TYPE_CUDA_GRAPH] @@ -76,32 +78,32 @@ def save(rollout_id): offload_tags.append(GPU_MEMORY_TYPE_KV_CACHE) if "weight" in args.offload_rollout_level: offload_tags.append(GPU_MEMORY_TYPE_WEIGHTS) - ray.get(rollout_manager.offload.remote(tags=offload_tags)) + await rollout_manager.offload.remote(tags=offload_tags) if args.use_critic: - critic_train_handle = critic_model.async_train(rollout_id, rollout_data_ref) + critic_task = await eager_create_task(critic_model.train(rollout_id, rollout_data_ref)) if rollout_id >= args.num_critic_only_steps: - ray.get(actor_model.async_train(rollout_id, rollout_data_ref)) - ray.get(critic_train_handle) + await actor_model.train(rollout_id, rollout_data_ref) + await critic_task else: - ray.get(actor_model.async_train(rollout_id, rollout_data_ref)) + await actor_model.train(rollout_id, rollout_data_ref) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): - save(rollout_id) + await save(rollout_id) - offload_train() + await offload_train() if args.offload_rollout: - ray.get(rollout_manager.onload_weights.remote()) - actor_model.update_weights() + await rollout_manager.onload_weights.remote() + await actor_model.update_weights() if args.offload_rollout: - ray.get(rollout_manager.onload_kv.remote()) + await rollout_manager.onload_kv.remote() if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch): - ray.get(rollout_manager.eval.remote(rollout_id)) + await rollout_manager.eval.remote(rollout_id) - ray.get(rollout_manager.dispose.remote()) + await rollout_manager.dispose.remote() if __name__ == "__main__": args = parse_args() - train(args) + asyncio.run(train(args)) diff --git a/train_async.py b/train_async.py index bef1d98abe8..e9e05a40629 100644 --- a/train_async.py +++ b/train_async.py @@ -1,14 +1,15 @@ -import ray +import asyncio from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from miles.utils.arguments import parse_args +from miles.utils.async_utils import eager_create_task from miles.utils.logging_utils import configure_logger from miles.utils.misc import should_run_periodic_action from miles.utils.tracking_utils import init_tracking # The framework supports other asynchronous approaches such as fully async (which is shown in examples/full_async). -def train(args): +async def train(args): assert not args.colocate, "Colocation is not supported for async training." configure_logger() # allocate the GPUs @@ -20,58 +21,58 @@ def train(args): rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) # create the actor and critic models - actor_model, critic_model = create_training_models(args, pgs, rollout_manager) + actor_model, critic_model = await create_training_models(args, pgs, rollout_manager) # always update weight first so that sglang has the loaded weights from training. - actor_model.update_weights() + await actor_model.update_weights() if args.check_weight_update_equal: - ray.get(rollout_manager.check_weights.remote(action="compare")) + await rollout_manager.check_weights.remote(action="compare") # async train loop. rollout_data_next_future = rollout_manager.generate.remote(args.start_rollout_id) for rollout_id in range(args.start_rollout_id, args.num_rollout): # Sync the last generation if rollout_data_next_future is not None: - rollout_data_curr_ref = ray.get(rollout_data_next_future) + rollout_data_curr_ref = await rollout_data_next_future # Start the next rollout early. if rollout_id + 1 < args.num_rollout: rollout_data_next_future = rollout_manager.generate.remote(rollout_id + 1) if args.use_critic: - critic_train_handle = critic_model.async_train(rollout_id, rollout_data_curr_ref) + critic_task = await eager_create_task(critic_model.train(rollout_id, rollout_data_curr_ref)) if rollout_id >= args.num_critic_only_steps: - ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref)) - ray.get(critic_train_handle) + await actor_model.train(rollout_id, rollout_data_curr_ref) + await critic_task else: - ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref)) + await actor_model.train(rollout_id, rollout_data_curr_ref) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): - actor_model.save_model( + await actor_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, ) if args.use_critic: - critic_model.save_model( + await critic_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, ) if args.rollout_global_dataset: - ray.get(rollout_manager.save.remote(rollout_id)) + await rollout_manager.save.remote(rollout_id) if (rollout_id + 1) % args.update_weights_interval == 0: # sync generate before update weights to prevent update weight in the middle of generation - rollout_data_curr_ref = ray.get(x) if (x := rollout_data_next_future) is not None else None + rollout_data_curr_ref = (await x) if (x := rollout_data_next_future) is not None else None rollout_data_next_future = None - actor_model.update_weights() + await actor_model.update_weights() if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch): - ray.get(rollout_manager.eval.remote(rollout_id)) + await rollout_manager.eval.remote(rollout_id) - ray.get(rollout_manager.dispose.remote()) + await rollout_manager.dispose.remote() if __name__ == "__main__": args = parse_args() - train(args) + asyncio.run(train(args))