Skip to content
Closed
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
105 changes: 105 additions & 0 deletions tests/experimental/orchestrator/weight_sync_driver_test.py
Original file line number Diff line number Diff line change
@@ -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()
121 changes: 121 additions & 0 deletions tests/experimental/rollout/manager_test.py
Original file line number Diff line number Diff line change
@@ -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()
1 change: 1 addition & 0 deletions tests/experimental/rollout/rollout_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
117 changes: 117 additions & 0 deletions tests/experimental/worker/rollout_worker_weight_sync_test.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading