Skip to content
Open
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
7 changes: 4 additions & 3 deletions miles/backends/fsdp_utils/update_weight_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,9 +71,10 @@ 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)
mode = self.args.pause_generation_mode
ray.get([engine.pause_generation.remote(mode=mode) for engine in self.rollout_engines])
if mode != "in_place":
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])
ray.get([engine.begin_weight_update.remote() for engine in self.rollout_engines])
dist.barrier(group=get_gloo_group())

Expand Down
62 changes: 52 additions & 10 deletions tests/fast/backends/test_fsdp_update_weight.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from argparse import Namespace
from types import SimpleNamespace

import pytest
import torch

from miles.backends.fsdp_utils import update_weight_utils
Expand All @@ -14,14 +15,15 @@ def __init__(self, fn, submissions, name):

def remote(self, *args, **kwargs):
self._submissions.append(self._name)
return _RemoteRef(self._fn, args, kwargs)
return _RemoteRef(self._fn, args, kwargs, self._name)


class _RemoteRef:
def __init__(self, fn, args, kwargs):
def __init__(self, fn, args, kwargs, name):
self._fn = fn
self._args = args
self._kwargs = kwargs
self.name = name

def resolve(self):
return self._fn(*self._args, **self._kwargs)
Expand All @@ -32,9 +34,10 @@ def __init__(self, name, events):
self.name = name
self.events = events
self.calls = []
self.pause_modes = []
self.submissions = []
self.session_open = False
self.pause_generation = self._remote(lambda: self._record("pause_generation"), "pause_generation")
self.pause_generation = self._remote(self._pause_generation, "pause_generation")
self.flush_cache = self._remote(lambda: self._record("flush_cache"), "flush_cache")
self.begin_weight_update = self._remote(self._begin_weight_update, "begin_weight_update")
self.update_weights_from_tensor = self._remote(
Expand All @@ -51,6 +54,10 @@ def _record(self, name):
self.calls.append(name)
self.events.append(f"{self.name}.{name}")

def _pause_generation(self, mode):
self.pause_modes.append(mode)
self._record("pause_generation")

def _begin_weight_update(self):
self._record("begin_weight_update")
assert not self.session_open
Expand Down Expand Up @@ -92,33 +99,54 @@ def _resolve_refs(value):
return value.resolve()


def _make_updater(model, rollout_engines):
def _make_updater(model, rollout_engines, pause_generation_mode="retract"):
updater = _SessionAwareUpdater(
Namespace(update_weight_buffer_size=1024),
Namespace(
pause_generation_mode=pause_generation_mode,
update_weight_buffer_size=1024,
),
SimpleNamespace(config=SimpleNamespace(model_type=""), state_dict=lambda: model),
)
updater.connect_rollout_engines(rollout_engines, None)
return updater


def test_fsdp_weight_updates_run_inside_engine_session(monkeypatch):
@pytest.mark.parametrize("pause_generation_mode", ["abort", "retract", "in_place"])
def test_fsdp_weight_updates_run_inside_engine_session(monkeypatch, pause_generation_mode):
events = []
ray_get_batches = []
engines = [_SessionEngine("engine0", events), _SessionEngine("engine1", events)]
updater = _make_updater({"weight": torch.ones(1)}, engines)
updater = _make_updater(
{"weight": torch.ones(1)},
engines,
pause_generation_mode=pause_generation_mode,
)

monkeypatch.setattr(update_weight_utils.ray, "get", _resolve_refs)
def resolve_and_record_refs(value):
refs = value if isinstance(value, list) else [value]
ray_get_batches.append([ref.name for ref in refs])
return _resolve_refs(value)

monkeypatch.setattr(update_weight_utils.ray, "get", resolve_and_record_refs)
monkeypatch.setattr(update_weight_utils.dist, "get_rank", lambda: 0)
monkeypatch.setattr(update_weight_utils.dist, "barrier", lambda **_kwargs: events.append("barrier"))
monkeypatch.setattr(update_weight_utils, "get_gloo_group", lambda: object())
monkeypatch.setattr(update_weight_utils, "gather_full_param", lambda param, async_op=False: param)

updater.update_weights()

expected_flush_events = (
[]
if pause_generation_mode == "in_place"
else [
"engine0.flush_cache",
"engine1.flush_cache",
]
)
assert events == [
"engine0.pause_generation",
"engine1.pause_generation",
"engine0.flush_cache",
"engine1.flush_cache",
*expected_flush_events,
"engine0.begin_weight_update",
"engine1.begin_weight_update",
"barrier",
Expand All @@ -130,6 +158,20 @@ def test_fsdp_weight_updates_run_inside_engine_session(monkeypatch):
"engine1.continue_generation",
"barrier",
]
assert engines[0].pause_modes == [pause_generation_mode]
assert engines[1].pause_modes == [pause_generation_mode]
expected_ray_get_batches = [["pause_generation", "pause_generation"]]
if pause_generation_mode != "in_place":
expected_ray_get_batches.append(["flush_cache", "flush_cache"])
expected_ray_get_batches.extend(
[
["begin_weight_update", "begin_weight_update"],
["update_weights_from_tensor"],
["end_weight_update", "end_weight_update"],
["continue_generation", "continue_generation"],
]
)
assert ray_get_batches == expected_ray_get_batches
assert engines[0].submissions == engines[0].calls
assert engines[1].submissions == engines[1].calls

Expand Down