Declare trainer concurrency groups once (#3077)

This commit is contained in:
fzyzcjy
2026-09-26 20:46:26 +08:00
committed by GitHub
parent fdb973993f
commit f580bef3bd
9 changed files with 82 additions and 194 deletions
-6
View File
@@ -36,11 +36,6 @@ TRAINER_CONTROLLER_ADDRS_FLAG = "--trainer-controller-addrs"
POOL_CATEGORY_TRAINER_ENGINE = "trainer_engine"
TRAINER_CONCURRENCY_GROUPS = {"heartbeat_status": 1, "default": 1, "fault_injector": 1, "kill_self": 1}
TRAINER_METHOD_CONCURRENCY_GROUPS = {
"get_heartbeat_status": "heartbeat_status",
"inject_fault": "fault_injector",
"kill_self": "kill_self",
}
TRAINER_CONTROLLER_WORKER_CLASS = "miles.ray.train.group.TrainerController"
@@ -241,7 +236,6 @@ def _compute_spec_trainer(
cell_index=ctx.cell_index,
),
concurrency_groups=TRAINER_CONCURRENCY_GROUPS if args.use_fault_tolerance else None,
method_concurrency_groups=TRAINER_METHOD_CONCURRENCY_GROUPS if args.use_fault_tolerance else None,
meta=lambda ctx: dict(role=config.role, cell_index=ctx.cell_index),
)
+11 -2
View File
@@ -488,9 +488,18 @@ class _ServeActorRayCommManager(_BaseActorManager[ServeWorkerSpec]):
)
def _compute_method_concurrency_groups(self, actor_class: type) -> dict[str, str]:
if self.spec.concurrency_groups is None:
if (groups := self.spec.concurrency_groups) is None:
return {}
return {**declared_concurrency_groups(actor_class), **(self.spec.method_concurrency_groups or {})}
method_groups = declared_concurrency_groups(actor_class)
assert method_groups, (
f"Worker {self.spec.name!r} declares concurrency groups {sorted(groups)} but no method of "
f"{actor_class.__name__} is annotated with @rpc(concurrency_group=...): threading the actor "
f"while every method stays in the default group buys nothing"
)
undeclared = sorted(set(method_groups.values()) - set(groups))
assert not undeclared, f"Worker {self.spec.name!r} routes methods to undeclared groups: {undeclared}"
return method_groups
async def post_setup(self) -> None:
pass
-16
View File
@@ -150,22 +150,6 @@ class ServeWorkerSpec(BaseWorkerSpec):
worker_class: str
ctor_kwargs: Callable[[WorkerCtorContext], dict[str, Any]]
concurrency_groups: dict[str, int] | None = None
method_concurrency_groups: dict[str, str] | None = None
@model_validator(mode="after")
def _require_the_groups_and_their_methods_together(self) -> "ServeWorkerSpec":
assert (self.concurrency_groups is None) == (self.method_concurrency_groups is None), (
f"Worker {self.name!r} must declare concurrency_groups and method_concurrency_groups "
f"together: groups nobody is assigned to are dead weight, and a method assigned to a "
f"group the actor never declares makes Ray reject the actor"
)
assert self.method_concurrency_groups is None or set(self.method_concurrency_groups.values()) <= set(
self.concurrency_groups
), (
f"Worker {self.name!r} routes methods to undeclared concurrency groups: "
f"{sorted(set(self.method_concurrency_groups.values()) - set(self.concurrency_groups))}"
)
return self
@model_validator(mode="before")
@classmethod
+3 -27
View File
@@ -17,7 +17,6 @@ from miles.ray.specs import train as train_specs
from miles.ray.specs.train import (
TRAINER_CONCURRENCY_GROUPS,
TRAINER_CONTROLLER_WORKER_CLASS,
TRAINER_METHOD_CONCURRENCY_GROUPS,
_compute_trainer_controller_provider,
compute_trainer_configs,
compute_trainer_controller_pool_id,
@@ -158,7 +157,6 @@ class TestScheduling:
def test_independent_dp_critic_cells_use_the_critic_gpu_shape(self, monkeypatch):
"""A critic sized differently from the actor must be split by its own GPU count."""
monkeypatch.setattr("miles.ray.specs.train.compute_megatron_world_size_except_dp", lambda _args: 2)
monkeypatch.setattr("miles.ray.specs.train._create_indep_dp_store_addr", lambda: "10.0.0.1:1234")
_actor_spec, critic_spec = specs_trainer(
_make_args(
@@ -300,36 +298,20 @@ class TestConcurrencyGroups:
assert spec.concurrency_groups == {"heartbeat_status": 1, "default": 1, "fault_injector": 1, "kill_self": 1}
def test_the_isolated_methods_travel_with_the_groups(self):
"""Declaring groups without routing any method to them leaves the heartbeat behind a train step."""
(spec,) = specs_trainer(_make_args(use_fault_tolerance=True))
assert spec.method_concurrency_groups == {
"get_heartbeat_status": "heartbeat_status",
"inject_fault": "fault_injector",
"kill_self": "kill_self",
}
def test_a_run_without_fault_tolerance_gets_a_plain_actor(self):
"""A threaded trainer actor runs NCCL setup off the main thread and deadlocked a non-FT run."""
(spec,) = specs_trainer(_make_args())
assert (spec.concurrency_groups, spec.method_concurrency_groups) == (None, None)
assert spec.concurrency_groups is None
def test_the_actor_is_not_annotated_statically(self):
"""A static @ray.method(concurrency_group=...) makes Ray reject the plain non-FT actor."""
annotations: list[str | None] = [
getattr(getattr(TrainRayActor, name), "__ray_concurrency_group__", None)
for name in TRAINER_METHOD_CONCURRENCY_GROUPS
for name in declared_concurrency_groups(TrainRayActor)
]
assert annotations == [None, None, None]
def test_every_routed_method_exists_on_the_trainer_actor(self):
"""A routed name the actor never defines only blows up when a fault-tolerant run launches."""
methods = [getattr(TrainRayActor, name, None) for name in TRAINER_METHOD_CONCURRENCY_GROUPS]
assert all(callable(method) for method in methods)
assert annotations and annotations == [None] * len(annotations)
def test_both_trainer_roles_follow_the_same_gate(self):
"""A critic threaded while its actor is not would deadlock exactly the run the gate protects."""
@@ -339,12 +321,6 @@ class TestConcurrencyGroups:
assert [spec.concurrency_groups is None for spec in fault_tolerant_specs] == [False, False]
assert [spec.concurrency_groups is None for spec in plain_specs] == [True, True]
def test_every_routed_group_is_declared(self):
"""Ray rejects an actor whose method names a concurrency group the class never declares."""
(spec,) = specs_trainer(_make_args(use_fault_tolerance=True))
assert set(spec.method_concurrency_groups.values()) <= set(TRAINER_CONCURRENCY_GROUPS)
def test_the_isolated_methods_are_annotated_on_the_actor(self):
"""Dropping an @rpc concurrency group would silently queue that call behind a train step."""
declared = declared_concurrency_groups(TrainRayActor)
+7 -1
View File
@@ -125,6 +125,7 @@ def _create_runner() -> _MiniFTControllerRunner:
api_server_url="http://127.0.0.1:8080",
poll_interval=10.0,
resume_delay=5.0,
cells_auto_resume=False,
)
@@ -704,10 +705,13 @@ class TestFtControllerDefaults:
class _FakeRunner:
def __init__(self, *, api_server_url: str, poll_interval: float, resume_delay: float) -> None:
def __init__(
self, *, api_server_url: str, poll_interval: float, resume_delay: float, cells_auto_resume: bool
) -> None:
self.api_server_url = api_server_url
self.poll_interval = poll_interval
self.resume_delay = resume_delay
self.cells_auto_resume = cells_auto_resume
self.ran = threading.Event()
self.thread_was_daemon: bool | None = None
self.thread_was_main_thread: bool | None = None
@@ -739,6 +743,7 @@ class TestMaybeStartMiniFtController:
api_server_port=18231,
mini_ft_controller_poll_interval=1.5,
mini_ft_controller_resume_delay=2.5,
cluster_backend="kubernetes",
)
)
@@ -747,6 +752,7 @@ class TestMaybeStartMiniFtController:
assert runner.api_server_url == "http://127.0.0.1:18231"
assert runner.poll_interval == 1.5
assert runner.resume_delay == 2.5
assert runner.cells_auto_resume is True
assert runner.ran.wait(timeout=5.0)
assert runner.thread_was_daemon is True
assert runner.thread_was_main_thread is False
+2 -1
View File
@@ -126,7 +126,7 @@ class TestGetModelUrl:
assert get_model_url(args, "unknown") == "http://10.0.0.1:3000/generate"
def test_get_model_url_no_routers(self):
"""get_model_url should work when sglang_model_routers is not set."""
"""get_model_url should work when sglang_model_routers is left at its parser default."""
from argparse import Namespace
from miles.rollout.sglang_rollout import get_model_url
@@ -134,6 +134,7 @@ class TestGetModelUrl:
args = Namespace(
sglang_router_ip="10.0.0.1",
sglang_router_port=3000,
sglang_model_routers=None,
)
assert get_model_url(args, "anything") == "http://10.0.0.1:3000/generate"
@@ -51,6 +51,10 @@ class _GroupedWorker:
@rpc(concurrency_group="heartbeat_status")
def outer_declaration_wins(self) -> None: ...
@rpc(concurrency_group="kill_self")
def isolated_echo(self, value: str, *, times: int = 1) -> str:
return value * times
def plain(self) -> None: ...
@@ -67,10 +71,7 @@ class DemoServeWorker:
_WORKER_CLASS_PATH = f"{DemoServeWorker.__module__}.{DemoServeWorker.__qualname__}"
_GROUPED_WORKER_CLASS_PATH = f"{_GroupedWorker.__module__}.{_GroupedWorker.__qualname__}"
_GROUPED_WORKER_ISOLATION = dict(
concurrency_groups={"kill_self": 1, "fault_injector": 1, "default": 1},
method_concurrency_groups={},
)
_GROUPED_WORKER_GROUPS = {"kill_self": 1, "fault_injector": 1, "default": 1}
_REPO_ROOT = Path(__file__).resolve().parents[4]
@@ -95,7 +96,6 @@ def _make_spec(
num_workers_per_cell: int = 1,
ctor_kwargs=None,
concurrency_groups: dict[str, int] | None = None,
method_concurrency_groups: dict[str, str] | None = None,
num_gpus_per_worker: float = 0,
num_cpus_per_worker: float = 0.2,
num_gpu_slots_per_worker: int = 0,
@@ -118,7 +118,6 @@ def _make_spec(
worker_class=worker_class,
ctor_kwargs=ctor_kwargs if ctor_kwargs is not None else (lambda _ctx: {}),
concurrency_groups=concurrency_groups,
method_concurrency_groups=method_concurrency_groups,
)
@@ -250,11 +249,9 @@ class TestServeWorkerClassFailures:
class TestServeSchedulingOptions:
async def test_concurrency_groups_reach_ray(self, fake_ray_cluster: FakeRayCluster):
"""The trainer heartbeat rpc must not queue behind a running train step."""
groups = {"heartbeat_status": 1, "default": 1}
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
await _launch([_make_spec(concurrency_groups=groups, method_concurrency_groups={"ping": "heartbeat_status"})])
assert _options(fake_ray_cluster)[0]["concurrency_groups"] == groups
assert _options(fake_ray_cluster)[0]["concurrency_groups"] == _GROUPED_WORKER_GROUPS
async def test_absent_concurrency_groups_are_not_passed_to_ray(self, fake_ray_cluster: FakeRayCluster):
"""Passing an empty group mapping would change how ray schedules the actor."""
@@ -264,18 +261,11 @@ class TestServeSchedulingOptions:
async def test_the_routed_methods_are_annotated_on_a_subclass(self, fake_ray_cluster: FakeRayCluster):
"""A declared group nobody is routed to leaves the isolated rpc queued behind the default group."""
await _launch(
[
_make_spec(
concurrency_groups={"probe": 1, "default": 1},
method_concurrency_groups={"ping": "probe"},
)
]
)
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
actor_class = _actor_classes(fake_ray_cluster)[0]
assert actor_class is not DemoServeWorker
assert actor_class.ping.__ray_concurrency_group__ == "probe"
assert actor_class is not bootstrapped_worker_class(_GROUPED_WORKER_CLASS_PATH)
assert actor_class.isolated.__ray_concurrency_group__ == "kill_self"
async def test_a_worker_without_groups_reaches_ray_unannotated(self, fake_ray_cluster: FakeRayCluster):
"""Ray refuses to build an actor whose method names a group the class never declares."""
@@ -295,22 +285,15 @@ class TestServeSchedulingOptions:
class TestServeConcurrencyGroupRouting:
async def test_the_declared_worker_class_stays_unannotated(self, fake_ray_cluster: FakeRayCluster):
"""Annotating the class itself would follow every later non-fault-tolerant run of that class."""
await _launch(
[
_make_spec(
concurrency_groups={"probe": 1, "default": 1},
method_concurrency_groups={"ping": "probe"},
)
]
)
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
assert not hasattr(DemoServeWorker.ping, "__ray_concurrency_group__")
assert not hasattr(_GroupedWorker.isolated, "__ray_concurrency_group__")
async def test_a_later_launch_without_groups_gets_a_class_no_earlier_launch_annotated(
self, fake_ray_cluster: FakeRayCluster
):
"""Ray rejects a plain actor whose method still names a group, so a fault-tolerant run must leave no mark."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
await _launch([_make_spec(name="plain-trainer", worker_class=_GROUPED_WORKER_CLASS_PATH)])
plain_actor_class = _actor_classes(fake_ray_cluster)[-1]
@@ -320,66 +303,66 @@ class TestServeConcurrencyGroupRouting:
async def test_each_routed_method_lands_in_its_own_group(self, fake_ray_cluster: FakeRayCluster):
"""Collapsing every routed method into one group serializes the heartbeat with the fault injector."""
await _launch(
[
_make_spec(
concurrency_groups={"probe": 1, "killer": 1, "default": 1},
method_concurrency_groups={"ping": "probe", "echo": "killer"},
)
]
)
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
actor_class = _actor_classes(fake_ray_cluster)[0]
assert (actor_class.ping.__ray_concurrency_group__, actor_class.echo.__ray_concurrency_group__) == (
"probe",
"killer",
)
async def test_an_unrouted_method_is_inherited_untouched(self, fake_ray_cluster: FakeRayCluster):
"""A train step pushed out of the default group would no longer block the group it must own."""
await _launch(
[
_make_spec(
concurrency_groups={"probe": 1, "default": 1},
method_concurrency_groups={"ping": "probe"},
)
]
)
actor_class = _actor_classes(fake_ray_cluster)[0]
assert actor_class.echo is DemoServeWorker.echo
assert (
actor_class.isolated.__ray_concurrency_group__,
actor_class.wrapped_isolated.__ray_concurrency_group__,
) == ("kill_self", "fault_injector")
async def test_a_routed_method_still_runs_the_original_body(self, fake_ray_cluster: FakeRayCluster):
"""A wrapper that swallowed the arguments or the return value would break every isolated rpc."""
await _launch(
[
_make_spec(
concurrency_groups={"probe": 1, "default": 1},
method_concurrency_groups={"echo": "probe"},
)
]
)
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
worker = _actor_classes(fake_ray_cluster)[0](ctor_kwargs=lambda _ctx: {}, context=_launch_context())
assert worker.echo("ab", times=2) == "abab"
assert worker.isolated_echo("ab", times=2) == "abab"
class TestConcurrencyGroupsAreValidatedAtLaunch:
async def test_groups_without_an_annotated_method_are_rejected(self, fake_ray_cluster: FakeRayCluster):
"""Threading the actor while every method stays in the default group buys nothing."""
spec = _make_spec(concurrency_groups={"heartbeat_status": 1, "default": 1})
manager = RayWorkerManager()
with pytest.raises(AssertionError, match="buys nothing"):
await manager.init(worker_manager_args(), [spec], {}, comm_backend=WorkerCommBackend.RAY)
assert fake_ray_cluster.handles == []
async def test_a_method_annotated_with_an_undeclared_group_is_rejected(self, fake_ray_cluster: FakeRayCluster):
"""Ray rejects the actor at creation time; the message must name the worker and the missing groups."""
spec = _make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups={"default": 1})
manager = RayWorkerManager()
with pytest.raises(AssertionError, match=r"'trainer'.*\['fault_injector', 'kill_self'\]"):
await manager.init(worker_manager_args(), [spec], {}, comm_backend=WorkerCommBackend.RAY)
async def test_a_declared_group_nobody_is_annotated_with_is_allowed(self, fake_ray_cluster: FakeRayCluster):
"""The trainer declares a default group precisely because no method is routed to it."""
groups = {**_GROUPED_WORKER_GROUPS, "spare": 1}
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=groups)])
assert _options(fake_ray_cluster)[0]["concurrency_groups"] == groups
class TestConcurrencyGroupsAreDeclaredOnce:
async def test_the_group_an_rpc_method_declares_reaches_ray(self, fake_ray_cluster: FakeRayCluster):
"""A method both wires isolate is declared once, and the launcher is what tells ray about it."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
assert _actor_classes(fake_ray_cluster)[0].isolated.__ray_concurrency_group__ == "kill_self"
async def test_a_group_declared_above_a_wrapper_still_reaches_ray(self, fake_ray_cluster: FakeRayCluster):
"""A group read off the wrapper alone would leave the wrapped method silently in the default group."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
assert _actor_classes(fake_ray_cluster)[0].wrapped_isolated.__ray_concurrency_group__ == "fault_injector"
async def test_a_default_group_method_is_left_undeclared(self, fake_ray_cluster: FakeRayCluster):
"""Ray rejects an actor naming a group its class never declares, and most methods name none."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
actor_class = _actor_classes(fake_ray_cluster)[0]
assert not hasattr(actor_class.plain, "__ray_concurrency_group__")
@@ -387,7 +370,7 @@ class TestConcurrencyGroupsAreDeclaredOnce:
async def test_both_wires_end_up_with_the_same_group(self, fake_ray_cluster: FakeRayCluster):
"""This is the whole point of declaring once: the two wires must not schedule a method differently."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
actor_class = _actor_classes(fake_ray_cluster)[0]
specs = collect_rpc_method_specs(_GroupedWorker)
@@ -399,20 +382,12 @@ class TestConcurrencyGroupsAreDeclaredOnce:
async def test_the_outermost_declaration_is_the_one_ray_hears(self, fake_ray_cluster: FakeRayCluster):
"""Two markers on one method must not resolve differently per wire, whichever one is meant to win."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, concurrency_groups=_GROUPED_WORKER_GROUPS)])
actor_class = _actor_classes(fake_ray_cluster)[0]
assert actor_class.outer_declaration_wins.__ray_concurrency_group__ == "kill_self"
assert collect_rpc_method_specs(_GroupedWorker)["outer_declaration_wins"].concurrency_group == "kill_self"
async def test_the_routed_body_of_a_declared_method_still_runs(self, fake_ray_cluster: FakeRayCluster):
"""The per-launch subclass wraps the method it annotates, so the wrapper must still call the original."""
await _launch([_make_spec(worker_class=_GROUPED_WORKER_CLASS_PATH, **_GROUPED_WORKER_ISOLATION)])
worker = _actor_classes(fake_ray_cluster)[0](ctor_kwargs=lambda _ctx: {}, context=_launch_context())
assert worker.isolated() is None
class TestServeWorkersAreStopped:
async def test_stopping_kills_the_actor_without_a_graceful_shutdown(self, fake_ray_cluster: FakeRayCluster):
@@ -428,69 +428,9 @@ class TestServeWorkerSpecExtraScheduling:
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
concurrency_groups={"heartbeat_status": 1, "default": 1},
method_concurrency_groups={"get_heartbeat_status": "heartbeat_status"},
)
assert spec.concurrency_groups == {"heartbeat_status": 1, "default": 1}
assert spec.method_concurrency_groups == {"get_heartbeat_status": "heartbeat_status"}
def test_groups_without_routed_methods_are_rejected(self):
"""Threading the actor while every method stays in the default group buys nothing."""
with pytest.raises(ValidationError, match="together"):
ServeWorkerSpec(
**_make_base_kwargs(),
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
concurrency_groups={"heartbeat_status": 1, "default": 1},
)
def test_routed_methods_without_groups_are_rejected(self):
"""Ray rejects an actor whose method names a concurrency group the class never declares."""
with pytest.raises(ValidationError, match="together"):
ServeWorkerSpec(
**_make_base_kwargs(),
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
method_concurrency_groups={"get_heartbeat_status": "heartbeat_status"},
)
def test_a_method_routed_to_an_undeclared_group_is_rejected(self):
"""Ray rejects the actor at creation time, long after the spec could have said why."""
with pytest.raises(ValidationError, match="undeclared concurrency groups"):
ServeWorkerSpec(
**_make_base_kwargs(),
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
concurrency_groups={"default": 1},
method_concurrency_groups={"get_heartbeat_status": "heartbeat_status"},
)
def test_a_declared_group_nobody_routes_to_is_allowed(self):
"""The trainer declares a default group precisely because no method is routed to it."""
spec = ServeWorkerSpec(
**_make_base_kwargs(),
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
concurrency_groups={"heartbeat_status": 1, "default": 1, "kill_self": 1},
method_concurrency_groups={"get_heartbeat_status": "heartbeat_status"},
)
assert set(spec.concurrency_groups) - set(spec.method_concurrency_groups.values()) == {"default", "kill_self"}
def test_the_rejection_names_the_worker_and_every_undeclared_group(self):
"""A message listing the declared groups instead of the missing ones sends the reader the wrong way."""
with pytest.raises(ValidationError, match=r"'demo-worker'.*\['fault_injector', 'kill_self'\]"):
ServeWorkerSpec(
**_make_base_kwargs(),
worker_class="miles.demo.Worker",
ctor_kwargs=lambda _ctx: {},
concurrency_groups={"heartbeat_status": 1, "default": 1},
method_concurrency_groups={
"get_heartbeat_status": "heartbeat_status",
"kill_self": "kill_self",
"inject_fault": "fault_injector",
},
)
def test_ctor_kwargs_receive_the_worker_position(self):
"""Each worker needs its own rank, so the callable is per worker."""
@@ -11,6 +11,7 @@ from miles.utils.workers.rpc.client.handle import RpcWorkerHandle
from miles.utils.workers.worker_provider.kubernetes.core import provider as core_provider
from miles.utils.workers.worker_provider.kubernetes.core.provider import KubernetesRunInfo, KubernetesWorkerProvider
from miles.utils.workers.worker_provider.kubernetes.helm.env import DEFAULT_LABEL_KEYS
from miles.utils.workers.worker_provider.utils import build_rpc_handle_of_worker_info
from miles.utils.workers.worker_spec import HostAndPort
NAMESPACE = "rl"
@@ -600,14 +601,16 @@ class TestGetWorkerInfos:
with pytest.raises(AssertionError, match="no observed worker pods"):
_worker_infos(_trainer_provider(FakePodApi()), cell_id="engine-00009")
def test_a_worker_that_is_not_served_has_no_handle(self):
"""A command worker has no RPC surface, so the provider must not give it a handle."""
def test_a_worker_that_is_not_served_names_no_class_and_refuses_a_handle(self):
"""A command worker has no rpc surface at all, so the refusal belongs where one is asked for."""
api = FakePodApi(pods=[make_pod(name="engine-0-0")])
provider = _provider(api, worker_ports={"engine": {"rpc": 8000}})
(info,) = _worker_infos(provider)
assert info.handle is None
assert info.worker_class is None
with pytest.raises(AssertionError, match="is not served"):
build_rpc_handle_of_worker_info(info)
def test_fans_a_pod_out_into_one_worker_per_rank_it_serves(self):
"""A supervised pod runs one worker process per rank, and each of them has to be driven separately."""