From 55e83994225aa48fe49e01a97a3cdee59d3e45aa Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Sun, 26 Jul 2026 17:33:12 +0800 Subject: [PATCH] Make the rollout-side offload/onload fan-out async to avoid blocking sync code These call sites are about to become direct HTTP calls to the engines' own servers. Turn them into `async def` + `await` now, while they still go through ray, so the switch is a one-line change per call site rather than a control-flow change. ServerGroup.offload / onload / onload_weights_from_disk stop handing their caller a list of pending handles and await their own fan-out; RolloutServer gathers over groups instead of over engines, so the whole server still offloads concurrently. --- miles/ray/rollout/rollout_server.py | 12 +-- miles/ray/rollout/server_group.py | 28 +++--- .../rollout/real_ray/test_rollout_server.py | 87 +++++++++++++++++++ 3 files changed, 107 insertions(+), 20 deletions(-) diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index 3c9ef3b1176..e6dc32615cd 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -266,16 +266,12 @@ async def recover(self): await asyncio.gather(*[g.recover(port_cursors=port_cursors) for g in self.server_groups]) async def offload(self, tags: list[str] | None = None): - handles = [] - for g in self.server_groups: - handles.extend(g.offload(tags=tags)) - return await asyncio.gather(*handles) + per_group = await asyncio.gather(*[g.offload(tags=tags) for g in self.server_groups]) + return [result for group_results in per_group for result in group_results] async def onload(self, tags: list[str] | None = None): - handles = [] - for g in self.server_groups: - handles.extend(g.onload(tags)) - return await asyncio.gather(*handles) + per_group = await asyncio.gather(*[g.onload(tags) for g in self.server_groups]) + return [result for group_results in per_group for result in group_results] async def check_weights( self, action: str, allow_quant_error: bool = False, selector: str = "all", skip_list: list[str] | None = None diff --git a/miles/ray/rollout/server_group.py b/miles/ray/rollout/server_group.py index 769be4fc372..2af8ac1da24 100644 --- a/miles/ray/rollout/server_group.py +++ b/miles/ray/rollout/server_group.py @@ -233,23 +233,27 @@ def mark_alive(self, engine_indices: list[int]): for engine_index in engine_indices: self.all_engines[engine_index].mark_alive() - def offload(self, tags: list[str] | None = None): + async def offload(self, tags: list[str] | None = None): if not self.needs_offload: return [] - return [ - engine.actor_handle.release_memory_occupation.remote(tags=tags) - for engine in self.engines - if engine.is_allocated - ] + return await asyncio.gather( + *[ + engine.actor_handle.release_memory_occupation.remote(tags=tags) + for engine in self.engines + if engine.is_allocated + ] + ) - def onload(self, tags: list[str] | None = None): + async def onload(self, tags: list[str] | None = None): if not self.needs_offload: return [] - return [ - engine.actor_handle.resume_memory_occupation.remote(tags=tags) - for engine in self.engines - if engine.is_allocated - ] + return await asyncio.gather( + *[ + engine.actor_handle.resume_memory_occupation.remote(tags=tags) + for engine in self.engines + if engine.is_allocated + ] + ) def onload_weights_from_disk(self): """Reload weights from ``model_path`` for non-updatable groups.""" diff --git a/tests/fast/ray/rollout/real_ray/test_rollout_server.py b/tests/fast/ray/rollout/real_ray/test_rollout_server.py index 197e40e1f81..dded11f2778 100644 --- a/tests/fast/ray/rollout/real_ray/test_rollout_server.py +++ b/tests/fast/ray/rollout/real_ray/test_rollout_server.py @@ -89,3 +89,90 @@ async def test_aggregates_across_groups_via_real_asyncio_gather( finally: _kill_group(a) _kill_group(b) + + +# ----------------------------- offload / onload ----------------------------- + + +@pytest.mark.asyncio +class TestOffloadOnloadAggregation: + async def test_offload_and_onload_reach_every_engine_of_every_group( + self, + patched_sglang_engine, + placement_group_factory, + ): + """Both fan out across groups and return one flat result per engine.""" + pg_a = placement_group_factory(2) + pg_b = placement_group_factory(3) + a = _build_group(pg_tuple=pg_a, num_engines=2, needs_offload=True) + b = _build_group(pg_tuple=pg_b, num_engines=3, needs_offload=True) + _start_group(a) + _start_group(b) + a.mark_alive([0, 1]) + b.mark_alive([0, 1, 2]) + + srv = RolloutServer(server_groups=[a, b]) + try: + offload_results = await srv.offload(tags=["weights"]) + onload_results = await srv.onload(["weights"]) + + assert len(offload_results) == 5 + assert len(onload_results) == 5 + + all_engines = [e for g in (a, b) for e in g.engines] + all_calls = ray.get([e.actor_handle.get_calls.remote() for e in all_engines]) + for calls in all_calls: + assert [name for name, _args, _kwargs in calls if name.endswith("_memory_occupation")] == [ + "release_memory_occupation", + "resume_memory_occupation", + ] + assert [kwargs for name, _args, kwargs in calls if name.endswith("_memory_occupation")] == [ + {"tags": ["weights"]}, + {"tags": ["weights"]}, + ] + finally: + _kill_group(a) + _kill_group(b) + + async def test_a_group_that_does_not_need_offload_is_skipped( + self, + patched_sglang_engine, + placement_group_factory, + ): + """Only the groups colocated with megatron give their memory back.""" + pg_a = placement_group_factory(2) + pg_b = placement_group_factory(2) + offloading = _build_group(pg_tuple=pg_a, num_engines=2, needs_offload=True) + resident = _build_group(pg_tuple=pg_b, num_engines=2, needs_offload=False) + _start_group(offloading) + _start_group(resident) + offloading.mark_alive([0, 1]) + resident.mark_alive([0, 1]) + + srv = RolloutServer(server_groups=[offloading, resident]) + try: + assert len(await srv.offload(tags=None)) == 2 + + resident_calls = ray.get([e.actor_handle.get_calls.remote() for e in resident.engines]) + assert all(not [c for c in calls if c[0] == "release_memory_occupation"] for calls in resident_calls) + finally: + _kill_group(offloading) + _kill_group(resident) + + async def test_a_dead_engine_is_not_addressed( + self, + patched_sglang_engine, + placement_group_factory, + ): + """Offload must not block forever on an engine the group already gave up on.""" + pg = placement_group_factory(2) + group = _build_group(pg_tuple=pg, num_engines=2, needs_offload=True) + _start_group(group) + group.mark_alive([0, 1]) + group.all_engines[1].mark_stopped() + + srv = RolloutServer(server_groups=[group]) + try: + assert len(await srv.offload(tags=None)) == 1 + finally: + _kill_group(group)