diff --git a/tests/experimental/rollout/manager_test.py b/tests/experimental/rollout/manager_test.py index d919d4502..18bf3c249 100644 --- a/tests/experimental/rollout/manager_test.py +++ b/tests/experimental/rollout/manager_test.py @@ -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): @@ -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() @@ -118,6 +137,5 @@ async def test_drain_timeout_returns(self): task.cancel() manager._active_tasks.pop("t0", None) - if __name__ == "__main__": absltest.main() diff --git a/tests/experimental/worker/rollout_worker_weight_sync_test.py b/tests/experimental/worker/rollout_worker_weight_sync_test.py index 535c141c9..ec0d2ea4b 100644 --- a/tests/experimental/worker/rollout_worker_weight_sync_test.py +++ b/tests/experimental/worker/rollout_worker_weight_sync_test.py @@ -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") @@ -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__": diff --git a/tests/experimental/worker/trainer_worker_weight_sync_test.py b/tests/experimental/worker/trainer_worker_weight_sync_test.py index e1ece4acd..8ea7f2fae 100644 --- a/tests/experimental/worker/trainer_worker_weight_sync_test.py +++ b/tests/experimental/worker/trainer_worker_weight_sync_test.py @@ -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): @@ -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() diff --git a/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py b/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py index d7f499f11..789169456 100644 --- a/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py +++ b/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py @@ -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.""" diff --git a/tunix/experimental/rollout/manager.py b/tunix/experimental/rollout/manager.py index f18824ca3..1820b4378 100644 --- a/tunix/experimental/rollout/manager.py +++ b/tunix/experimental/rollout/manager.py @@ -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.""" diff --git a/tunix/experimental/rollout/sampler.py b/tunix/experimental/rollout/sampler.py index be59dc7aa..7a6b86661 100644 --- a/tunix/experimental/rollout/sampler.py +++ b/tunix/experimental/rollout/sampler.py @@ -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: diff --git a/tunix/experimental/rollout/vanilla_sampler_adapter.py b/tunix/experimental/rollout/vanilla_sampler_adapter.py index d3dc91224..338f075cc 100644 --- a/tunix/experimental/rollout/vanilla_sampler_adapter.py +++ b/tunix/experimental/rollout/vanilla_sampler_adapter.py @@ -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 diff --git a/tunix/experimental/worker/rollout_worker.py b/tunix/experimental/worker/rollout_worker.py index c32b636a2..a06a16b7f 100644 --- a/tunix/experimental/worker/rollout_worker.py +++ b/tunix/experimental/worker/rollout_worker.py @@ -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.""" @@ -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 diff --git a/tunix/experimental/worker/trainer_worker.py b/tunix/experimental/worker/trainer_worker.py index 1400841f9..1450b821c 100644 --- a/tunix/experimental/worker/trainer_worker.py +++ b/tunix/experimental/worker/trainer_worker.py @@ -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