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
19 changes: 7 additions & 12 deletions tests/experimental/orchestrator/orchestrator_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def test_create_engine(self):
)
self.assertIs(engine._inference_workers[datatypes.Role.REFERENCE], mock_ref)
self.assertSequenceEqual(
orch.worker_infos(), [actor_info, critic_info, ref_info, rollout_info]
orch.worker_infos(), [rollout_info, actor_info, critic_info, ref_info]
)

def test_create_engine_with_weight_sync_shim_registrations(self):
Expand All @@ -115,34 +115,29 @@ def test_create_engine_with_weight_sync_shim_registrations(self):
)
orch.register_worker_handle("actor-0", [datatypes.Role.ACTOR], mock_actor)

# Local handle fallback, no worker ID in _remote_worker_handles_by_id
# Local handle fallback
local_actor = remote_execution.InProcessActorHandle(
remote_execution.InProcessRemoteExecutionServer(mock.MagicMock())
)
orch._remote_worker_handles[datatypes.Role.ACTOR.value].append(local_actor)
orch.register_worker_handle(
"local-actor-0", [datatypes.Role.ACTOR], local_actor
)

engine = orch._create_engine()
self.assertIsNotNone(engine._weight_sync_coordinator)

# Assert they are properly shimmed in the registry
self.assertIn("actor-0", orch.registry.worker_ids())
self.assertIn("rollout-0", orch.registry.worker_ids())
self.assertIn("local-actor-0", orch.registry.worker_ids())
self.assertEqual(
type(orch.registry.get("actor-0")).__name__, "RemoteWorkerShim"
)
self.assertEqual(
type(orch.registry.get("rollout-0")).__name__, "RemoteWorkerShim"
)

local_actor_id = [
w_id
for w_id in orch.registry.worker_ids()
if w_id.startswith("local-actor-")
]
self.assertEqual(len(local_actor_id), 1)

self.assertEqual(
type(orch.registry.get(local_actor_id[0])).__name__, "RemoteWorkerShim"
type(orch.registry.get("local-actor-0")).__name__, "RemoteWorkerShim"
)

def test_bring_up_and_shutdown_remote_worker_handles(self):
Expand Down
232 changes: 218 additions & 14 deletions tests/experimental/orchestrator/worker_registry_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,16 @@

"""Tests for the WorkerRegistry and WorkerGroup."""

import pickle
import threading
import time
from unittest import mock

from absl.testing import absltest
from tunix.experimental.common import datatypes
from tunix.experimental.orchestrator import worker_registry
from tunix.experimental.worker import mock_worker
from tunix.experimental.worker import remote_execution


class WorkerRegistryTest(absltest.TestCase):
Expand All @@ -35,14 +42,24 @@ def test_register_and_group_by_role(self):
rollout_group = registry.group("rollout")
self.assertEqual(rollout_group.role, "rollout")
self.assertLen(rollout_group, 1)
self.assertEqual(list(rollout_group), [rollout])

r0_handle = registry.get("r0")
self.assertIsInstance(r0_handle, remote_execution.InProcessActorHandle)
self.assertEqual(list(rollout_group), [r0_handle])
self.assertEqual(rollout_group.handles(), [r0_handle])
self.assertEqual(rollout_group.members(), [r0_handle])
self.assertEqual(rollout_group.worker_ids(), ["r0"])
self.assertEqual(rollout_group.infos(), [registry.info("r0")])

trainer_group = registry.group("trainer")
self.assertEqual(trainer_group.role, "trainer")
self.assertLen(trainer_group, 2)
self.assertEqual(trainer_group.members(), [trainer0, trainer1])
t0_handle = registry.get("t0")
t1_handle = registry.get("t1")
self.assertEqual(trainer_group.members(), [t0_handle, t1_handle])
self.assertEqual(trainer_group.handles(), [t0_handle, t1_handle])
self.assertEqual(trainer_group.worker_ids(), ["t0", "t1"])

self.assertIs(registry.get("r0"), rollout)
self.assertLen(registry, 3)
self.assertIn("t0", registry)

Expand All @@ -58,26 +75,32 @@ def test_worker_group_properties(self):
self.assertEqual(rollout_group.role, "rollout")
self.assertFalse(rollout_group.is_empty())
self.assertLen(rollout_group, 1)
self.assertLen(list(rollout_group), 1)
self.assertEqual(rollout_group[0], registry.get("r0"))

self.assertEqual(trainer_group.role, "trainer")
self.assertFalse(trainer_group.is_empty())
self.assertLen(trainer_group, 2)
self.assertLen(list(trainer_group), 2)
self.assertEqual(trainer_group[0], registry.get("t0"))
self.assertEqual(trainer_group[1], registry.get("t1"))

