diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index 45ee325d5d..e9e9fe8417 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -277,6 +277,7 @@ def _setup_trainer(self) -> RayPPOTrainer: # NOTE (sumanthrh): Instantiate tracker before trainer init. # We have custom validation before this step to give better error messages. tracker = self.get_tracker() + self.tracker = tracker inference_engine_client = self.get_inference_client() @@ -327,10 +328,25 @@ def _setup_trainer(self) -> RayPPOTrainer: def run(self): self.trainer = None + self.tracker = None + status = "failed" try: trainer = self._setup_trainer() + # Start the training loop - asyncio.run(trainer.train()) + async def train_and_finalize(): + run_status = "failed" + try: + await trainer.train() + run_status = "success" + finally: + try: + await trainer.finalize_metrics(run_status) + except Exception as finalization_error: + logger.warning(f"Could not finalize run metrics: {finalization_error}") + + asyncio.run(train_and_finalize()) + status = "success" except Exception as e: # OOMs raised inside actor init (e.g. FSDPPolicyWorkerBase.init_model) # surface here as RayTaskError. Without this they only land in Ray @@ -345,6 +361,14 @@ def run(self): else: logger.error(f"Setup failed before tracker was initialized:\n{e}") raise + finally: + if self.tracker is not None: + try: + self.tracker.run_status = status + self.tracker.update_summary({"run_status": status}) + self.tracker.finish(exit_code=0 if status == "success" else 1) + except Exception as finalization_error: + logger.warning(f"Could not finish run tracking: {finalization_error}") @ray.remote(num_cpus=1) diff --git a/skyrl/train/fully_async_trainer.py b/skyrl/train/fully_async_trainer.py index 457562d4ba..c3685cf969 100644 --- a/skyrl/train/fully_async_trainer.py +++ b/skyrl/train/fully_async_trainer.py @@ -482,6 +482,9 @@ async def train(self): if self._ray_gpu_monitor is not None: self._ray_gpu_monitor.start() + if self._vllm_metrics_scraper is not None: + await self._vllm_metrics_scraper.sample() + # Eval before training if self.cfg.trainer.eval_interval > 0 and self.cfg.trainer.eval_before_train: with self._phase_gauge.timed_phase("eval", self.all_timings): @@ -731,8 +734,7 @@ async def _watch_generators_done(tasks=generator_tasks, event=all_generators_don if self.has_critic: self.dispatch.finalize_pending_saves("critic") - if self._vllm_metrics_scraper is not None: - await self._vllm_metrics_scraper.aclose() + await self.finalize_metrics("success") self.tracker.finish() logger.info("Training done!") diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index b3ad7023fd..49a874d9d5 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -1,3 +1,4 @@ +import asyncio import math import os import shutil @@ -144,6 +145,9 @@ def __init__( VLLMMetricsScraper() if cfg.generator.inference_engine.enable_ray_prometheus_stats else None ) + self._metrics_finalized = False + self._resumed_from_checkpoint = False + self._ray_gpu_monitor = RayGpuMonitor() if cfg.trainer.enable_ray_gpu_monitor else None # trajectory logger is installed after construction if needed @@ -232,6 +236,22 @@ def _build_train_dataloader_and_compute_training_steps(self): if self.cfg.trainer.max_training_steps is not None: self.total_training_steps = min(self.total_training_steps, self.cfg.trainer.max_training_steps) + async def finalize_metrics(self, status: str) -> None: + """Finalize observations before the tracker closes, including failed runs.""" + if self._metrics_finalized: + return + self._metrics_finalized = True + self.tracker.run_status = status + 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: + self.tracker.update_summary(summary) + except Exception as e: + logger.warning(f"Could not finalize vLLM metrics: {e}") + finally: + self.tracker.update_summary({"run_status": status}) + @torch.no_grad() async def eval(self, vllm_metrics_scraper: Optional[VLLMMetricsScraper] = None) -> Dict[str, float]: """ @@ -598,9 +618,7 @@ async def train(self): if self.has_critic: self.dispatch.finalize_pending_saves("critic") - if self._vllm_metrics_scraper is not None: - await self._vllm_metrics_scraper.aclose() - + await self.finalize_metrics("success") if self._ray_gpu_monitor is not None: self._ray_gpu_monitor.stop() @@ -1819,6 +1837,7 @@ def load_checkpoints(self) -> Tuple[int, str]: # 1. Load and validate trainer state with io.open_file(trainer_state_path, "rb") as f: trainer_state = torch.load(f, map_location="cpu", weights_only=False) + self._resumed_from_checkpoint = True saved_global_step = trainer_state.get("global_step", global_step) logger.info("Successfully loaded trainer state") if saved_global_step != global_step: diff --git a/skyrl/train/utils/tracking.py b/skyrl/train/utils/tracking.py index 0c2b42fbf7..4f5310f07f 100644 --- a/skyrl/train/utils/tracking.py +++ b/skyrl/train/utils/tracking.py @@ -75,14 +75,44 @@ def __init__( self.logger = ConsoleLogger() self._exception_logged = False + self._finished = False + self.run_status = "running" + self._vllm_history_metrics = set() + + def update_summary(self, data) -> None: + """Write run summaries without adding another training history step.""" + try: + if self.backend == "wandb" and not self._finished and self.logger.run is not None: + self.logger.run.summary.update(data) + except Exception as e: + logger.warning(f"Could not update run summary: {e}") def log(self, data, step, commit=False): if self.backend == "wandb": + for name in data: + if name.startswith("vllm/") and name not in self._vllm_history_metrics: + self.logger.define_metric(name, summary="none") + self._vllm_history_metrics.add(name) self.logger.log(data=data, step=step, commit=commit) else: self.logger.log(data=data, step=step) - def finish(self): + def finish(self, exit_code: int = 0): + if self._finished: + return + if self.run_status == "running": + self.run_status = "success" if exit_code == 0 else "failed" + if self.backend == "wandb" and self.logger.run is not None: + # Explicit removals also clear summaries already inferred from history. + for name in self._vllm_history_metrics: + try: + del self.logger.run.summary[name] + except KeyError: + pass + except Exception as e: + logger.warning(f"Could not remove automatic metric summary {name}: {e}") + self.update_summary({"run_status": self.run_status}) + self._finished = True if self.backend == "console": return # NOTE (sumanthrh): We use a try-except block here while finishing tracking. @@ -90,7 +120,7 @@ def finish(self): # https://github.com/wandb/wandb/issues/6449 try: if self.backend == "wandb": - self.logger.finish(exit_code=0) + self.logger.finish(exit_code=exit_code) else: self.logger.finish() except Exception as e: @@ -113,6 +143,7 @@ def log_exception(self, e: BaseException, step: int = 0) -> None: if self._exception_logged: return self._exception_logged = True + self.run_status = "failed" tb_str = traceback.format_exc()[-10000:] logger.error(f"Training failed at step {step} with {type(e).__name__}:\n{tb_str}") if self.backend == "wandb": @@ -128,7 +159,7 @@ def log_exception(self, e: BaseException, step: int = 0) -> None: # Tables upload asynchronously. Finish the run so the upload # completes before the caller re-raises and the process dies. try: - self.finish() + self.finish(exit_code=1) except Exception as finish_exc: logger.warning(f"tracker.finish() raised after logging exception: {finish_exc}") except Exception as log_exc: diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 0b8d884ecc..a2a5c3693a 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -22,6 +22,7 @@ from loguru import logger from skyrl.backends.skyrl_train.inference_servers.common import format_http_url +from skyrl.train.utils.vllm_run_statistics import RunStatistics from skyrl.train.utils.vllm_window_statistics import latency_metrics # vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`) @@ -39,6 +40,7 @@ _COUNTER_PREEMPTIONS = "ray_vllm_num_preemptions_total" _COUNTER_EXTERNAL_PREFIX_QUERIES = "ray_vllm_external_prefix_cache_queries_total" _COUNTER_EXTERNAL_PREFIX_HITS = "ray_vllm_external_prefix_cache_hits_total" +_COUNTER_KV_OFFLOAD_STORE_BYTES = "ray_vllm_kv_offload_store_bytes_total" _COUNTER_KV_OFFLOAD_LOAD_BYTES = "ray_vllm_kv_offload_load_bytes_total" _HIST_TTFT_SUM = "ray_vllm_time_to_first_token_seconds_sum" _HIST_TTFT_COUNT = "ray_vllm_time_to_first_token_seconds_count" @@ -69,6 +71,7 @@ _COUNTER_PREEMPTIONS, _COUNTER_EXTERNAL_PREFIX_QUERIES, _COUNTER_EXTERNAL_PREFIX_HITS, + _COUNTER_KV_OFFLOAD_STORE_BYTES, _COUNTER_KV_OFFLOAD_LOAD_BYTES, _HIST_REQUEST_TPOT_SUM, _HIST_REQUEST_TPOT_COUNT, @@ -211,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.run_statistics = RunStatistics() self._prev_aggregated: Optional[Dict[str, float]] = None self._prev_timestamp: Optional[float] = None self._client: Optional[httpx.AsyncClient] = None @@ -306,6 +310,8 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: counters = engine_counters.setdefault(engine, {}) counters[name] = counters.get(name, 0.0) + value self._engine_snapshot = engine_counters + if self._worker_ids is not None and {engine[0] for engine in engine_counters} != self._worker_ids: + return None buckets = {} schemas = {} for (name, labels), value in parsed.items(): @@ -328,14 +334,29 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: # Ray counters skip zero increments. A live engine gauge confirms the exporter exists. if any(name == _GAUGE_NUM_RUNNING for name, _ in parsed): sums.setdefault(_COUNTER_PREEMPTIONS, 0.0) + sums.setdefault(_COUNTER_PROMPT_TOKENS, 0.0) + sums.setdefault(_COUNTER_GENERATION_TOKENS, 0.0) for hits, queries in ( (_COUNTER_PREFIX_HITS, _COUNTER_PREFIX_QUERIES), (_COUNTER_EXTERNAL_PREFIX_HITS, _COUNTER_EXTERNAL_PREFIX_QUERIES), ): if queries in sums: sums.setdefault(hits, 0.0) + if _COUNTER_KV_OFFLOAD_STORE_BYTES in sums: + sums.setdefault(_COUNTER_KV_OFFLOAD_LOAD_BYTES, 0.0) return {**sums, **means, **per_pos, **buckets} + async def finalize(self) -> Dict[str, float]: + """Attempt one final collection and close the HTTP client.""" + try: + if self._label is not None: + await self.stop() + else: + await self.sample() + return self.run_statistics.summary() + finally: + await self.aclose() + async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, float]: """Return ``vllm/...`` scalars for the current step (empty if unavailable). @@ -345,6 +366,7 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f """ snapshot = await self._read_snapshot() if snapshot is None: + self.run_statistics.incomplete.add("combined") self._prev_aggregated = None self._prev_timestamp = None self._previous_engines = {} @@ -355,6 +377,7 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f dt = max(now - self._prev_timestamp, 1e-9) window = generation_time_s if (generation_time_s is not None and generation_time_s > 0) else dt out = self._window_metrics(self._prev_aggregated, snapshot, window, "vllm/") + self.run_statistics.add("combined", self._prev_aggregated, snapshot, window) else: out = self._window_metrics(None, snapshot, None, "vllm/") # gauges only @@ -418,6 +441,7 @@ async def stop(self) -> Dict[str, float]: self._window_prev = None self._active_since = None self._paused = False + self.run_statistics.add(label.removeprefix("vllm/"), prev, new_snapshot, window) if new_snapshot is None: return {} result = self._window_metrics(prev, new_snapshot, window, f"{label}/") diff --git a/skyrl/train/utils/vllm_run_statistics.py b/skyrl/train/utils/vllm_run_statistics.py new file mode 100644 index 0000000000..143b8ba6c2 --- /dev/null +++ b/skyrl/train/utils/vllm_run_statistics.py @@ -0,0 +1,55 @@ +"""Weighted summaries of observed, non-overlapping vLLM windows.""" + +from skyrl.train.utils.vllm_window_statistics import counter_deltas + + +class RunStatistics: + """Sum raw counters and durations, keeping sync train and eval separate.""" + + def __init__(self): + self.windows = {} + self.seconds = {} + self.incomplete = set() + + def add(self, scope, previous, current, duration): + """Omit a scope with a missing/reset window; never average step rates.""" + deltas = counter_deltas(previous, current) + if deltas is None: + self.incomplete.add(scope) + return + if scope not in self.windows: + self.windows[scope] = deltas + else: + total = self.windows[scope] + for name in deltas: + if "_bucket::" in name and total.get(name.split("_bucket::", 1)[0] + "_count") == 0: + total.setdefault(name, 0.0) + # Only counters observed over every window have full-run denominators. + self.windows[scope] = {name: value + deltas[name] for name, value in total.items() if name in deltas} + self.seconds[scope] = self.seconds.get(scope, 0.0) + duration + + def summary(self): + """Derive rates, request-weighted means and merged histogram P90 values.""" + from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper + + result = {} + for scope, deltas in self.windows.items(): + if scope in self.incomplete: + continue + prefix = f"vllm_correct_aggregate/{scope}/" + metrics = VLLMMetricsScraper._derive(deltas, dict.fromkeys(deltas, 0), self.seconds[scope], prefix) + result.update( + {key: value for key, value in metrics.items() if "draft_num_" not in key and "_pos_" not in key} + ) + result[prefix + "measurement_seconds"] = self.seconds[scope] + for counter, public in ( + ("generation_tokens", "output_tokens_total"), + ("prompt_tokens", "prompt_tokens_total"), + ("num_preemptions", "preemptions_total"), + ("kv_offload_store_bytes", "kv_offload_store_bytes_total"), + ("kv_offload_load_bytes", "kv_offload_load_bytes_total"), + ): + value = deltas.get(f"ray_vllm_{counter}_total") + if value is not None: + result[prefix + public] = value + return result diff --git a/tests/train/test_metrics_lifecycle.py b/tests/train/test_metrics_lifecycle.py new file mode 100644 index 0000000000..ef237ce371 --- /dev/null +++ b/tests/train/test_metrics_lifecycle.py @@ -0,0 +1,175 @@ +"""Tests for metrics finalization on success and controlled failures.""" + +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +import pytest +import torch + +from skyrl.train.entrypoints.main_base import BasePPOExp +from skyrl.train.trainer import RayPPOTrainer +from skyrl.train.utils.tracking import Tracking +from skyrl.train.utils.trainer_utils import ResumeMode + + +@pytest.mark.parametrize("fail", [False, True]) +def test_entrypoint_finalizes_before_exception_logging(fail): + events = [] + tracker = Tracking("test", "test", backend="console") + tracker.log_exception = Mock(side_effect=lambda *args, **kwargs: events.append("exception")) + trainer = SimpleNamespace(tracker=tracker, global_step=0, flush_pending_metrics=Mock()) + + async def train(): + if fail: + raise ValueError("controlled training failure") + + async def finalize(status): + events.append(status) + tracker.run_status = status + + trainer.train = train + trainer.finalize_metrics = finalize + exp = BasePPOExp.__new__(BasePPOExp) + + def setup(): + exp.trainer = trainer + exp.tracker = tracker + return trainer + + exp._setup_trainer = setup + if fail: + with pytest.raises(ValueError, match="controlled training failure"): + exp.run() + assert events == ["failed", "exception"] + else: + exp.run() + assert events == ["success"] + assert tracker.run_status == ("failed" if fail else "success") + + +@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): + 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})) + 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: + trainer.tracker.update_summary.assert_called_once_with({"run_status": "failed"}) + else: + assert trainer.tracker.update_summary.call_args_list[0].args == ({"tokens": 10},) + + +@pytest.mark.asyncio +async def test_failed_terminal_collection_does_not_publish_aggregates(): + trainer = RayPPOTrainer.__new__(RayPPOTrainer) + trainer._metrics_finalized = False + trainer.tracker = Tracking("test", "test", backend="console") + trainer._resumed_from_checkpoint = False + trainer.tracker.update_summary = Mock() + trainer._vllm_metrics_scraper = SimpleNamespace( + finalize=AsyncMock(side_effect=RuntimeError("scrape failed")), + ) + await trainer.finalize_metrics("failed") + trainer.tracker.update_summary.assert_called_once_with({"run_status": "failed"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [RuntimeError("scrape failed"), asyncio.CancelledError()]) +async def test_terminal_collection_always_closes_http_client(error): + from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper + + scraper = VLLMMetricsScraper(urls=["test"]) + scraper._prev_timestamp = 0 + client = SimpleNamespace(aclose=AsyncMock()) + scraper._client = client + scraper.sample = AsyncMock(side_effect=error) + with pytest.raises(type(error)): + await scraper.finalize() + client.aclose.assert_awaited_once() + assert scraper._client is None + + +@pytest.mark.asyncio +async def test_final_collection_subtracts_reused_external_engine_baseline(): + from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper + + counter = "ray_vllm_generation_tokens_total" + scraper = VLLMMetricsScraper(urls=["test"]) + scraper._read_snapshot = AsyncMock(side_effect=[{counter: 10000}, {counter: 10100}]) + with patch("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", side_effect=[10, 12]): + await scraper.sample() + summary = await scraper.finalize() + assert summary["vllm_correct_aggregate/combined/output_tokens_total"] == 100 + assert summary["vllm_correct_aggregate/combined/generation_throughput_tok_s"] == 50 + assert scraper._read_snapshot.await_count == 2 + + +@pytest.mark.asyncio +async def test_unavailable_initial_export_does_not_create_a_partial_run_total(): + from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper + + counter = "ray_vllm_generation_tokens_total" + scraper = VLLMMetricsScraper(urls=["test"]) + scraper._read_snapshot = AsyncMock(side_effect=[None, {counter: 100}, {counter: 200}]) + await scraper.sample() + await scraper.sample(generation_time_s=2) + await scraper.sample(generation_time_s=2) + assert scraper.run_statistics.summary() == {} + + +@pytest.mark.parametrize("enable_pd", [False, True]) +def test_pd_preserves_step_metrics_collection(enable_pd): + from unittest.mock import patch + + from skyrl.train.config import SkyRLTrainConfig + + cfg = SkyRLTrainConfig() + cfg.generator.inference_engine.enable_ray_prometheus_stats = True + cfg.generator.inference_engine.enable_pd = enable_pd + cfg.trainer.enable_ray_gpu_monitor = False + with patch("skyrl.train.trainer.VLLMMetricsScraper") as scraper: + trainer = RayPPOTrainer( + cfg, + Tracking("test", "test", backend="console"), + tokenizer=Mock(), + train_dataset=None, + inference_engine_client=Mock(), + generator=Mock(), + ) + assert trainer._vllm_metrics_scraper is not None + assert scraper.call_count == 1 + + +@pytest.mark.parametrize("checkpoint_exists", [False, True]) +def test_resume_guard_distinguishes_missing_latest_from_loaded_step_zero(tmp_path, checkpoint_exists): + from skyrl.train.config import SkyRLTrainConfig + + trainer = RayPPOTrainer.__new__(RayPPOTrainer) + trainer.cfg = SkyRLTrainConfig() + trainer.cfg.trainer.ckpt_path = str(tmp_path) + trainer.cfg.trainer.critic.model.path = "" + trainer._resumed_from_checkpoint = False + trainer.dispatch = Mock() + if checkpoint_exists: + checkpoint = tmp_path / "global_step_0" + checkpoint.mkdir() + torch.save({"global_step": 0}, checkpoint / "trainer_state.pt") + trainer.cfg.trainer.resume_path = str(checkpoint) + trainer.resume_mode = ResumeMode.FROM_PATH + else: + trainer.resume_mode = ResumeMode.LATEST + + step, path = trainer.load_checkpoints() + assert step == 0 + assert trainer._resumed_from_checkpoint == checkpoint_exists + assert (path is not None) == checkpoint_exists diff --git a/tests/train/test_tracking.py b/tests/train/test_tracking.py index 7018ea920e..1a3884fb4f 100644 --- a/tests/train/test_tracking.py +++ b/tests/train/test_tracking.py @@ -39,3 +39,29 @@ def test_wandb_init_tags_default_none(): wandb_mock.init.assert_called_once() assert wandb_mock.init.call_args.kwargs["tags"] is None + + +def test_vllm_history_metrics_disable_automatic_summaries_once(): + with patch.dict("sys.modules", {"wandb": MagicMock()}) as mocked: + wandb_mock = mocked["wandb"] + tracker = Tracking("proj", "exp", backend="wandb", config={}) + tracker.log({"vllm/train/generation_throughput_tok_s": 50, "train/loss": 1}, step=1) + tracker.log({"vllm/train/generation_throughput_tok_s": 25}, step=2) + wandb_mock.define_metric.assert_called_once_with("vllm/train/generation_throughput_tok_s", summary="none") + assert wandb_mock.log.call_count == 2 + assert wandb_mock.log.call_args.kwargs["data"] == {"vllm/train/generation_throughput_tok_s": 25} + tracker.finish() + + +def test_finish_removes_regular_sdk_summary_but_preserves_other_metrics(): + with patch.dict("sys.modules", {"wandb": MagicMock()}) as mocked: + wandb_mock = mocked["wandb"] + summary = {"vllm/train/rate": 25, "train/loss": 1} + wandb_mock.run.summary = summary + tracker = Tracking("proj", "exp", backend="wandb", config={}) + tracker.log({"vllm/train/rate": 25}, step=1) + tracker.update_summary({"vllm_correct_aggregate/train/rate": 30}) + tracker.finish() + assert summary == {"train/loss": 1, "vllm_correct_aggregate/train/rate": 30, "run_status": "success"} + wandb_mock.Api.assert_not_called() + wandb_mock.finish.assert_called_once() diff --git a/tests/train/test_vllm_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index 9befb25189..4e4a2ac605 100644 --- a/tests/train/test_vllm_metrics_scraper.py +++ b/tests/train/test_vllm_metrics_scraper.py @@ -699,6 +699,48 @@ async def test_worker_filter_and_merged_latency_buckets(): assert snapshot["ray_vllm_time_to_first_token_seconds_bucket::1"] == 2 +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", [False, True]) +async def test_cold_start_tokens_are_included_in_step_and_run_metrics(monkeypatch, sync): + phase = 0 + now = 0.0 + monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time", Mock(monotonic=lambda: now)) + + def respond(request): + text = 'ray_vllm_num_requests_running{WorkerId="worker"} 0\n' + if phase: + generated, prompted = {1: (100, 200), 2: (300, 500)}[phase] + text += ( + f'ray_vllm_generation_tokens_total{{WorkerId="worker"}} {generated}\n' + f'ray_vllm_prompt_tokens_total{{WorkerId="worker"}} {prompted}\n' + ) + return httpx.Response(200, text=text) + + scraper = VLLMMetricsScraper(urls=["http://test/metrics"], worker_ids=["worker"]) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + scraper._client = client + if not sync: + await scraper.sample() + for next_phase, seconds in ((1, 2), (2, 8)): + if sync: + await scraper.start("vllm/train") + phase = next_phase + now += seconds + metrics = await scraper.stop() if sync else await scraper.sample() + if phase == 1: + prefix = "vllm/train/" if sync else "vllm/" + assert metrics[prefix + "generation_throughput_tok_s"] == 50 + assert metrics[prefix + "prompt_throughput_tok_s"] == 100 + summary = scraper.run_statistics.summary() + + prefix = "vllm_correct_aggregate/train/" if sync else "vllm_correct_aggregate/combined/" + assert summary[prefix + "output_tokens_total"] == 300 + assert summary[prefix + "prompt_tokens_total"] == 500 + assert summary[prefix + "generation_throughput_tok_s"] == 30 + assert summary[prefix + "prompt_throughput_tok_s"] == 50 + assert summary[prefix + "measurement_seconds"] == 10 + + @pytest.mark.asyncio async def test_failed_node_scrape_skips_partial_totals_and_reestablishes_baseline(): step = 0 diff --git a/tests/train/test_vllm_run_statistics.py b/tests/train/test_vllm_run_statistics.py new file mode 100644 index 0000000000..697679a486 --- /dev/null +++ b/tests/train/test_vllm_run_statistics.py @@ -0,0 +1,69 @@ +"""Tests for weighted summaries from observed counter windows.""" + +import pytest + +from skyrl.train.utils.vllm_run_statistics import RunStatistics + + +def test_weighted_rates_and_cache_hits_use_raw_totals(): + counter = "ray_vllm_generation_tokens_total" + queries = "ray_vllm_external_prefix_cache_queries_total" + hits = "ray_vllm_external_prefix_cache_hits_total" + stats = RunStatistics() + first = {counter: 100, queries: 10, hits: 5} + stats.add("train", dict.fromkeys(first, 0), first, 2) + stats.add("train", first, {counter: 300, queries: 100, hits: 14}, 8) + summary = stats.summary() + assert summary["vllm_correct_aggregate/train/generation_throughput_tok_s"] == 30 + assert summary["vllm_correct_aggregate/train/external_prefix_cache_hit_rate"] == pytest.approx(0.14) + assert summary["vllm_correct_aggregate/train/output_tokens_total"] == 300 + assert summary["vllm_correct_aggregate/train/measurement_seconds"] == 10 + + +@pytest.mark.parametrize("terminal", [None, {"ray_vllm_generation_tokens_total": 1}]) +def test_missing_or_reset_window_omits_scope(terminal): + counter = "ray_vllm_generation_tokens_total" + stats = RunStatistics() + stats.add("train", {counter: 0}, {counter: 100}, 2) + stats.add("train", {counter: 100}, terminal, 3) + assert stats.summary() == {} + + +def test_missing_counter_is_not_reported_with_a_full_run_denominator(): + counter = "ray_vllm_generation_tokens_total" + stats = RunStatistics() + stats.add("train", {counter: 0}, {counter: 100}, 2) + stats.add("train", {}, {counter: 300}, 3) + assert "vllm_correct_aggregate/train/output_tokens_total" not in stats.summary() + assert "vllm_correct_aggregate/train/generation_throughput_tok_s" not in stats.summary() + + +def test_run_tpot_is_request_weighted_and_tracker_metrics_are_pruned(): + stats = RunStatistics() + base = "ray_vllm_request_time_per_output_token_seconds" + ttft = "ray_vllm_time_to_first_token_seconds" + for count, total in [(1, 0.8), (3, 0.6)]: + deltas = { + base + "_count": count, + base + "_sum": total, + base + "_bucket::1": count / 2, + base + "_bucket::2": count, + base + "_bucket::+Inf": count, + ttft + "_count": count, + ttft + "_sum": total, + ttft + "_bucket::1": count / 2, + ttft + "_bucket::2": count, + ttft + "_bucket::+Inf": count, + "ray_vllm_inter_token_latency_seconds_count": 100, + "ray_vllm_inter_token_latency_seconds_sum": 0.1, + "ray_vllm_kv_offload_store_bytes_total": 1000, + } + stats.add("train", dict.fromkeys(deltas, 0), deltas, 2) + summary = stats.summary() + assert summary["vllm_correct_aggregate/train/tpot_seconds_avg"] == pytest.approx(1.4 / 4) + assert summary["vllm_correct_aggregate/train/tpot_seconds_p90"] == pytest.approx(1.8) + assert summary["vllm_correct_aggregate/train/ttft_seconds_p90"] == pytest.approx(1.8) + assert summary["vllm_correct_aggregate/train/kv_offload_store_bytes_total"] == 2000 + assert not any( + "itl" in key or "request_tpot" in key or "p50" in key or "store_throughput" in key for key in summary + )