mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Declare trainer concurrency groups once (#3077)
This commit is contained in:
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user