empty_group = registry.group("inference")
self.assertEqual(empty_group.role, "inference")
self.assertTrue(empty_group.is_empty())
self.assertEmpty(empty_group)
self.assertEmpty(list(empty_group))
self.assertEmpty(empty_group.handles())
self.assertEmpty(empty_group.members())
self.assertEmpty(empty_group.infos())
self.assertEmpty(empty_group.worker_ids())

def test_fused_worker_joins_every_role(self):
registry = worker_registry.WorkerRegistry()
fused = mock_worker.MockWorker("f0", {"trainer", "inference"})
registry.register(fused)

self.assertEqual(registry.group("trainer").members(), [fused])
self.assertEqual(registry.group("inference").members(), [fused])
fused_handle = registry.get("f0")
self.assertEqual(registry.group("trainer").members(), [fused_handle])
self.assertEqual(registry.group("inference").members(), [fused_handle])

def test_duplicate_worker_id_raises(self):
registry = worker_registry.WorkerRegistry()
Expand Down Expand Up @@ -120,9 +143,10 @@ def test_unregister_retains_role_if_members_remain(self):
registry.register(mock_worker.MockWorker(worker_id="t0", roles={"trainer"}))
t1 = mock_worker.MockWorker(worker_id="t1", roles={"trainer"})
registry.register(t1)
t1_handle = registry.get("t1")
registry.unregister("t0")
self.assertIn("trainer", registry.roles())
self.assertEqual(registry.group("trainer").members(), [t1])
self.assertEqual(registry.group("trainer").members(), [t1_handle])

def test_registry_retrieval_methods(self):
registry = worker_registry.WorkerRegistry()
Expand All @@ -131,35 +155,215 @@ def test_registry_retrieval_methods(self):
registry.register(t0)
registry.register(r0)

r0_handle = registry.get("r0")
t0_handle = registry.get("t0")

self.assertEqual(registry.info("r0").worker_id, "r0")
self.assertEqual(registry.worker_ids(), ["r0", "t0"])
self.assertEqual(registry.workers(), [r0, t0])
self.assertEqual(registry.infos(), [r0.info(), t0.info()])
self.assertEqual(registry.get_handle("r0"), r0_handle)
self.assertEqual(registry.get_handle("t0"), t0_handle)
with self.assertRaises(KeyError):
registry.get_handle("non-existent")
self.assertEqual(registry.worker_ids(), ["t0", "r0"])
self.assertEqual(registry.handles(), [t0_handle, r0_handle])
self.assertEqual(registry.workers(), [t0_handle, r0_handle])
self.assertEqual(registry.infos(), [t0.info(), r0.info()])

def test_register_override_cleans_up_empty_roles(self):
registry = worker_registry.WorkerRegistry()
t0 = mock_worker.MockWorker(worker_id="t0", roles={"trainer"})
t1 = mock_worker.MockWorker(worker_id="t1", roles={"trainer"})
registry.register(t0)
registry.register(t1)
t1_handle = registry.get("t1")

# Override t0 with a new worker that no longer has the "trainer" role
t0_new = mock_worker.MockWorker(worker_id="t0", roles={"rollout"})
registry.register(t0_new, override=True)
t0_new_handle = registry.get("t0")

# The new t0_new should be removed from the "trainer" group
self.assertNotIn(t0_new, registry.group("trainer").members())
self.assertEqual(registry.group("trainer").members(), [t1])
self.assertNotIn(t0_new_handle, registry.group("trainer").members())
self.assertEqual(registry.group("trainer").members(), [t1_handle])
self.assertIn("trainer", registry.roles())

# Then override t1 to rollout as well to empty the role
t1_new = mock_worker.MockWorker(worker_id="t1", roles={"rollout"})
registry.register(t1_new, override=True)
t1_new_handle = registry.get("t1")
self.assertNotIn("trainer", registry.roles())

# The worker should just be "rollout" now
self.assertEqual(registry.roles(), {"rollout"})
self.assertCountEqual(registry.group("rollout").members(), [t0_new, t1_new])
self.assertCountEqual(
registry.group("rollout").members(), [t0_new_handle, t1_new_handle]
)

def test_register_handle_direct(self):
registry = worker_registry.WorkerRegistry()
mock_handle = mock.MagicMock(spec=remote_execution.ActorHandle)
info = registry.register_handle(
worker_id="actor-0",
roles=[datatypes.Role.ACTOR],
handle=mock_handle,
resources={"cores": 8},
)

self.assertEqual(info.worker_id, "actor-0")
self.assertEqual(info.roles, frozenset({"actor"}))
self.assertEqual(info.resources, {"remote": True, "cores": 8})
self.assertIs(registry.get("actor-0"), mock_handle)
self.assertEqual(registry.handles("actor"), [mock_handle])
self.assertEqual(registry.handles(datatypes.Role.ACTOR), [mock_handle])

