Skip to content
Draft
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
33 changes: 29 additions & 4 deletions docs/content/docs/checkpointing-logging/vllm-metrics.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -97,10 +97,35 @@ Loading a checkpoint disables run aggregates while preserving step metrics.
`LATEST` with no checkpoint still counts as a fresh run. A hard kill cannot
publish final summaries or status.

Prefill/decode deployments keep step collection and raw Ray metrics. Generic
run summaries are omitted: both roles can report the same prompt and request,
and prefill can emit a placeholder output token. These need separate accounting
before their totals can be used for run comparisons.
### Prefill/decode metrics

For SkyRL-managed prefill/decode deployments, metrics are collected separately
from each role's server actors. Summary keys add a role after the scope:
`vllm_correct_aggregate/eval/prefill/prompt_tokens_total` and
`vllm_correct_aggregate/eval/decode/output_tokens_total`, for example. Step keys
use the same role segment, such as `vllm/eval/prefill/ttft_seconds_avg`.

Each role uses the same metrics and reductions as a regular deployment,
including running/waiting requests, KV usage, throughput and latency. Engine
imbalance is computed within each role. Separate queue charts can show whether
requests are waiting at prefill or decode. Fully async step keys use
`vllm/prefill/...` and `vllm/decode/...`; sync keys also include `train` or `eval`.
Scheduler gauges remain step snapshots and are not added to run summaries.

These are engine observations: both roles can count the same prompt, and
prefill can report a placeholder output token. SkyRL keeps those observations
in their role scopes without adding the two roles' totals or merging their
latency histograms. Missing series and undefined ratios remain omitted.

Role TTFT measures time within that engine's request. Prefill and decode TTFT
describe different parts of serving and cannot be added to obtain client TTFT.
End-to-end latency includes the proxy and transfer path and needs separate
client instrumentation.

Both roles retain the fresh-run, baseline and finalization rules above. A
missing worker or reset omits the affected role's summary. External PD
deployments without known prefill/decode worker membership keep raw/step
collection but omit run summaries.

## Mark runs in Grafana

