diff --git a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx index 980cef0c8b..d29c3bbded 100644 --- a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx +++ b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx @@ -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 diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index 8fb7259427..765b714c07 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -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( diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index 49a874d9d5..c94e7280d3 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -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}") diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index a2a5c3693a..3338026475 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -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 @@ -239,6 +240,34 @@ 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: @@ -246,6 +275,7 @@ async def _get_client(self) -> httpx.AsyncClient: 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 @@ -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: @@ -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") @@ -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 @@ -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: @@ -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: @@ -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 diff --git a/tests/train/generators/chat_templating_test_constants.py b/tests/train/generators/chat_templating_test_constants.py index 55767e7538..dc27cc7bf1 100644 --- a/tests/train/generators/chat_templating_test_constants.py +++ b/tests/train/generators/chat_templating_test_constants.py @@ -3,8 +3,6 @@ tests/train/generators/test_skyrl_gym_generator_chat_templating.py::test_skyrl_gym_generator_chat_templating_exact """ -from datetime import date - # Produced by expected_str = tokenizer.apply_chat_template(expected_chat_history, tokenize=False) # where expected_chat_history is: @@ -35,10 +33,10 @@ def get_expected_chat_history(mock_response_text: str): b<|im_end|> """ -LLAMA3_2_EXPECTED_STR = f"""<|begin_of_text|><|start_header_id|>system<|end_header_id|> +LLAMA3_2_EXPECTED_STR = """<|begin_of_text|><|start_header_id|>system<|end_header_id|> Cutting Knowledge Date: December 2023 -Today Date: {date.today().strftime("%d %b %Y")} +Today Date: 26 Jul 2024 <|eot_id|><|start_header_id|>user<|end_header_id|> diff --git a/tests/train/generators/test_skyrl_gym_generator_chat_templating.py b/tests/train/generators/test_skyrl_gym_generator_chat_templating.py index 6b210f4092..6ffc428021 100644 --- a/tests/train/generators/test_skyrl_gym_generator_chat_templating.py +++ b/tests/train/generators/test_skyrl_gym_generator_chat_templating.py @@ -2,6 +2,7 @@ uv run --extra dev --isolated pytest tests/train/generators/test_skyrl_gym_generator_chat_templating.py """ +from datetime import datetime from pathlib import Path from typing import Any, Dict from unittest.mock import AsyncMock, MagicMock @@ -134,7 +135,7 @@ def _make_input_batch(prompt, extras): "qwen3-custom_chat_template_builtin", ], ) -async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenization_codepath, expected_str): +async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenization_codepath, expected_str, monkeypatch): """ Tests the behavior of chat templating for various models in multi-turn conversation. @@ -144,6 +145,11 @@ async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenizatio We hardcode the expected string in the constants file, so it is easier to check. But we also double check that those expected strings are correct by applying the chat template on the expected chat history. """ + # Keep the template's clock consistent with the fixed expected date across midnight. + template_datetime = MagicMock() + template_datetime.now.return_value = datetime(2024, 7, 26) + monkeypatch.setattr("transformers.utils.chat_template_utils.datetime", template_datetime) + # 1. Preparations to mock the generation. _register_test_env_if_needed() # Register only when needed tokenizer = AutoTokenizer.from_pretrained(model_name) diff --git a/tests/train/test_metrics_lifecycle.py b/tests/train/test_metrics_lifecycle.py index ef237ce371..22b405a100 100644 --- a/tests/train/test_metrics_lifecycle.py +++ b/tests/train/test_metrics_lifecycle.py @@ -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},) diff --git a/tests/train/test_vllm_pd_metrics.py b/tests/train/test_vllm_pd_metrics.py new file mode 100644 index 0000000000..c97261b1fa --- /dev/null +++ b/tests/train/test_vllm_pd_metrics.py @@ -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()