# Rejects non-ActorHandle
with self.assertRaises(TypeError):
registry.register_handle(
worker_id="bad",
roles=["actor"],
handle="not_a_handle", # pytype: disable=wrong-arg-types
)

# Rejects empty roles
with self.assertRaises(ValueError):
registry.register_handle(
worker_id="no-roles",
roles=[],
handle=mock_handle,
)

# Rejects duplicate worker_id unless override
with self.assertRaises(ValueError):
registry.register_handle(
worker_id="actor-0",
roles=["actor"],
handle=mock_handle,
)

@mock.patch.object(remote_execution.ActorHandle, "from_address")
def test_register_from_hostname(self, mock_from_address):
mock_handle = mock.MagicMock(spec=remote_execution.ActorHandle)
mock_from_address.return_value = mock_handle

registry = worker_registry.WorkerRegistry()
meta = pickle.dumps({
"service_type": "trainer",
"service_port": 5001,
"worker_id": "trainer-0",
})
info = registry.register_from_hostname(
hostname="test-host",
port=5000,
metadata=meta,
rpc_timeout_s=60.0,
)

mock_from_address.assert_called_once_with(
"grpc://test-host:5001", rpc_timeout_s=60.0
)
self.assertEqual(info.worker_id, "trainer-0")
self.assertEqual(info.roles, frozenset({"actor"}))
self.assertEqual(
info.resources, {"remote": True, "address": "test-host:5001"}
)
self.assertIs(registry.get("trainer-0"), mock_handle)

def test_register_from_hostname_unknown_service_type(self):
registry = worker_registry.WorkerRegistry()
meta = pickle.dumps({
"service_type": "unknown_service",
"service_port": 5000,
"worker_id": "bad-0",
})
with self.assertRaisesRegex(
RuntimeError, "unknown service type unknown_service"
):
registry.register_from_hostname("host", 0, meta)

def test_role_normalization(self):
registry = worker_registry.WorkerRegistry()
handle_actor = mock.MagicMock(spec=remote_execution.ActorHandle)
handle_rollout = mock.MagicMock(spec=remote_execution.ActorHandle)

registry.register_handle(
worker_id="a0",
roles=[datatypes.Role.ACTOR],
handle=handle_actor,
)
registry.register_handle(
worker_id="r0",
roles=["rollout"],
handle=handle_rollout,
)

# Lookup by enum Role.ACTOR or string "actor" both work
self.assertEqual(
registry.group(datatypes.Role.ACTOR).handles(), [handle_actor]
)
self.assertEqual(registry.group("actor").handles(), [handle_actor])
self.assertEqual(registry.handles(datatypes.Role.ACTOR), [handle_actor])
self.assertEqual(registry.handles("actor"), [handle_actor])

# Lookup by enum Role.ROLLOUT or string "rollout" both work
self.assertEqual(
registry.group(datatypes.Role.ROLLOUT).handles(), [handle_rollout]
)
self.assertEqual(registry.group("rollout").handles(), [handle_rollout])
self.assertEqual(registry.handles(datatypes.Role.ROLLOUT), [handle_rollout])
self.assertEqual(registry.handles("rollout"), [handle_rollout])

def test_wait_for_workers_already_available(self):
registry = worker_registry.WorkerRegistry()
registry.register_handle(
"a0",
[datatypes.Role.ACTOR],
mock.MagicMock(spec=remote_execution.ActorHandle),
)
registry.register_handle(
"r0",
[datatypes.Role.ROLLOUT],
mock.MagicMock(spec=remote_execution.ActorHandle),
)

# Should return immediately without timing out
registry.wait_for_workers(
{
datatypes.Role.ACTOR: 1,
datatypes.Role.ROLLOUT: 1,
datatypes.Role.REFERENCE: 0,
},
timeout=1.0,
)

def test_wait_for_workers_delayed_registration(self):
registry = worker_registry.WorkerRegistry()
mock_actor = mock.MagicMock(spec=remote_execution.ActorHandle)

def register_later():
time.sleep(0.05)
registry.register_handle(
worker_id="actor-0",
roles=[datatypes.Role.ACTOR],
handle=mock_actor,
)

t = threading.Thread(target=register_later)
t.start()
registry.wait_for_workers(
min_workers={datatypes.Role.ACTOR: 1},
timeout=2.0,
poll_interval_s=0.01,
)
t.join()
self.assertLen(registry.handles(datatypes.Role.ACTOR), 1)

def test_wait_for_workers_timeout(self):
registry = worker_registry.WorkerRegistry()
with self.assertRaises(TimeoutError):
registry.wait_for_workers(
{datatypes.Role.ACTOR: 1},
timeout=0.05,
poll_interval_s=0.01,
)


if __name__ == "__main__":
Expand Down
Loading
Loading