Expand Down
15 changes: 13 additions & 2 deletions skyrl/train/entrypoints/main_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,12 +294,23 @@ def _setup_trainer(self) -> RayPPOTrainer:
colocate_pg=self.colocate_pg,
)
if trainer._vllm_metrics_scraper is not None:
groups = self._server_groups or ((self._prefill_server_groups or []) + (self._decode_server_groups or []))
enable_pd = self.cfg.generator.inference_engine.enable_pd
groups = (
(self._prefill_server_groups or []) + (self._decode_server_groups or [])
if enable_pd
else self._server_groups
)
actors = [actor for group in (groups or []) for actor in group.get_actors()]
if actors and self.cfg.generator.inference_engine.backend == "vllm":
try:
worker_ids = ray.get([actor.get_ray_worker_id.remote() for actor in actors], timeout=10)
trainer._vllm_metrics_scraper.set_worker_ids(worker_ids)
if enable_pd:
num_prefill = sum(len(group.get_actors()) for group in (self._prefill_server_groups or []))
trainer._vllm_metrics_scraper.set_worker_roles(
{"prefill": worker_ids[:num_prefill], "decode": worker_ids[num_prefill:]}
)
else:
trainer._vllm_metrics_scraper.set_worker_ids(worker_ids)
except Exception as error:
trainer._vllm_metrics_scraper.set_worker_ids([])
logger.warning(
Expand Down
4 changes: 3 additions & 1 deletion skyrl/train/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,7 +245,9 @@ async def finalize_metrics(self, status: str) -> None:
try:
if self._vllm_metrics_scraper is not None:
summary = await asyncio.wait_for(self._vllm_metrics_scraper.finalize(), timeout=10)
if not self._resumed_from_checkpoint and not self.cfg.generator.inference_engine.enable_pd:
if not self._resumed_from_checkpoint and (
not self.cfg.generator.inference_engine.enable_pd or self._vllm_metrics_scraper.has_worker_roles
):
self.tracker.update_summary(summary)
except Exception as e:
logger.warning(f"Could not finalize vLLM metrics: {e}")
Expand Down
58 changes: 58 additions & 0 deletions skyrl/train/utils/vllm_metrics_scraper.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,7 @@ def __init__(
self._urls = urls if urls is not None else discover_ray_metrics_urls()
self._timeout = request_timeout_s
self._worker_ids = None if worker_ids is None else frozenset(worker_ids)
self._role_scrapers: Dict[str, VLLMMetricsScraper] = {}
self.run_statistics = RunStatistics()
self._prev_aggregated: Optional[Dict[str, float]] = None
self._prev_timestamp: Optional[float] = None
Expand All @@ -239,13 +240,42 @@ def __init__(
def set_worker_ids(self, worker_ids: Iterable[str]) -> None:
"""Restrict snapshots to the fixed set of servers launched for this run."""
self._worker_ids = frozenset(worker_ids)
self._role_scrapers = {}

@property
def has_worker_roles(self) -> bool:
"""Whether both PD roles have known, separate worker membership."""
return bool(self._role_scrapers)

def set_worker_roles(self, worker_ids_by_role: Dict[str, Iterable[str]]) -> None:
"""Collect prefill and decode independently using their launched workers."""
groups = {role: frozenset(ids) for role, ids in worker_ids_by_role.items()}
if set(groups) != {"prefill", "decode"} or not all(groups.values()):
raise ValueError("PD metrics require nonempty prefill and decode worker groups")
if not groups["prefill"].isdisjoint(groups["decode"]):
raise ValueError("Prefill and decode metric workers must be disjoint")
self._worker_ids = groups["prefill"] | groups["decode"]
self._role_scrapers = {
role: VLLMMetricsScraper(urls=self._urls, request_timeout_s=self._timeout, worker_ids=ids)
for role, ids in groups.items()
}

def _role_metrics(self, results: List[Dict[str, float]]) -> Dict[str, float]:
"""Keep the existing metrics under separate prefill and decode scopes."""
out = {}
for role, metrics in zip(self._role_scrapers, results):
for key, value in metrics.items():
prefix, name = key.rsplit("/", 1)
out[f"{prefix}/{role}/{name}"] = value
return out

async def _get_client(self) -> httpx.AsyncClient:
if self._client is None:
self._client = httpx.AsyncClient(timeout=self._timeout)
return self._client

async def aclose(self) -> None:
await asyncio.gather(*(scraper.aclose() for scraper in self._role_scrapers.values()))
if self._client is not None:
await self._client.aclose()
self._client = None
Expand Down Expand Up @@ -349,6 +379,10 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]:
async def finalize(self) -> Dict[str, float]:
"""Attempt one final collection and close the HTTP client."""
try:
if self._role_scrapers:
return self._role_metrics(
await asyncio.gather(*(scraper.finalize() for scraper in self._role_scrapers.values()))
)
if self._label is not None:
await self.stop()
else:
Expand All @@ -364,6 +398,10 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f
time since the previous call); ``None`` falls back to the wall-clock
interval (fully-async overlap).
"""
if self._role_scrapers:
return self._role_metrics(
await asyncio.gather(*(scraper.sample(generation_time_s) for scraper in self._role_scrapers.values()))
)
snapshot = await self._read_snapshot()
if snapshot is None:
self.run_statistics.incomplete.add("combined")
Expand Down Expand Up @@ -398,6 +436,14 @@ async def start(self, label: str) -> None:
"""
if self._label is not None:
raise ValueError(f"`start({label!r})` called while window {self._label!r} is still open")
if self._role_scrapers:
await asyncio.gather(*(scraper.start(label) for scraper in self._role_scrapers.values()))
# Both roles begin timing after their baselines are ready.
started = time.monotonic()
for scraper in self._role_scrapers.values():
scraper._active_since = started
self._label = label
return
self._window_prev = await self._read_snapshot()
self._window_engines = self._engine_snapshot
self._label = label
Expand All @@ -407,6 +453,10 @@ async def start(self, label: str) -> None:

def pause(self) -> None:
"""Stop accumulating active time until the next :meth:`resume`."""
if self._role_scrapers:
for scraper in self._role_scrapers.values():
scraper.pause()
return
if self._label is None:
raise ValueError("`pause` called without an open window")
if self._paused:
Expand All @@ -417,6 +467,10 @@ def pause(self) -> None:

def resume(self) -> None:
"""Resume accumulating active time after a :meth:`pause`."""
if self._role_scrapers:
for scraper in self._role_scrapers.values():
scraper.resume()
return
if self._label is None:
raise ValueError("`resume` called without an open window")
if not self._paused:
Expand All @@ -433,6 +487,10 @@ async def stop(self) -> Dict[str, float]:
"""
if self._label is None:
raise ValueError("`stop` called without an open window")
if self._role_scrapers:
result = self._role_metrics(await asyncio.gather(*(s.stop() for s in self._role_scrapers.values())))
self._label = None
return result
new_snapshot = await self._read_snapshot()
if not self._paused:
self._window_time_s += time.monotonic() - self._active_since
Expand Down
11 changes: 7 additions & 4 deletions tests/train/test_metrics_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,21 +49,24 @@ def setup():


@pytest.mark.asyncio
@pytest.mark.parametrize("resumed,enable_pd", [(False, False), (True, False), (False, True)])
async def test_trainer_finalizes_once_and_omits_resumed_or_pd_aggregates(resumed, enable_pd):
@pytest.mark.parametrize(
"resumed,enable_pd,known_roles",
[(False, False, False), (True, False, False), (False, True, False), (False, True, True), (True, True, True)],
)
async def test_trainer_finalizes_once_and_omits_resumed_or_unknown_pd_aggregates(resumed, enable_pd, known_roles):
trainer = RayPPOTrainer.__new__(RayPPOTrainer)
trainer._metrics_finalized = False
trainer.tracker = Tracking("test", "test", backend="console")
trainer._resumed_from_checkpoint = resumed
trainer.cfg = SimpleNamespace(generator=SimpleNamespace(inference_engine=SimpleNamespace(enable_pd=enable_pd)))
trainer.tracker.update_summary = Mock()
scraper = SimpleNamespace(finalize=AsyncMock(return_value={"tokens": 10}))
scraper = SimpleNamespace(finalize=AsyncMock(return_value={"tokens": 10}), has_worker_roles=known_roles)
trainer._vllm_metrics_scraper = scraper
await trainer.finalize_metrics("failed")
await trainer.finalize_metrics("failed")
scraper.finalize.assert_awaited_once()
assert trainer.tracker.run_status == "failed"
if resumed or enable_pd:
if resumed or (enable_pd and not known_roles):
trainer.tracker.update_summary.assert_called_once_with({"run_status": "failed"})
else:
assert trainer.tracker.update_summary.call_args_list[0].args == ({"tokens": 10},)
Expand Down
155 changes: 155 additions & 0 deletions tests/train/test_vllm_pd_metrics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,155 @@
"""Role-specific collection for prefill/decode inference servers."""

from unittest.mock import Mock

import httpx
import pytest

from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper


def _exports(phase, missing_worker=None):
lines = []
for worker, prompt, output, latency in [
("p0", 60, 1, 0.1),
("p1", 40, 1, 0.1),
("d0", 60, 30, 0.3),
("d1", 40, 70, 0.3),
("unrelated", 10000, 10000, 100),
]:
if worker == missing_worker:
continue
labels = f'WorkerId="{worker}",engine="0"'
for name, value in {
"num_requests_running": 2 if worker.startswith("p") else 3,
"num_requests_waiting": 5 if worker.startswith("p") else 1,
"generation_tokens_total": 1000 + phase * output,
"prompt_tokens_total": 1000 + phase * prompt,
"prefix_cache_queries_total": 1000 + phase * prompt,
"prefix_cache_hits_total": 100 + phase * prompt / 2,
}.items():
lines.append(f"ray_vllm_{name}{{{labels}}} {value}")
for name, mean, bound in [
("time_to_first_token_seconds", latency, 0.2 if worker.startswith("p") else 1),
("request_time_per_output_token_seconds", 0.03, 0.05),
]:
count = 1 + phase
lines.extend(
[
f"ray_vllm_{name}_sum{{{labels}}} {mean * count}",
f"ray_vllm_{name}_count{{{labels}}} {count}",
f'ray_vllm_{name}_bucket{{{labels},le="{bound}"}} {count}',
f'ray_vllm_{name}_bucket{{{labels},le="+Inf"}} {count}',
]
)
return "\n".join(lines)


@pytest.mark.asyncio
@pytest.mark.parametrize("sync", [False, True])
async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync):
clock = {"time": 0, "phase": 0}
monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["time"])
scraper = VLLMMetricsScraper(urls=["http://test/metrics"])
scraper.set_worker_roles({"prefill": ["p0", "p1"], "decode": ["d0", "d1"]})

def transport(request):
if sync and clock["phase"] == 0 and request.headers["role"] == "decode":
clock["time"] += 3
return httpx.Response(200, text=_exports(clock["phase"]))

for role, child in scraper._role_scrapers.items():
child._client = httpx.AsyncClient(headers={"role": role}, transport=httpx.MockTransport(transport))
try:
if not sync:
await scraper.sample()
for phase, seconds in [(1, 2), (2, 8)]:
if sync:
await scraper.start("vllm/train")
scraper.pause()
clock["time"] += 5
scraper.resume()
clock["phase"] = phase
clock["time"] += seconds
step = await scraper.stop() if sync else await scraper.sample()
prefix = "vllm/train/" if sync else "vllm/"
assert step[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(100 / seconds)
assert step[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(100 / seconds)
assert step[prefix + "prefill/prompt_throughput_cv"] == pytest.approx(0.2)
assert step[prefix + "decode/generation_throughput_cv"] == pytest.approx(0.4)
assert step[prefix + "prefill/num_requests_waiting"] == 10
assert step[prefix + "decode/num_requests_waiting"] == 2
# The sync window is closed; finalization still closes both role clients.
summary = await scraper.finalize()
finally:
await scraper.aclose()
scope = "train" if sync else "combined"
prefix = f"vllm_correct_aggregate/{scope}/"
assert summary[prefix + "prefill/prompt_tokens_total"] == 200
assert summary[prefix + "decode/output_tokens_total"] == 200
assert summary[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(20)
assert summary[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(20)
assert summary[prefix + "prefill/ttft_seconds_avg"] == pytest.approx(0.1)
assert summary[prefix + "decode/ttft_seconds_avg"] == pytest.approx(0.3)
assert summary[prefix + "prefill/ttft_seconds_p90"] == pytest.approx(0.18)
assert summary[prefix + "decode/ttft_seconds_p90"] == pytest.approx(0.9)
assert summary[prefix + "decode/tpot_seconds_avg"] == pytest.approx(0.03)
assert summary[prefix + "prefill/prefix_cache_hit_rate"] == pytest.approx(0.5)
assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.5)
assert summary[prefix + "prefill/output_tokens_total"] == 4
assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.03)
assert summary[prefix + "decode/prompt_tokens_total"] == 200
assert not any(key.rsplit("/", 2)[-2] not in {"prefill", "decode"} for key in summary)
assert all(child._client is None for child in scraper._role_scrapers.values())


@pytest.mark.asyncio
async def test_missing_prefill_worker_omits_only_prefill_summary(monkeypatch):
clock = {"phase": 0}
monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["phase"])
scraper = VLLMMetricsScraper(urls=["http://test/metrics"])
scraper.set_worker_roles({"prefill": ["p0", "p1"], "decode": ["d0", "d1"]})
for child in scraper._role_scrapers.values():
child._client = httpx.AsyncClient(
transport=httpx.MockTransport(
lambda request: httpx.Response(
200, text=_exports(clock["phase"], "p1" if clock["phase"] == 2 else None)
)
)
)
await scraper.sample()
clock["phase"] = 1
await scraper.sample()
clock["phase"] = 2
summary = await scraper.finalize()
assert not any("/prefill/" in key for key in summary)
assert summary["vllm_correct_aggregate/combined/decode/output_tokens_total"] == 200


def test_setup_uses_server_groups_to_assign_worker_roles(tmp_path, monkeypatch):
from skyrl.train.config import SkyRLTrainConfig
from skyrl.train.entrypoints.main_base import BasePPOExp

cfg = SkyRLTrainConfig()
cfg.generator.inference_engine.enable_pd = True
cfg.trainer.fully_async.simulate_training = True
cfg.trainer.export_path = str(tmp_path / "export")
cfg.trainer.ckpt_path = str(tmp_path / "checkpoints")
exp = BasePPOExp.__new__(BasePPOExp)
exp.cfg = cfg
exp.tokenizer = Mock()
exp.train_dataset = exp.eval_dataset = exp.colocate_pg = None
actors = [Mock() for _ in range(3)]
exp._server_groups = None
exp._prefill_server_groups = [Mock(get_actors=Mock(return_value=actors[:2]))]
exp._decode_server_groups = [Mock(get_actors=Mock(return_value=actors[2:]))]
trainer = Mock()
exp.get_trainer = Mock(return_value=trainer)
for method in ("get_tracker", "get_inference_client", "get_generator", "get_trajectory_logger"):
setattr(exp, method, Mock())
lookup = Mock(return_value=["p0", "p1", "d0"])
monkeypatch.setattr("skyrl.train.entrypoints.main_base.ray.get", lookup)
assert exp._setup_trainer() is trainer
lookup.assert_called_once_with([a.get_ray_worker_id.remote.return_value for a in actors], timeout=10)
trainer._vllm_metrics_scraper.set_worker_roles.assert_called_once_with({"prefill": ["p0", "p1"], "decode": ["d0"]})
trainer._vllm_metrics_scraper.set_worker_ids.assert_not_called()
Loading