diff --git a/tests/experimental/orchestrator/weight_sync_driver_test.py b/tests/experimental/orchestrator/weight_sync_driver_test.py new file mode 100644 index 000000000..c13b69c1c --- /dev/null +++ b/tests/experimental/orchestrator/weight_sync_driver_test.py @@ -0,0 +1,105 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for WeightSyncDriver.""" + +import asyncio + +from absl.testing import absltest +from tunix.experimental.orchestrator import weight_sync_driver + + +class _FakeCoordinator: + + def __init__(self): + self.calls = [] + self.loops = [] + self.fail_next = False + + async def sync(self, policy_version=0, **kwargs): + self.calls.append(policy_version) + self.loops.append(asyncio.get_running_loop()) + if self.fail_next: + self.fail_next = False + raise RuntimeError("round failed") + return f"committed-v{policy_version}" + + +class _FakeComponents: + + def __init__(self, coordinator): + self.coordinator = coordinator + self.closed = False + self.built_on = asyncio.get_running_loop() + + async def close(self): + self.closed = True + + +class WeightSyncDriverTest(absltest.TestCase): + + def _driver(self, coordinator): + self.factory_uuids = [] + + async def factory(initial_uuid): + self.factory_uuids.append(initial_uuid) + return _FakeComponents(coordinator) + + return weight_sync_driver.WeightSyncDriver( + factory, initial_uuid=7, initial_policy_version=3 + ) + + def test_factory_gets_initial_uuid(self): + driver = self._driver(_FakeCoordinator()) + self.assertEqual(self.factory_uuids, [7]) + driver.close() + + def test_sync_advances_version(self): + coordinator = _FakeCoordinator() + driver = self._driver(coordinator) + self.assertEqual(driver.sync_weights(), "committed-v4") + self.assertEqual(driver.sync_weights(), "committed-v5") + self.assertEqual(driver.policy_version, 5) + self.assertEqual(coordinator.calls, [4, 5]) + driver.close() + + def test_failed_round_keeps_version(self): + coordinator = _FakeCoordinator() + coordinator.fail_next = True + driver = self._driver(coordinator) + with self.assertRaises(RuntimeError): + driver.sync_weights() + self.assertEqual(driver.policy_version, 3) + self.assertEqual(driver.sync_weights(), "committed-v4") + driver.close() + + def test_everything_runs_on_one_loop(self): + coordinator = _FakeCoordinator() + driver = self._driver(coordinator) + driver.sync_weights() + driver.sync_weights() + loops = set(coordinator.loops) + self.assertLen(loops, 1) + driver.close() + + def test_close_closes_components(self): + coordinator = _FakeCoordinator() + driver = self._driver(coordinator) + components = driver._components + driver.close() + self.assertTrue(components.closed) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/rollout/manager_test.py b/tests/experimental/rollout/manager_test.py new file mode 100644 index 000000000..d1c4ecc20 --- /dev/null +++ b/tests/experimental/rollout/manager_test.py @@ -0,0 +1,121 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import asyncio +import unittest + +from absl.testing import absltest +from tunix.experimental.rollout import manager as manager_lib +from tunix.experimental.rollout import sampler as sampler_lib + + +class _FakeSampler(sampler_lib.Sampler): + + def __init__(self, metadata): + self._metadata = metadata + self.calls = [] + + async def get_weight_sync_metadata(self, **kwargs): + self.calls.append(kwargs) + return self._metadata + + +class GetWeightSyncMetadataTest(unittest.IsolatedAsyncioTestCase): + + async def test_delegates_to_sampler(self): + sampler = _FakeSampler([{"unit": "sampler0"}]) + manager = manager_lib.RolloutManager( + sampler=sampler, tokenizer="mock", chat_parser="mock" + ) + result = await manager.get_weight_sync_metadata() + self.assertEqual(result, [{"unit": "sampler0"}]) + + async def test_forwards_kwargs(self): + sampler = _FakeSampler([]) + manager = manager_lib.RolloutManager( + sampler=sampler, tokenizer="mock", chat_parser="mock" + ) + await manager.get_weight_sync_metadata(timeout_s=5) + self.assertEqual(sampler.calls, [{"timeout_s": 5}]) + + async def test_default_sampler_raises_not_implemented(self): + manager = manager_lib.RolloutManager(tokenizer="mock", chat_parser="mock") + with self.assertRaises(NotImplementedError): + await manager.get_weight_sync_metadata() + + +class _FakeSyncSampler(_FakeSampler): + + async def pre_weight_sync(self, sync_request=None, **kwargs): + return "ok" + + async def post_weight_sync(self, sync_request=None, **kwargs): + return "ok" + + +class AdmissionGateTest(unittest.IsolatedAsyncioTestCase): + + def _manager(self, **kwargs): + return manager_lib.RolloutManager( + sampler=_FakeSyncSampler([]), + tokenizer="mock", + chat_parser="mock", + **kwargs, + ) + + async def test_pre_closes_admission(self): + manager = self._manager() + await manager.pre_weight_sync() + self.assertFalse(manager._traffic.is_admission_open()) + + async def test_post_reopens_admission(self): + manager = self._manager() + await manager.pre_weight_sync() + await manager.post_weight_sync() + self.assertTrue(manager._traffic.is_admission_open()) + + async def test_repeated_pre_is_allowed(self): + manager = self._manager() + await manager.pre_weight_sync() + await manager.pre_weight_sync() + self.assertFalse(manager._traffic.is_admission_open()) + + async def test_pre_waits_for_inflight_work(self): + manager = self._manager() + done = asyncio.Event() + + async def work(): + await done.wait() + + task = asyncio.create_task(work()) + manager._active_tasks["t0"] = task + pre = asyncio.create_task(manager.pre_weight_sync()) + await asyncio.sleep(0.01) + self.assertFalse(pre.done()) + done.set() + await task + manager._active_tasks.pop("t0", None) + await pre + + async def test_drain_timeout_returns(self): + manager = self._manager(drain_timeout_s=0.05) + task = asyncio.create_task(asyncio.Event().wait()) + manager._active_tasks["t0"] = task + await manager.pre_weight_sync() + task.cancel() + manager._active_tasks.pop("t0", None) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/rollout/rollout_test.py b/tests/experimental/rollout/rollout_test.py index 725940101..cf14f54c4 100644 --- a/tests/experimental/rollout/rollout_test.py +++ b/tests/experimental/rollout/rollout_test.py @@ -241,6 +241,7 @@ async def _run_test(): await handle.asubmit("pre_weight_sync", metadata) v = await handle.asubmit("weight_sync", metadata) self.assertEqual(v, 333) + await handle.asubmit("post_weight_sync", metadata) # Direct coroutine execution of generate over handle req = datatypes.RolloutRequest( diff --git a/tests/experimental/worker/rollout_worker_weight_sync_test.py b/tests/experimental/worker/rollout_worker_weight_sync_test.py new file mode 100644 index 000000000..2375ccb0f --- /dev/null +++ b/tests/experimental/worker/rollout_worker_weight_sync_test.py @@ -0,0 +1,117 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for RolloutWorker weight sync phases.""" + +import unittest + +from absl.testing import absltest +from tunix.experimental.common import datatypes +from tunix.experimental.worker import rollout_worker as rollout_worker_lib + +WorkerState = datatypes.WorkerState + + +class _FakeManager: + + def __init__(self): + self.calls = [] + self.admission_open = True + + async def pre_weight_sync(self, sync_request=None, **kwargs): + self.calls.append("pre") + return "ok" + + async def weight_sync(self, sync_request=None, **kwargs): + self.calls.append("sync") + return 1 + + async def post_weight_sync(self, sync_request=None, **kwargs): + self.calls.append("post") + return "ok" + + async def get_weight_sync_metadata(self, **kwargs): + self.calls.append("metadata") + return [{"unit": "u0"}] + + def resume_all(self): + self.calls.append("resume") + + def reopen_admission(self): + self.admission_open = True + return True + + +class _Request: + + def __init__(self, req_id, uuid): + self.extra_config = {"req_id": req_id, "uuid": uuid} + + +class WeightSyncPhasesTest(unittest.IsolatedAsyncioTestCase): + + def _worker(self): + worker = rollout_worker_lib.RolloutWorker( + worker_id="w0", tokenizer="mock", chat_parser="mock" + ) + worker.manager = _FakeManager() + worker._state = WorkerState.READY + return worker + + async def test_pre_leaves_worker_syncing(self): + worker = self._worker() + await worker.pre_weight_sync(_Request("r1", 1)) + self.assertEqual(worker.state, WorkerState.SYNCING) + + async def test_post_restores_ready(self): + worker = self._worker() + await worker.pre_weight_sync(_Request("r1", 1)) + await worker.weight_sync(_Request("r1", 1)) + await worker.post_weight_sync(_Request("r1", 1)) + self.assertEqual(worker.state, WorkerState.READY) + + async def test_status_reports_round(self): + worker = self._worker() + await worker.pre_weight_sync(_Request("r1", 1)) + status = await worker.get_weight_sync_status() + self.assertEqual(status["req_id"], "r1") + self.assertEqual(status["uuid"], 1) + self.assertEqual(status["phase"], "prepared") + + async def test_abort_resumes_serving(self): + worker = self._worker() + await worker.pre_weight_sync(_Request("r1", 1)) + await worker.abort_weight_sync(_Request("r1", 1)) + self.assertEqual(worker.state, WorkerState.READY) + status = await worker.get_weight_sync_status() + self.assertEqual(status["phase"], "aborted") + + async def test_metadata_delegates_to_manager(self): + worker = self._worker() + result = await worker.get_weight_sync_metadata() + self.assertEqual(result, [{"unit": "u0"}]) + + async def test_full_round_call_order(self): + worker = self._worker() + req = _Request("r1", 1) + await worker.bind_weight_sync() + await worker.get_weight_sync_metadata() + 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"]) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/worker/traffic_controller_test.py b/tests/experimental/worker/traffic_controller_test.py new file mode 100644 index 000000000..48b7fbf75 --- /dev/null +++ b/tests/experimental/worker/traffic_controller_test.py @@ -0,0 +1,96 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for traffic_controller.""" + +import asyncio +import unittest + +from absl.testing import absltest +from tunix.experimental.common import datatypes +from tunix.experimental.worker import traffic_controller + +WorkerState = datatypes.WorkerState + + +class TrafficControllerTest(unittest.IsolatedAsyncioTestCase): + + def setUp(self): + super().setUp() + self.controller = traffic_controller.TrafficController() + + def test_initial_state(self): + self.assertEqual(self.controller.state, WorkerState.READY) + self.assertTrue(self.controller.is_admission_open()) + self.assertEqual(len(self.controller.get_active_tasks()), 0) + + async def test_try_admit_success(self): + async def dummy_task(): + await asyncio.sleep(0.01) + + task = asyncio.create_task(dummy_task()) + self.assertTrue(self.controller.try_admit(task)) + self.assertIn(task, self.controller.get_active_tasks()) + + await task + + # Task should be removed automatically + self.assertEqual(len(self.controller.get_active_tasks()), 0) + + async def test_try_admit_rejected_when_syncing(self): + self.controller.transition_to_syncing() + self.assertEqual(self.controller.state, WorkerState.SYNCING) + self.assertFalse(self.controller.is_admission_open()) + + async def dummy_task(): + pass + + task = asyncio.create_task(dummy_task()) + self.assertFalse(self.controller.try_admit(task)) + await asyncio.sleep(0) + self.assertTrue(task.cancelled()) + self.assertEqual(len(self.controller.get_active_tasks()), 0) + + async def test_reopen(self): + self.controller.transition_to_syncing() + self.assertTrue(self.controller.reopen()) + self.assertEqual(self.controller.state, WorkerState.READY) + self.assertTrue(self.controller.is_admission_open()) + + async def test_stop_and_cancel_all(self): + async def dummy_task(): + await asyncio.sleep(1.0) + + task1 = asyncio.create_task(dummy_task()) + task2 = asyncio.create_task(dummy_task()) + + self.controller.try_admit(task1) + self.controller.try_admit(task2) + + cancelled_tasks = self.controller.stop_and_cancel_all() + self.assertCountEqual(cancelled_tasks, [task1, task2]) + await asyncio.sleep(0) + self.assertTrue(task1.cancelled()) + self.assertTrue(task2.cancelled()) + + self.assertEqual(self.controller.state, WorkerState.STOPPED) + self.assertFalse(self.controller.is_admission_open()) + + # Cannot reopen after stop + self.assertFalse(self.controller.reopen()) + self.assertEqual(self.controller.state, WorkerState.STOPPED) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/worker/trainer_worker_weight_sync_test.py b/tests/experimental/worker/trainer_worker_weight_sync_test.py new file mode 100644 index 000000000..e1ece4acd --- /dev/null +++ b/tests/experimental/worker/trainer_worker_weight_sync_test.py @@ -0,0 +1,88 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for TrainerWorker weight sync staging.""" + +from absl.testing import absltest +from tunix.experimental.common import datatypes +from tunix.experimental.worker import trainer_worker as trainer_worker_lib + +WorkerState = datatypes.WorkerState + + +class _FakeTrainer: + + def __init__(self): + self.calls = [] + + def prepare_weight_sync(self, sync_request=None, **kwargs): + self.calls.append("prepare") + return [{"unit": "trainer0"}] + + +class _ReleasingTrainer(_FakeTrainer): + + def release_weight_sync(self, **kwargs): + self.calls.append("release") + return "released" + + +class _FailingTrainer(_FakeTrainer): + + def prepare_weight_sync(self, sync_request=None, **kwargs): + raise RuntimeError("boom") + + +class WeightSyncStagingTest(absltest.TestCase): + + def _worker(self, trainer): + worker = trainer_worker_lib.TrainerWorker( + trainer_factory=lambda: trainer, worker_id="t0" + ) + worker.initialize() + worker._state = WorkerState.READY + return worker + + def test_prepare_stays_syncing(self): + worker = self._worker(_FakeTrainer()) + worker.prepare_weight_sync() + self.assertEqual(worker.state, WorkerState.SYNCING) + + def test_prepare_returns_trainer_metadata(self): + worker = self._worker(_FakeTrainer()) + self.assertEqual(worker.prepare_weight_sync(), [{"unit": "trainer0"}]) + + def test_prepare_failure_sets_error_state(self): + worker = self._worker(_FailingTrainer()) + with self.assertRaises(RuntimeError): + worker.prepare_weight_sync() + self.assertEqual(worker.state, WorkerState.ERROR) + + def test_release_restores_ready(self): + trainer = _ReleasingTrainer() + worker = self._worker(trainer) + worker.prepare_weight_sync() + self.assertEqual(worker.release_weight_sync(), "released") + self.assertEqual(worker.state, WorkerState.READY) + self.assertEqual(trainer.calls, ["prepare", "release"]) + + def test_release_without_trainer_hook(self): + worker = self._worker(_FakeTrainer()) + worker.prepare_weight_sync() + self.assertIsNone(worker.release_weight_sync()) + self.assertEqual(worker.state, WorkerState.READY) + + +if __name__ == "__main__": + absltest.main() diff --git a/tunix/experimental/orchestrator/weight_sync_driver.py b/tunix/experimental/orchestrator/weight_sync_driver.py new file mode 100644 index 000000000..55f467a90 --- /dev/null +++ b/tunix/experimental/orchestrator/weight_sync_driver.py @@ -0,0 +1,71 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Runs weight sync rounds from synchronous code on a resident event loop.""" + +import asyncio +import threading +from typing import Any, Callable, Coroutine + + +class WeightSyncDriver: + """Bridges a synchronous training loop to the async weight sync coordinator. + + gRPC channels are bound to the loop that created them, so the driver owns + one loop for its whole lifetime and builds every component on it. + """ + + def __init__( + self, + components_factory: Callable[[int], Coroutine[Any, Any, Any]], + *, + initial_uuid: int, + initial_policy_version: int, + round_timeout_s: float = 3600.0, + ): + self._round_timeout_s = round_timeout_s + self._policy_version = initial_policy_version + self._loop = asyncio.new_event_loop() + self._thread = threading.Thread( + target=self._loop.run_forever, name="weight-sync-driver", daemon=True + ) + self._thread.start() + self._components = self._run(components_factory(initial_uuid)) + + @property + def policy_version(self) -> int: + return self._policy_version + + def _run(self, coro: Coroutine[Any, Any, Any]) -> Any: + """Runs a coroutine on the driver loop and blocks for its result.""" + future = asyncio.run_coroutine_threadsafe(coro, self._loop) + return future.result(self._round_timeout_s) + + def sync_weights(self, **kwargs) -> Any: + """Runs one round and advances the policy version on commit.""" + version = self._policy_version + 1 + result = self._run( + self._components.coordinator.sync(policy_version=version, **kwargs) + ) + self._policy_version = version + return result + + def close(self) -> None: + """Closes the components on the driver loop, then retires the loop.""" + close = getattr(self._components, "close", None) + if close is not None: + self._run(close()) + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join() + self._loop.close() diff --git a/tunix/experimental/rollout/manager.py b/tunix/experimental/rollout/manager.py index e2478e193..a289b22b4 100644 --- a/tunix/experimental/rollout/manager.py +++ b/tunix/experimental/rollout/manager.py @@ -22,6 +22,7 @@ from tunix.experimental.rollout import sampler as sampler_lib from tunix.experimental.rollout import vanilla_sampler_adapter from tunix.experimental.trajectory import trajectory as trajectory_lib +from tunix.experimental.worker import traffic_controller as traffic_controller_lib from tunix.rl.rollout import base_rollout TrajectoryOrError = Union[ @@ -45,6 +46,7 @@ def __init__( max_concurrency: int = 64, tokenizer: Any = None, chat_parser: Any = None, + drain_timeout_s: float = 300.0, ): self.config = config if sampler is None: @@ -89,6 +91,8 @@ def __init__( ] = {} self._active_tasks: Dict[str, asyncio.Task[Any]] = {} self._completed_queue: asyncio.Queue[TrajectoryOrError] = asyncio.Queue() + self._traffic = traffic_controller_lib.TrafficController() + self._drain_timeout_s = drain_timeout_s async def _generate_one( self, @@ -96,6 +100,8 @@ async def _generate_one( on_complete: Optional[Callable[[TrajectoryOrError], None]] = None, ) -> TrajectoryOrError: """Spawns an async task running the multi-turn episode loop concurrently.""" + if not self._traffic.is_admission_open(): + raise RuntimeError("rollout admission is closed during weight sync") loop = asyncio.get_running_loop() future: asyncio.Future[TrajectoryOrError] = loop.create_future() @@ -245,10 +251,18 @@ def cancel_all(self) -> None: for task in self._active_tasks.values(): task.cancel() + def reopen_admission(self) -> bool: + """Reopens request admission after a weight sync round.""" + return self._traffic.reopen() + async def pre_weight_sync( self, sync_request: sampler_lib.WeightSyncRequest | Any = None, **kwargs ) -> Any: - """Phase 3 Barrier 1: Pauses active collectors and checks staging.""" + """Phase 3 Barrier 1: Closes admission and drains in-flight work.""" + self._traffic.transition_to_syncing() + tasks = list(self._active_tasks.values()) + if tasks: + await asyncio.wait(tasks, timeout=self._drain_timeout_s) self.pause_all() if self.sampler: return await self.sampler.pre_weight_sync(sync_request, **kwargs) @@ -273,4 +287,11 @@ async def post_weight_sync( if self.sampler: res = await self.sampler.post_weight_sync(sync_request, **kwargs) self.resume_all() + self._traffic.reopen() return res + + async def get_weight_sync_metadata(self, **kwargs) -> Any: + """Returns the sampler's transport metadata for weight sync registration.""" + if self.sampler: + return await self.sampler.get_weight_sync_metadata(**kwargs) + return [] diff --git a/tunix/experimental/worker/rollout_worker.py b/tunix/experimental/worker/rollout_worker.py index d8e4d401b..c9e7e751b 100644 --- a/tunix/experimental/worker/rollout_worker.py +++ b/tunix/experimental/worker/rollout_worker.py @@ -74,6 +74,8 @@ def __init__( self.worker_id = worker_id self.config = config self._policy_version = 0 + self._state = datatypes.WorkerState.PENDING + self._sync_round = {"req_id": None, "uuid": 0, "phase": "idle"} if tokenizer is None or chat_parser is None: raise ValueError( "RolloutWorker requires valid tokenizer and chat_parser arguments" @@ -447,29 +449,57 @@ async def as_completed_stream( yield self._to_rollout_response(res) async def pre_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: - """Prepares the worker for an upcoming weight synchronization step.""" + """Quiesces the worker; it stays SYNCING until post or abort.""" if self.state == WorkerState.PENDING: self.initialize() self.state = WorkerState.SYNCING - try: - return await self.manager.pre_weight_sync(sync_request, **kwargs) - finally: - self.state = WorkerState.READY + self._record_round(sync_request, "idle") + result = await self.manager.pre_weight_sync(sync_request, **kwargs) + self._record_round(sync_request, "prepared") + return result async def weight_sync(self, sync_request: Any = None, **kwargs) -> Any: - """Synchronizes the worker's internal model weights.""" + """Materializes the received weights; the worker stays SYNCING.""" if self.state == WorkerState.PENDING: self.initialize() self.state = WorkerState.SYNCING - try: - metadata = kwargs.pop("metadata", None) - request = sync_request if sync_request is not None else metadata - result = await self.manager.weight_sync(request, **kwargs) - self._policy_version += 1 - return result - finally: - self.state = WorkerState.READY + metadata = kwargs.pop("metadata", None) + request = sync_request if sync_request is not None else metadata + result = await self.manager.weight_sync(request, **kwargs) + self._policy_version += 1 + self._record_round(sync_request, "h2d_done") + return result async def post_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: - """Finalizes policy weight update and resumes workers.""" - return await self.manager.post_weight_sync(sync_request, **kwargs) + """Publishes the new weights and resumes serving.""" + result = await self.manager.post_weight_sync(sync_request, **kwargs) + self.state = WorkerState.READY + self._record_round(sync_request, "committed") + return result + + async def bind_weight_sync(self, **kwargs) -> Any: + """No-op; the sampler binds its transport at engine init.""" + return None + + async def get_weight_sync_metadata(self, **kwargs) -> Any: + """Returns the sampler's transport metadata via the manager.""" + return await self.manager.get_weight_sync_metadata(**kwargs) + + 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 + + async def get_weight_sync_status(self, **kwargs) -> Any: + """Returns this worker's view of the current weight sync round.""" + return dict(self._sync_round, policy_version=self._policy_version) + + def _record_round(self, sync_request: Any, phase: str) -> None: + extra = getattr(sync_request, "extra_config", None) or {} + if extra.get("req_id") is not None: + self._sync_round["req_id"] = extra.get("req_id") + self._sync_round["uuid"] = extra.get("uuid", 0) + self._sync_round["phase"] = phase diff --git a/tunix/experimental/worker/traffic_controller.py b/tunix/experimental/worker/traffic_controller.py new file mode 100644 index 000000000..b3764c612 --- /dev/null +++ b/tunix/experimental/worker/traffic_controller.py @@ -0,0 +1,125 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Traffic controller for rollout worker admission and state management.""" + +import asyncio +import threading +from typing import Any + +from tunix.experimental.common import datatypes + +WorkerState = datatypes.WorkerState + + +class TrafficController: + """Encapsulates worker lifecycle state and task admission control. + + Provides a single source of truth for the rollout worker's state and active + async tasks, avoiding distributed locks across the worker and manager. + """ + + def __init__(self): + self._lock = threading.RLock() + self._state = WorkerState.READY + self._admission_open = asyncio.Event() + self._admission_open.set() + self._active_tasks: set[asyncio.Task[Any]] = set() + + @property + def state(self) -> WorkerState: + """Returns the current worker state.""" + with self._lock: + return self._state + + async def wait_for_admission(self) -> None: + """Blocks until admission is open.""" + await self._admission_open.wait() + + def is_admission_open(self) -> bool: + """Returns True if admission is currently open.""" + with self._lock: + return self._admission_open.is_set() + + def try_admit(self, task: asyncio.Task[Any]) -> bool: + """Attempts to admit a task if the gate is open. + + If admission is open and the state is READY, the task is tracked. + If not, the task is immediately cancelled. + + Args: + task: An asyncio.Task to admit. + + Returns: + True if admitted, False if rejected. + """ + with self._lock: + if not self._admission_open.is_set() or self._state != WorkerState.READY: + task.cancel() + return False + self._active_tasks.add(task) + task.add_done_callback(self._on_task_done) + return True + + def _on_task_done(self, task: asyncio.Task[Any]) -> None: + with self._lock: + self._active_tasks.discard(task) + + def get_active_tasks(self) -> list[asyncio.Task[Any]]: + """Returns a snapshot of currently running tasks.""" + with self._lock: + return list(self._active_tasks) + + def transition_to_syncing(self) -> None: + """Transitions state to SYNCING and closes admission. + + Raises: + RuntimeError: If transitioning from an invalid state. + """ + with self._lock: + if self._state not in (WorkerState.READY, WorkerState.SYNCING): + raise RuntimeError( + f"Cannot transition to SYNCING from {self._state.value}" + ) + self._state = WorkerState.SYNCING + self._admission_open.clear() + + def reopen(self) -> bool: + """Reopens admission and sets state to READY. + + Returns: + True if successfully reopened, False if the worker was STOPPED. + """ + with self._lock: + if self._state == WorkerState.STOPPED: + return False + self._state = WorkerState.READY + self._admission_open.set() + return True + + def stop_and_cancel_all(self) -> list[asyncio.Task[Any]]: + """Permanently sets state to STOPPED, closes admission, and cancels tasks. + + Returns: + A list of the cancelled tasks. + """ + with self._lock: + self._state = WorkerState.STOPPED + self._admission_open.clear() + tasks = list(self._active_tasks) + + # Cancel outside the lock to avoid reentrancy if callbacks are invoked immediately + for task in tasks: + task.cancel() + return tasks diff --git a/tunix/experimental/worker/trainer_worker.py b/tunix/experimental/worker/trainer_worker.py index c53f97513..1400841f9 100644 --- a/tunix/experimental/worker/trainer_worker.py +++ b/tunix/experimental/worker/trainer_worker.py @@ -235,13 +235,14 @@ def restore_checkpoint(self, **kwargs) -> Any: """Restore state from latest checkpoint and return the metadata pytree.""" return self._trainer.restore_checkpoint(**kwargs) - def prepare_weight_sync(self, **kwargs) -> Any: - """Stages weights for transfer and returns coordinates/metadata.""" + def prepare_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: + """Stages weights for transfer and returns their metadata.""" self._ensure_ready() self.state = WorkerState.SYNCING try: + if sync_request is not None: + kwargs["sync_request"] = sync_request metadata = self._trainer.prepare_weight_sync(**kwargs) - self.state = WorkerState.READY self._last_error = None if metadata is not None: return metadata @@ -251,6 +252,14 @@ def prepare_weight_sync(self, **kwargs) -> Any: self.state = WorkerState.ERROR raise + 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 + if self.state == WorkerState.SYNCING: + self.state = WorkerState.READY + return result + def get_metrics(self) -> Any: """Returns and clears the recently collected step metric records.""" return self._trainer.get_metrics()