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
34 changes: 34 additions & 0 deletions rlix/scheduler/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -585,6 +585,32 @@ async def clear_progress(self, *, pipeline_id: str) -> None:
def _clear_rollout_intent_locked(self, pipeline_id: str) -> None:
self._state.rollout_open_pipelines.pop(pipeline_id, None)

def _warn_gen_request_without_demand_locked(
self, *, cluster_id: str, pipeline_id: str, step_target_estimate: Optional[int]
) -> None:
"""Warn when a GENERATION request carries no usable demand signal.

The gap-ratio planner skips pipelines whose registry entry has no
positive ``step_target_estimate`` AND no stored progress (see the
bootstrap-estimate branch in ``planner.plan_generation_gap_ratio``).
Such a request is never signaled and the caller blocks indefinitely.
Warn-only: progress may still arrive later on some backends, so
rejecting outright could break legitimate request-then-report flows.
"""
if step_target_estimate is not None and int(step_target_estimate) > 0:
return
_, step_target = self._pipeline_progress_totals_locked(pipeline_id=pipeline_id)
if step_target > 0.0:
return
logger.warning(
"request_gpus(GENERATION) for cluster_id=%r has no step_target_estimate "
"(got %r) and no stored progress; the gap-ratio planner will skip this "
"pipeline and the request may block indefinitely until progress or a "
"positive estimate arrives",
cluster_id,
step_target_estimate,
)

async def request_gpus(
self,
*,
Expand Down Expand Up @@ -636,6 +662,9 @@ async def request_gpus(
self._state.pending_bucket(priority).append(pending)
if priority == Priority.GENERATION:
self._state.rollout_open_pipelines[pipeline_id] = step_target_estimate
self._warn_gen_request_without_demand_locked(
cluster_id=cluster_id, pipeline_id=pipeline_id, step_target_estimate=step_target_estimate
)
# Queue Tracing: Track enqueue AFTER append (depth is correct)
self._tracer.trace_queue_enqueue(
cluster_id, priority, lora_name, bucket_depth=len(self._state.pending_bucket(priority))
Expand Down Expand Up @@ -737,6 +766,11 @@ async def notify_release_then_request_gpus(
if request_priority == Priority.GENERATION:
pipeline_id_gen, _ = parse_cluster_id(request_cluster_id)
self._state.rollout_open_pipelines[pipeline_id_gen] = request_step_target_estimate
self._warn_gen_request_without_demand_locked(
cluster_id=request_cluster_id,
pipeline_id=pipeline_id_gen,
step_target_estimate=request_step_target_estimate,
)
# Queue Tracing: Track enqueue AFTER append (depth is correct)
self._tracer.trace_queue_enqueue(
request_cluster_id,
Expand Down
113 changes: 113 additions & 0 deletions tests/test_scheduler_gen_request_warning.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,113 @@
"""PR ingress-validation test for the scheduler.

A GENERATION request with neither a positive ``step_target_estimate`` nor any
stored progress is skipped by the gap-ratio planner (bootstrap-estimate
branch) and blocks its caller indefinitely. The scheduler now warns at
enqueue time; these tests pin the warn / no-warn matrix.
"""

from __future__ import annotations

import importlib
import logging
import sys
import types
from pathlib import Path

import pytest

REPO_ROOT = Path(__file__).resolve().parents[1]
RLIX_ROOT = REPO_ROOT / "rlix"


def _install_import_stubs(monkeypatch: pytest.MonkeyPatch) -> None:
for module_name in list(sys.modules):
if module_name == "ray" or module_name.startswith("rlix"):
monkeypatch.delitem(sys.modules, module_name, raising=False)

ray_stub = types.ModuleType("ray")

def _remote(*args, **kwargs):
def _decorate(obj):
return obj

return _decorate

ray_stub.remote = _remote
ray_stub.get_actor = lambda *args, **kwargs: None
ray_stub.get = lambda value: value
monkeypatch.setitem(sys.modules, "ray", ray_stub)

package_roots = {
"rlix": RLIX_ROOT,
"rlix.protocol": RLIX_ROOT / "protocol",
"rlix.scheduler": RLIX_ROOT / "scheduler",
"rlix.utils": RLIX_ROOT / "utils",
}
for module_name, module_path in package_roots.items():
package_module = types.ModuleType(module_name)
package_module.__path__ = [str(module_path)] # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, module_name, package_module)


def _make_scheduler(monkeypatch: pytest.MonkeyPatch):
_install_import_stubs(monkeypatch)
scheduler_module = importlib.import_module("rlix.scheduler.scheduler")
protocol_types = importlib.import_module("rlix.protocol.types")
return scheduler_module, protocol_types, scheduler_module.SchedulerImpl()


_PIPELINE_ID = "miles_test"
_CLUSTER_ID = f"{_PIPELINE_ID}_actor_infer"


def test_warns_when_no_estimate_and_no_progress(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
_, _, scheduler = _make_scheduler(monkeypatch)
with caplog.at_level(logging.WARNING):
scheduler._warn_gen_request_without_demand_locked(
cluster_id=_CLUSTER_ID, pipeline_id=_PIPELINE_ID, step_target_estimate=None
)
assert "may block indefinitely" in caplog.text


def test_warns_on_non_positive_estimate(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
_, _, scheduler = _make_scheduler(monkeypatch)
with caplog.at_level(logging.WARNING):
scheduler._warn_gen_request_without_demand_locked(
cluster_id=_CLUSTER_ID, pipeline_id=_PIPELINE_ID, step_target_estimate=0
)
assert "may block indefinitely" in caplog.text


def test_silent_with_positive_estimate(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
_, _, scheduler = _make_scheduler(monkeypatch)
with caplog.at_level(logging.WARNING):
scheduler._warn_gen_request_without_demand_locked(
cluster_id=_CLUSTER_ID, pipeline_id=_PIPELINE_ID, step_target_estimate=8
)
assert caplog.text == ""


def test_silent_when_progress_already_stored(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
scheduler_module, protocol_types, scheduler = _make_scheduler(monkeypatch)
report = protocol_types.ProgressReport(
pipeline_id=_PIPELINE_ID,
step_target_trajectories=8,
metrics={"completed": 2},
)
scheduler._state.latest_progress_by_pipeline[_PIPELINE_ID] = {
"train": {"__full_finetune__": report}
}
with caplog.at_level(logging.WARNING):
scheduler._warn_gen_request_without_demand_locked(
cluster_id=_CLUSTER_ID, pipeline_id=_PIPELINE_ID, step_target_estimate=None
)
assert caplog.text == ""