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
20 changes: 19 additions & 1 deletion tests/experimental/rollout/manager_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,13 @@ async def get_weight_sync_metadata(self, **kwargs):
self.calls.append(kwargs)
return self._metadata

async def bind_weight_sync(self, **kwargs):
self.calls.append("bind")
return None

async def pre_weight_sync(self, sync_request=None, **kwargs):
return "ok"


class GetWeightSyncMetadataTest(unittest.IsolatedAsyncioTestCase):

Expand Down Expand Up @@ -84,6 +91,18 @@ async def test_post_reopens_admission(self):
await manager.pre_weight_sync()
await manager.post_weight_sync()
self.assertTrue(manager._traffic.is_admission_open())

async def test_reopen_admission_after_abort(self):
manager = self._manager()
await manager.pre_weight_sync()
self.assertTrue(manager.reopen_admission())
self.assertTrue(manager._traffic.is_admission_open())

async def test_bind_delegates_to_sampler(self):
sampler = _FakeSyncSampler([])
manager = manager_lib.RolloutManager(
sampler=sampler, tokenizer="mock", chat_parser="mock")
await manager.bind_weight_sync()

async def test_repeated_pre_is_allowed(self):
manager = self._manager()
Expand Down Expand Up @@ -118,6 +137,5 @@ async def test_drain_timeout_returns(self):
task.cancel()
manager._active_tasks.pop("t0", None)


if __name__ == "__main__":
absltest.main()
19 changes: 18 additions & 1 deletion tests/experimental/worker/rollout_worker_weight_sync_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,10 @@ async def weight_sync(self, sync_request=None, **kwargs):
async def post_weight_sync(self, sync_request=None, **kwargs):
self.calls.append("post")
return "ok"

async def bind_weight_sync(self, **kwargs):
self.calls.append("bind")
return None

async def get_weight_sync_metadata(self, **kwargs):
self.calls.append("metadata")
Expand Down Expand Up @@ -117,7 +121,20 @@ async def test_full_round_call_order(self):
await worker.pre_weight_sync(req)
await worker.weight_sync(req)
await worker.post_weight_sync(req)
self.assertEqual(worker.manager.calls, ["metadata", "pre", "sync", "post"])
self.assertEqual(worker.manager.calls, ["bind", "metadata", "pre", "sync", "post"])

async def test_bind_delegates_to_manager(self):
worker = self._worker()
await worker.bind_weight_sync()
self.assertIn("bind", worker.manager.calls)

async def test_abort_reopens_admission(self):
worker = self._worker()
await worker.pre_weight_sync(_Request("r1", 1))
worker.manager.admission_open = False
await worker.abort_weight_sync(_Request("r1", 1))
self.assertTrue(worker.manager.admission_open)
self.assertEqual(worker.state, WorkerState.READY)


if __name__ == "__main__":
Expand Down
10 changes: 9 additions & 1 deletion tests/experimental/worker/trainer_worker_weight_sync_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,9 @@ class _ReleasingTrainer(_FakeTrainer):

def release_weight_sync(self, **kwargs):
self.calls.append("release")
self.release_kwargs = kwargs
return "released"


class _FailingTrainer(_FakeTrainer):

def prepare_weight_sync(self, sync_request=None, **kwargs):
Expand Down Expand Up @@ -83,6 +83,14 @@ def test_release_without_trainer_hook(self):
self.assertIsNone(worker.release_weight_sync())
self.assertEqual(worker.state, WorkerState.READY)

def test_release_forwards_sync_request(self):
trainer = _ReleasingTrainer()
worker = self._worker(trainer)
request = object()
worker.prepare_weight_sync()
worker.release_weight_sync(request)
self.assertIs(trainer.release_kwargs["sync_request"], request)


if __name__ == "__main__":
absltest.main()
4 changes: 4 additions & 0 deletions tunix/experimental/rollout/legacy_vllm_sampler_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,6 +278,10 @@ async def get_weight_sync_metadata(self, **kwargs) -> Any:
raise NotImplementedError(
"get_weight_sync_metadata() not implemented for this SamplerServer."
)

async def bind_weight_sync(self, **kwargs) -> Any:
del kwargs
return None

async def pre_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
"""Prepares staging handshake prior to policy weight update."""
Expand Down
8 changes: 8 additions & 0 deletions tunix/experimental/rollout/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,6 +295,14 @@ async def post_weight_sync(
self.resume_all()
self._traffic.reopen()
return res

def reopen_admission(self) -> bool:
"""Reopens rollout admission after an aborted round."""
return self._traffic.reopen()

async def bind_weight_sync(self, **kwargs) -> Any:
"""Binds the sampler's destination-side transport for this round."""
return await self.sampler.bind_weight_sync(**kwargs)

async def get_weight_sync_metadata(self, **kwargs) -> Any:
"""Returns the sampler's transport metadata for weight sync registration."""
Expand Down
4 changes: 4 additions & 0 deletions tunix/experimental/rollout/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,6 +174,10 @@ async def sample(
) -> list[SamplingResponse] | Any:
"""Generates completions for a batch of prompt conversations concurrently."""
...

async def bind_weight_sync(self, **kwargs) -> Any:
"""Binds destination-side transport resources. Idempotent per round."""
...

# --- Weight Synchronization ---
async def get_weight_sync_metadata(self, **kwargs) -> Any:
Expand Down
4 changes: 4 additions & 0 deletions tunix/experimental/rollout/vanilla_sampler_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -301,6 +301,10 @@ async def get_weight_sync_metadata(self, **kwargs) -> Any:
"get_weight_sync_metadata() not implemented for this SamplerServer."
)

async def bind_weight_sync(self, **kwargs) -> Any:
del kwargs
return None

async def pre_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
"""Prepares staging handshake prior to policy weight update."""
del sync_request, kwargs
Expand Down
5 changes: 3 additions & 2 deletions tunix/experimental/worker/rollout_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -478,8 +478,8 @@ async def post_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
return result

async def bind_weight_sync(self, **kwargs) -> Any:
"""No-op; the sampler binds its transport at engine init."""
return None
"""Binds the destination-side transport via the manager."""
return await self.manager.bind_weight_sync(**kwargs)

async def get_weight_sync_metadata(self, **kwargs) -> Any:
"""Returns the sampler's transport metadata via the manager."""
Expand All @@ -488,6 +488,7 @@ async def get_weight_sync_metadata(self, **kwargs) -> Any:
async def abort_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
"""Discards the round and resumes serving the previous weights."""
self.manager.resume_all()
self.manager.reopen_admission()
self.state = WorkerState.READY
self._record_round(sync_request, "aborted")
return None
Expand Down
2 changes: 1 addition & 1 deletion tunix/experimental/worker/trainer_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,7 +255,7 @@ def prepare_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
def release_weight_sync(self, sync_request: Any = None, **kwargs) -> Any:
"""Releases this round's staging and restores READY."""
release = getattr(self._trainer, "release_weight_sync", None)
result = release(**kwargs) if release else None
result = release(sync_request=sync_request, **kwargs) if release else None
if self.state == WorkerState.SYNCING:
self.state = WorkerState.READY
return result
Expand Down
Loading