Resolve the eval fleet's inference controller handle once (#3024)

This commit is contained in:
fzyzcjy
2026-09-26 20:20:14 +08:00
committed by GitHub
parent 6c5e42dc8a
commit 04e603ea54
2 changed files with 21 additions and 13 deletions
+5 -4
View File
@@ -18,7 +18,7 @@ from miles.rollout.checkpoint_eval import EvalSkip, retarget_args
from miles.rollout.inference_rollout.inference_rollout_common import GenerateState
from miles.utils.http_utils import wait_http_ok
from miles.utils.workers.rpc.client.misc import ServerRestartedError
from miles.utils.workers.worker_handle import WorkerUnreachableError
from miles.utils.workers.worker_handle import BaseWorkerHandle, WorkerUnreachableError
from miles.utils.workers.worker_provider.base import BaseWorkerProvider
from miles.utils.workers.worker_spec import HostAndPort
@@ -48,15 +48,16 @@ class RolloutExecutorEvalFleet:
def __init__(
self, args: Namespace, *, info: EvalFleetInfo, inference_controller_provider: BaseWorkerProvider
) -> None:
self._inference_controller_provider = inference_controller_provider
self._inference_controller: BaseWorkerHandle = inference_controller_provider.get_handle(
inference_controller_worker_name()
)
self._state = GenerateState(
retarget_args(args, info.router.host, info.router.port, info.num_gpus, info.num_gpus_per_engine)
)
async def pin(self, checkpoint_dir: str, weight_version: str) -> GenerateState:
try:
inference_controller = self._inference_controller_provider.get_handle(inference_controller_worker_name())
pin = await inference_controller.pin_eval_fleet(
pin = await self._inference_controller.pin_eval_fleet(
checkpoint_dir=checkpoint_dir, weight_version=weight_version
)
except UNREACHABLE_CONTROLLER_ERRORS as e:
+16 -9
View File
@@ -360,10 +360,10 @@ class TestRolloutExecutorEvalFleet:
assert exc.value.reason == "pin_violation"
async def test_the_controller_is_resolved_again_for_every_point(self, fleet_states):
"""A controller that restarted answers on a new handle, and a session that kept the old one never heals."""
async def test_the_controller_is_resolved_once_and_the_handle_reused(self, fleet_states):
"""Resolving a fresh handle per point would hand every point a fresh boot-uuid pin, so it is resolved once."""
first, second = (
FakeInferenceController([EvalFleetPin(skip_reason=None)]),
FakeInferenceController([EvalFleetPin(skip_reason=None), EvalFleetPin(skip_reason=None)]),
FakeInferenceController([EvalFleetPin(skip_reason=None)]),
)
provider = FakeControllerProvider([first, second])
@@ -372,7 +372,7 @@ class TestRolloutExecutorEvalFleet:
await session.pin("/snap/step_5", "5")
await session.pin("/snap/step_6", "6")
assert (provider.lookups, len(first.calls), len(second.calls)) == (2, 1, 1)
assert (provider.lookups, len(first.calls), len(second.calls)) == (1, 2, 0)
async def test_a_controller_that_cannot_be_reached_skips_the_point(self, fleet_states):
"""Losing the controller must skip one eval point, not raise into the driver's rollout loop."""
@@ -399,9 +399,16 @@ class TestRolloutExecutorEvalFleet:
assert exc.value.reason == "controller_unreachable"
async def test_a_controller_that_restarted_skips_the_point(self, fleet_states):
"""A restarted server is a transport failure too, and eval must degrade rather than crash the run."""
session = make_session_over(FakeControllerProvider([ServerRestartedError("boot uuid changed")]))
async def test_the_pin_keeps_reporting_a_restarted_controller_across_points(self, fleet_states):
"""The kept handle keeps its boot-uuid pin, so a later point still sees the restart instead of a fresh pin."""
controller = FakeInferenceController(
[ServerRestartedError("boot uuid changed"), ServerRestartedError("boot uuid changed")]
)
provider = FakeControllerProvider([controller])
session = make_session_over(provider)
with pytest.raises(EvalSkip):
await session.pin("/snap/step_5", "5")
for checkpoint_dir, weight_version in (("/snap/step_5", "5"), ("/snap/step_6", "6")):
with pytest.raises(EvalSkip):
await session.pin(checkpoint_dir, weight_version)
assert (provider.lookups, len(controller.calls)) == (1, 2)