From 49686b2a40b2744c0ccc0d27753f124b2316942a Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 30 Sep 2026 16:18:05 +0000 Subject: [PATCH 1/6] feat: collect scoped vLLM offload and latency window metrics Signed-off-by: Codex Signed-off-by: SumanthRH --- .../inference_servers/vllm_server_actor.py | 6 ++ skyrl/train/entrypoints/main_base.py | 6 ++ skyrl/train/utils/vllm_metrics_scraper.py | 78 +++++++++++++++++- skyrl/train/utils/vllm_window_statistics.py | 79 +++++++++++++++++++ tests/train/test_vllm_metrics_scraper.py | 34 ++++++++ tests/train/test_vllm_window_statistics.py | 51 ++++++++++++ 6 files changed, 253 insertions(+), 1 deletion(-) create mode 100644 skyrl/train/utils/vllm_window_statistics.py create mode 100644 tests/train/test_vllm_window_statistics.py diff --git a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py index 525951284f..7ca5a444aa 100644 --- a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py +++ b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py @@ -376,6 +376,12 @@ def _setup_mooncake_port(self, mooncake_server_port: int) -> None: f"host={self._ip}, port={mooncake_server_port}, engine_id={engine_id}" ) + def get_metrics_worker_id(self) -> str: + """Return the Ray worker identity used by this server's metrics logger.""" + import ray + + return ray.get_runtime_context().get_worker_id() + def get_server_info(self) -> ServerInfo: """Get the server's IP and port info.""" return ServerInfo( diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index b50efdbe47..b3bd60b76c 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -292,6 +292,12 @@ def _setup_trainer(self) -> RayPPOTrainer: generator=generator, 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 [])) + actors = [actor for group in (groups or []) for actor in group.get_actors()] + if actors and self.cfg.generator.inference_engine.backend == "vllm": + worker_ids = ray.get([actor.get_metrics_worker_id.remote() for actor in actors]) + trainer._vllm_metrics_scraper.set_worker_ids(worker_ids) # Install the trajectory logger after construction trainer.trajectory_logger = self.get_trajectory_logger() # Expose the trainer on self so callers can log exceptions raised diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 9178c47715..5c65e5b102 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -20,6 +20,7 @@ from loguru import logger from skyrl.backends.skyrl_train.inference_servers.common import format_http_url +from skyrl.train.utils.vllm_window_statistics import WindowStatistics # vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`) # AND the `ray_` prefix that Ray's metrics agent adds to every custom metric. @@ -60,6 +61,16 @@ _COUNTER_SPEC_DRAFT_TOKENS, _COUNTER_SPEC_ACCEPTED_TOKENS, ) +_ADDITIONAL_COUNTERS = ( + "ray_vllm_num_preemptions_total", + "ray_vllm_external_prefix_cache_queries_total", + "ray_vllm_external_prefix_cache_hits_total", + "ray_vllm_kv_offload_store_bytes_total", + "ray_vllm_kv_offload_load_bytes_total", + "ray_vllm_request_time_per_output_token_seconds_sum", + "ray_vllm_request_time_per_output_token_seconds_count", +) +_SUM_METRICS += _ADDITIONAL_COUNTERS _MEAN_METRICS = (_GAUGE_KV_CACHE_USAGE,) ParsedSamples = Dict[Tuple[str, FrozenSet[Tuple[str, str]]], float] @@ -193,9 +204,12 @@ def __init__( self, urls: Optional[List[str]] = None, request_timeout_s: float = 2.0, + worker_ids: Optional[Iterable[str]] = None, ): 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.last_window = WindowStatistics(valid=False) self._prev_aggregated: Optional[Dict[str, float]] = None self._prev_timestamp: Optional[float] = None self._client: Optional[httpx.AsyncClient] = None @@ -213,6 +227,10 @@ def __init__( "engine metrics will not appear in wandb." ) + 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) + async def _get_client(self) -> httpx.AsyncClient: if self._client is None: self._client = httpx.AsyncClient(timeout=self._timeout) @@ -262,10 +280,47 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: ) self._warned_empty = True + if self._worker_ids is not None: + parsed = {key: value for key, value in parsed.items() if dict(key[1]).get("WorkerId") in self._worker_ids} + if not parsed: + return None + buckets = {} + schemas = {} + for (name, labels), value in parsed.items(): + if name not in { + "ray_vllm_time_to_first_token_seconds_bucket", + "ray_vllm_request_time_per_output_token_seconds_bucket", + }: + continue + label_dict = dict(labels) + bound = label_dict.get("le") + if bound is None: + continue + owner = frozenset((k, v) for k, v in labels if k != "le") + schemas.setdefault(name, {}).setdefault(owner, set()).add(bound) + key = f"{name}::{bound}" + buckets[key] = buckets.get(key, 0.0) + value + for name, owners in schemas.items(): + if len({frozenset(bounds) for bounds in owners.values()}) != 1: + buckets = {key: value for key, value in buckets.items() if not key.startswith(name + "::")} + # Ray counters skip zero increments. A live engine gauge confirms the + # exporter exists; connector queries/store counters establish support. + if any(name == _GAUGE_NUM_RUNNING for name, _ in parsed): + buckets.setdefault("ray_vllm_num_preemptions_total", 0.0) sums = aggregate(parsed, _SUM_METRICS, how="sum") means = aggregate(parsed, _MEAN_METRICS, how="mean") per_pos = sum_by_position(parsed, _COUNTER_SPEC_ACCEPTED_PER_POS) - return {**sums, **means, **per_pos} + if "ray_vllm_num_preemptions_total" in sums: + buckets.pop("ray_vllm_num_preemptions_total", None) + for hits, queries in ( + ("ray_vllm_prefix_cache_hits_total", _COUNTER_PREFIX_QUERIES), + ("ray_vllm_external_prefix_cache_hits_total", "ray_vllm_external_prefix_cache_queries_total"), + ): + if queries in sums: + sums.setdefault(hits, 0.0) + if "ray_vllm_kv_offload_store_bytes_total" in sums: + sums.setdefault("ray_vllm_kv_offload_load_bytes_total", 0.0) + return {**sums, **means, **per_pos, **buckets} async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, float]: """Return ``vllm/...`` scalars for the current step (empty if unavailable). @@ -286,6 +341,9 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f else: out = self._window_metrics(None, snapshot, None, "vllm/") # gauges only + self.last_window = WindowStatistics.between( + self._prev_aggregated, snapshot, window if self._prev_timestamp is not None else None + ) self._prev_aggregated = snapshot self._prev_timestamp = now return out @@ -342,6 +400,7 @@ async def stop(self) -> Dict[str, float]: self._window_prev = None self._active_since = None self._paused = False + self.last_window = WindowStatistics.between(prev, new_snapshot, window) if new_snapshot is None: return {} return self._window_metrics(prev, new_snapshot, window, f"{label}/") @@ -406,6 +465,23 @@ def delta(name: str) -> Optional[float]: itl_count_d = delta(_HIST_ITL_COUNT) if itl_sum_d is not None and itl_count_d is not None and itl_count_d > 0: out[f"{prefix}tpot_seconds_avg"] = itl_sum_d / itl_count_d + out[f"{prefix}itl_seconds_avg"] = itl_sum_d / itl_count_d + + stats = WindowStatistics.between(prev, cur, throughput_window_s) + out.update(stats.latency_metrics(prefix)) + preemptions = delta("ray_vllm_num_preemptions_total") + if preemptions is not None: + out[prefix + "num_preemptions"] = preemptions + if gen_d is not None and gen_d > 0: + out[prefix + "preemptions_per_million_tokens"] = preemptions * 1e6 / gen_d + external_q = delta("ray_vllm_external_prefix_cache_queries_total") + external_h = delta("ray_vllm_external_prefix_cache_hits_total") + if external_q is not None and external_q > 0 and external_h is not None: + out[prefix + "external_prefix_cache_hit_rate"] = external_h / external_q + for direction in ("store", "load"): + byte_delta = delta(f"ray_vllm_kv_offload_{direction}_bytes_total") + if byte_delta is not None and has_window: + out[f"{prefix}kv_offload_{direction}_throughput_bytes_s"] = byte_delta / throughput_window_s # Speculative-decoding (MTP draft) acceptance. Counters, so pure deltas over the window -- # no throughput denominator needed. Keys mirror the legacy metrics: raw draft/accept counts, diff --git a/skyrl/train/utils/vllm_window_statistics.py b/skyrl/train/utils/vllm_window_statistics.py new file mode 100644 index 0000000000..cf10a4bf6c --- /dev/null +++ b/skyrl/train/utils/vllm_window_statistics.py @@ -0,0 +1,79 @@ +"""Internal counter and histogram statistics for a vLLM observation window.""" + +import math +from dataclasses import dataclass, field +from typing import Dict, Optional + + +def histogram_quantile(q: float, buckets: Dict[float, float]) -> Optional[float]: + """Estimate a quantile from cumulative classic-histogram bucket counts.""" + bounds = sorted(buckets) + if len(bounds) < 2 or bounds[-1] != math.inf or buckets[math.inf] <= 0: + return None + counts = [buckets[b] for b in bounds] + if any(not math.isfinite(c) or c < 0 for c in counts): + return None + if any(a > b for a, b in zip(counts, counts[1:])): + return None + rank = q * counts[-1] + lower, previous = 0.0, 0.0 + for i, upper in enumerate(bounds): + count = counts[i] + if count >= rank: + if upper == math.inf: + return bounds[i - 1] + if i == 0 and upper <= 0: + return upper + return lower + (upper - lower) * (rank - previous) / (count - previous) if count > previous else upper + lower, previous = upper, count + return None + + +@dataclass +class WindowStatistics: + """Keep denominators and histogram buckets out of the public scalar payload.""" + + duration_seconds: float = 0.0 + deltas: Dict[str, float] = field(default_factory=dict) + valid: bool = True + + @classmethod + def between(cls, previous, current, duration): + """Difference cumulative snapshots; omit missing or decreasing counters.""" + result = cls(duration_seconds=max(duration or 0.0, 0.0), valid=previous is not None and current is not None) + if not result.valid: + return result + for name, value in current.items(): + if not name.endswith(("_total", "_sum", "_count")) and "::" not in name: + continue + if name not in previous: + continue + delta = value - previous[name] + if not math.isfinite(delta) or delta < 0: + result.valid = False + continue + result.deltas[name] = delta + return result + + def latency_metrics(self, prefix): + """Reduce merged histogram deltas to latency means and quantiles.""" + out = {} + for exported, public in ( + ("time_to_first_token_seconds", "ttft_seconds"), + ("request_time_per_output_token_seconds", "request_tpot_seconds"), + ): + base = f"ray_vllm_{exported}" + count = self.deltas.get(base + "_count", 0) + total = self.deltas.get(base + "_sum") + if count > 0 and total is not None: + out[prefix + public + "_avg"] = total / count + buckets = { + float(name.split("::", 1)[1]): value + for name, value in self.deltas.items() + if name.startswith(base + "_bucket::") + } + for q in (0.5, 0.9): + value = histogram_quantile(q, buckets) + if value is not None: + out[f"{prefix}{public}_p{int(q * 100)}"] = value + return out diff --git a/tests/train/test_vllm_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index 4aaf4be6af..606d523ded 100644 --- a/tests/train/test_vllm_metrics_scraper.py +++ b/tests/train/test_vllm_metrics_scraper.py @@ -677,3 +677,37 @@ def test_discover_ray_metrics_urls_filters_dead_and_missing(monkeypatch): "http://10.0.0.1:51001/metrics", "http://10.0.0.4:51004/metrics", ] + + +@pytest.mark.asyncio +async def test_worker_filter_and_merged_latency_buckets(): + from unittest.mock import AsyncMock + + scraper = VLLMMetricsScraper(urls=["http://test/metrics"], worker_ids=["ours"]) + text = """ray_vllm_generation_tokens_total{WorkerId="ours"} 10 +ray_vllm_generation_tokens_total{WorkerId="other"} 900 +ray_vllm_time_to_first_token_seconds_bucket{WorkerId="ours",le="1"} 2 +ray_vllm_time_to_first_token_seconds_bucket{WorkerId="ours",le="+Inf"} 3 +ray_vllm_time_to_first_token_seconds_bucket{WorkerId="other",le="1"} 900 +""" + scraper._fetch_all = AsyncMock(return_value=parse_metrics_text(text)) + snapshot = await scraper._read_snapshot() + assert snapshot["ray_vllm_generation_tokens_total"] == 10 + assert snapshot["ray_vllm_time_to_first_token_seconds_bucket::1"] == 2 + + +def test_offload_and_preemption_scalar_reductions(): + current = { + "ray_vllm_num_preemptions_total": 2, + "ray_vllm_generation_tokens_total": 100, + "ray_vllm_external_prefix_cache_queries_total": 30, + "ray_vllm_external_prefix_cache_hits_total": 12, + "ray_vllm_kv_offload_store_bytes_total": 1000, + "ray_vllm_kv_offload_load_bytes_total": 500, + } + metrics = VLLMMetricsScraper._derive(current, dict.fromkeys(current, 0), 5, "vllm/") + assert metrics["vllm/num_preemptions"] == 2 + assert metrics["vllm/preemptions_per_million_tokens"] == 20000 + assert metrics["vllm/external_prefix_cache_hit_rate"] == 0.4 + assert metrics["vllm/kv_offload_store_throughput_bytes_s"] == 200 + assert metrics["vllm/kv_offload_load_throughput_bytes_s"] == 100 diff --git a/tests/train/test_vllm_window_statistics.py b/tests/train/test_vllm_window_statistics.py new file mode 100644 index 0000000000..c782667d96 --- /dev/null +++ b/tests/train/test_vllm_window_statistics.py @@ -0,0 +1,51 @@ +"""Tests for vLLM counter windows and merged latency histograms.""" + +import math + +import pytest + +from skyrl.train.utils.vllm_window_statistics import ( + WindowStatistics, + histogram_quantile, +) + + +def test_histogram_quantiles(): + buckets = {1.0: 2, 2.0: 8, math.inf: 10} + assert histogram_quantile(0.5, buckets) == pytest.approx(1.5) + assert histogram_quantile(0.9, buckets) == 2 + assert histogram_quantile(0.5, {1: 0, math.inf: 0}) is None + assert histogram_quantile(0.5, {1: 3, 2: 2, math.inf: 4}) is None + + +def test_window_omits_missing_and_reset_counters(): + result = WindowStatistics.between( + {"tokens_total": 10, "bad_total": 20}, {"tokens_total": 25, "bad_total": 1, "new_total": 4}, 3 + ) + assert result.deltas == {"tokens_total": 15} + assert not result.valid + assert result.duration_seconds == 3 + + +def test_request_tpot_and_ttft_are_separate(): + current = {} + for name, count, total in [ + ("time_to_first_token_seconds", 10, 12), + ("request_time_per_output_token_seconds", 4, 2), + ]: + base = "ray_vllm_" + name + current.update( + { + base + "_count": count, + base + "_sum": total, + base + "_bucket::1": count / 2, + base + "_bucket::2": count, + base + "_bucket::+Inf": count, + } + ) + result = WindowStatistics.between(dict.fromkeys(current, 0), current, 5) + metrics = result.latency_metrics("vllm/") + assert metrics["vllm/ttft_seconds_avg"] == 1.2 + assert metrics["vllm/request_tpot_seconds_avg"] == 0.5 + assert metrics["vllm/request_tpot_seconds_p50"] == 1 + assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(1.8) From a1ebe273e77ce4ab039d1913c22df5a883fbb0a1 Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 30 Sep 2026 18:02:46 +0000 Subject: [PATCH 2/6] fix: include first observations from empty latency histograms Signed-off-by: Codex Signed-off-by: SumanthRH --- skyrl/train/utils/vllm_window_statistics.py | 9 +++++++-- tests/train/test_vllm_window_statistics.py | 18 ++++++++++++++++++ 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/skyrl/train/utils/vllm_window_statistics.py b/skyrl/train/utils/vllm_window_statistics.py index cf10a4bf6c..37569473c1 100644 --- a/skyrl/train/utils/vllm_window_statistics.py +++ b/skyrl/train/utils/vllm_window_statistics.py @@ -46,9 +46,14 @@ def between(cls, previous, current, duration): for name, value in current.items(): if not name.endswith(("_total", "_sum", "_count")) and "::" not in name: continue - if name not in previous: + baseline = previous.get(name) + if baseline is None and "_bucket::" in name: + count = name.split("_bucket::", 1)[0] + "_count" + if previous.get(count) == 0: + baseline = 0.0 + if baseline is None: continue - delta = value - previous[name] + delta = value - baseline if not math.isfinite(delta) or delta < 0: result.valid = False continue diff --git a/tests/train/test_vllm_window_statistics.py b/tests/train/test_vllm_window_statistics.py index c782667d96..d9a53be46f 100644 --- a/tests/train/test_vllm_window_statistics.py +++ b/tests/train/test_vllm_window_statistics.py @@ -49,3 +49,21 @@ def test_request_tpot_and_ttft_are_separate(): assert metrics["vllm/request_tpot_seconds_avg"] == 0.5 assert metrics["vllm/request_tpot_seconds_p50"] == 1 assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(1.8) + + +def test_lazy_histogram_buckets_use_confirmed_empty_baseline(): + base = "ray_vllm_time_to_first_token_seconds" + previous = {base + "_count": 0, base + "_sum": 0} + current = { + base + "_count": 2, + base + "_sum": 0.15, + base + "_bucket::0.1": 1, + base + "_bucket::0.2": 2, + base + "_bucket::+Inf": 2, + } + metrics = WindowStatistics.between(previous, current, 1).latency_metrics("vllm/") + assert metrics["vllm/ttft_seconds_avg"] == pytest.approx(0.075) + assert metrics["vllm/ttft_seconds_p50"] == pytest.approx(0.1) + assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(0.18) + # Absence without an explicit empty count is unknown. + assert not WindowStatistics.between({}, current, 1).latency_metrics("vllm/") From 314eac2597ebaec20c36773f9c08a8a75ca2d58b Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 30 Sep 2026 18:54:26 +0000 Subject: [PATCH 3/6] Use request TPOT and reduce tracker latency metrics Signed-off-by: Codex Signed-off-by: SumanthRH --- skyrl/train/utils/vllm_metrics_scraper.py | 26 ++++++++---------- tests/train/test_vllm_metrics_scraper.py | 33 ++++++++++++++++++++--- 2 files changed, 41 insertions(+), 18 deletions(-) diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 5c65e5b102..8a64839aae 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -36,8 +36,6 @@ _COUNTER_GENERATION_TOKENS = "ray_vllm_generation_tokens_total" _HIST_TTFT_SUM = "ray_vllm_time_to_first_token_seconds_sum" _HIST_TTFT_COUNT = "ray_vllm_time_to_first_token_seconds_count" -_HIST_ITL_SUM = "ray_vllm_inter_token_latency_seconds_sum" -_HIST_ITL_COUNT = "ray_vllm_inter_token_latency_seconds_count" # Speculative-decoding (MTP draft) counters. The per-position counter additionally carries a # `position` label ("0".."k-1"); it is summed per-position in `sum_by_position` rather than through # `_SUM_METRICS` (which would collapse the label and lose the per-depth breakdown). @@ -55,8 +53,6 @@ _COUNTER_GENERATION_TOKENS, _HIST_TTFT_SUM, _HIST_TTFT_COUNT, - _HIST_ITL_SUM, - _HIST_ITL_COUNT, _COUNTER_SPEC_DRAFTS, _COUNTER_SPEC_DRAFT_TOKENS, _COUNTER_SPEC_ACCEPTED_TOKENS, @@ -461,14 +457,15 @@ def delta(name: str) -> Optional[float]: if ttft_sum_d is not None and ttft_count_d is not None and ttft_count_d > 0: out[f"{prefix}ttft_seconds_avg"] = ttft_sum_d / ttft_count_d - itl_sum_d = delta(_HIST_ITL_SUM) - itl_count_d = delta(_HIST_ITL_COUNT) - if itl_sum_d is not None and itl_count_d is not None and itl_count_d > 0: - out[f"{prefix}tpot_seconds_avg"] = itl_sum_d / itl_count_d - out[f"{prefix}itl_seconds_avg"] = itl_sum_d / itl_count_d - stats = WindowStatistics.between(prev, cur, throughput_window_s) - out.update(stats.latency_metrics(prefix)) + latencies = stats.latency_metrics(prefix) + for suffix in ("avg", "p90"): + request_key = f"{prefix}request_tpot_seconds_{suffix}" + if request_key in latencies: + out[f"{prefix}tpot_seconds_{suffix}"] = latencies[request_key] + ttft_key = f"{prefix}ttft_seconds_{suffix}" + if ttft_key in latencies: + out[ttft_key] = latencies[ttft_key] preemptions = delta("ray_vllm_num_preemptions_total") if preemptions is not None: out[prefix + "num_preemptions"] = preemptions @@ -478,10 +475,9 @@ def delta(name: str) -> Optional[float]: external_h = delta("ray_vllm_external_prefix_cache_hits_total") if external_q is not None and external_q > 0 and external_h is not None: out[prefix + "external_prefix_cache_hit_rate"] = external_h / external_q - for direction in ("store", "load"): - byte_delta = delta(f"ray_vllm_kv_offload_{direction}_bytes_total") - if byte_delta is not None and has_window: - out[f"{prefix}kv_offload_{direction}_throughput_bytes_s"] = byte_delta / throughput_window_s + load_bytes = delta("ray_vllm_kv_offload_load_bytes_total") + if load_bytes is not None and has_window: + out[f"{prefix}kv_offload_load_throughput_bytes_s"] = load_bytes / throughput_window_s # Speculative-decoding (MTP draft) acceptance. Counters, so pure deltas over the window -- # no throughput denominator needed. Keys mirror the legacy metrics: raw draft/accept counts, diff --git a/tests/train/test_vllm_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index 606d523ded..c0e2457c8d 100644 --- a/tests/train/test_vllm_metrics_scraper.py +++ b/tests/train/test_vllm_metrics_scraper.py @@ -241,7 +241,8 @@ async def fake_fetch_all(): assert out["vllm/prompt_throughput_tok_s"] == pytest.approx(100.0) assert out["vllm/prefix_cache_hit_rate"] == pytest.approx(40.0 / 50.0) assert out["vllm/ttft_seconds_avg"] == pytest.approx(1.0 / 5.0) - assert out["vllm/tpot_seconds_avg"] == pytest.approx(0.5 / 100.0) + assert "vllm/tpot_seconds_avg" not in out + assert "vllm/itl_seconds_avg" not in out # Gauges still flow through. assert out["vllm/num_requests_running"] == pytest.approx(5) assert out["vllm/kv_cache_usage_perc"] == pytest.approx(0.35) @@ -362,7 +363,8 @@ async def fake_fetch_all(): # Time-independent derived metrics are unchanged by the window choice. assert out["vllm/prefix_cache_hit_rate"] == pytest.approx(40.0 / 50.0) assert out["vllm/ttft_seconds_avg"] == pytest.approx(1.0 / 5.0) - assert out["vllm/tpot_seconds_avg"] == pytest.approx(0.5 / 100.0) + assert "vllm/tpot_seconds_avg" not in out + assert "vllm/itl_seconds_avg" not in out @pytest.mark.asyncio @@ -709,5 +711,30 @@ def test_offload_and_preemption_scalar_reductions(): assert metrics["vllm/num_preemptions"] == 2 assert metrics["vllm/preemptions_per_million_tokens"] == 20000 assert metrics["vllm/external_prefix_cache_hit_rate"] == 0.4 - assert metrics["vllm/kv_offload_store_throughput_bytes_s"] == 200 + assert "vllm/kv_offload_store_throughput_bytes_s" not in metrics assert metrics["vllm/kv_offload_load_throughput_bytes_s"] == 100 + + +def test_logged_tpot_uses_request_histogram_and_only_mean_p90(): + current = {} + for name, count, total in [ + ("time_to_first_token_seconds", 10, 12), + ("request_time_per_output_token_seconds", 4, 2), + ("inter_token_latency_seconds", 100, 0.1), + ]: + base = "ray_vllm_" + name + current.update( + { + base + "_count": count, + base + "_sum": total, + base + "_bucket::1": count / 2, + base + "_bucket::2": count, + base + "_bucket::+Inf": count, + } + ) + metrics = VLLMMetricsScraper._derive(current, dict.fromkeys(current, 0), 5, "vllm/") + assert metrics["vllm/tpot_seconds_avg"] == pytest.approx(0.5) + assert metrics["vllm/tpot_seconds_p90"] == pytest.approx(1.8) + assert metrics["vllm/ttft_seconds_avg"] == pytest.approx(1.2) + assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(1.8) + assert not any("itl" in key or "request_tpot" in key or "p50" in key for key in metrics) From ff9197a395049ababc3156e660322abf6292e06a Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Thu, 1 Oct 2026 18:32:37 +0000 Subject: [PATCH 4/6] Trim window metrics and reject incomplete or invalid observations Signed-off-by: SumanthRH --- .../checkpointing-logging/vllm-metrics.mdx | 26 ++++++++ skyrl/train/entrypoints/main_base.py | 10 ++- skyrl/train/utils/vllm_metrics_scraper.py | 16 ++--- skyrl/train/utils/vllm_window_statistics.py | 11 ++-- tests/train/test_vllm_metrics_scraper.py | 61 ++++++++++++++++++- tests/train/test_vllm_window_statistics.py | 30 +++------ 6 files changed, 113 insertions(+), 41 deletions(-) diff --git a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx index 71862ed7df..5206b4e0d6 100644 --- a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx +++ b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx @@ -31,6 +31,11 @@ output into the wandb log payload — the same payload used for training metrics, so the keys appear under whatever logger backend is configured (`wandb`, `mlflow`, `swanlab`, `tensorboard`, or `console`). +Snapshots are filtered by the WorkerIds of inference actors launched by SkyRL. +This membership is fixed; replacements and late joins are not supported. +External deployments without launched actors retain unfiltered collection. +If launched worker identities cannot be obtained, collection is disabled. + Both `Trainer` and `FullyAsyncTrainer` log these: | Key | Source | Aggregation | @@ -41,12 +46,33 @@ Both `Trainer` and `FullyAsyncTrainer` log these: | `vllm/generation_throughput_tok_s` | counter delta / Δt | summed before differencing | | `vllm/prompt_throughput_tok_s` | counter delta / Δt | summed before differencing | | `vllm/prefix_cache_hit_rate` | hits Δ / queries Δ | summed before ratio | +| `vllm/external_prefix_cache_hit_rate` | external hits Δ / queries Δ | summed before ratio | +| `vllm/num_preemptions` | preemption-event delta | sum across replicas | +| `vllm/preemptions_per_million_tokens` | preemptions × 1e6 / output-token Δ | summed before ratio | +| `vllm/kv_offload_load_throughput_bytes_s` | load-byte Δ / Δt | summed before differencing | | `vllm/ttft_seconds_avg` | histogram sum Δ / count Δ | summed before ratio | | `vllm/tpot_seconds_avg` | histogram sum Δ / count Δ | summed before ratio | +| `vllm/ttft_seconds_p90`, `vllm/tpot_seconds_p90` | histogram bucket deltas | merged before quantile | + +TTFT uses `time_to_first_token_seconds`; TPOT uses the request-weighted +`request_time_per_output_token_seconds` histogram, matching Serve LLM. +Older runs used inter-token latency for `tpot_seconds_avg` and are not +comparable to this definition. ITL, latency P50, duplicate request-TPOT keys, +and offload store throughput are omitted from tracker output. Their raw +Prometheus metrics remain available for Grafana. + +Throughput divides token/byte deltas by the supplied generation duration; +without that duration, sampling uses elapsed wall time. Percentiles interpolate +compatible cumulative bucket deltas after merging engines. A previously +confirmed zero histogram count supplies a zero baseline for its first bucket +exports; an absent count supplies no baseline. Invalid histogram windows do +not produce latency metrics. Rate- and ratio-style metrics need two consecutive samples to take a delta, so they appear starting from the **second** training step. Counter resets (e.g. engine restart) are skipped rather than reported as negative rates. +If any node endpoint fails, the partial snapshot is skipped and sampling +reestablishes its baseline at the next complete snapshot. The full set of vLLM metrics is still available via the Prometheus endpoints themselves — only this curated subset is forwarded to wandb. The selection diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index b3bd60b76c..6533661cec 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -296,8 +296,14 @@ def _setup_trainer(self) -> RayPPOTrainer: groups = self._server_groups or ((self._prefill_server_groups or []) + (self._decode_server_groups or [])) actors = [actor for group in (groups or []) for actor in group.get_actors()] if actors and self.cfg.generator.inference_engine.backend == "vllm": - worker_ids = ray.get([actor.get_metrics_worker_id.remote() for actor in actors]) - trainer._vllm_metrics_scraper.set_worker_ids(worker_ids) + try: + worker_ids = ray.get([actor.get_metrics_worker_id.remote() for actor in actors], timeout=10) + trainer._vllm_metrics_scraper.set_worker_ids(worker_ids) + except Exception as error: + trainer._vllm_metrics_scraper.set_worker_ids([]) + logger.warning( + f"vLLM metrics disabled: could not identify launched workers ({type(error).__name__})" + ) # Install the trajectory logger after construction trainer.trajectory_logger = self.get_trajectory_logger() # Expose the trainer on self so callers can log exceptions raised diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 8a64839aae..69d3c426a0 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -205,11 +205,11 @@ 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.last_window = WindowStatistics(valid=False) self._prev_aggregated: Optional[Dict[str, float]] = None self._prev_timestamp: Optional[float] = None self._client: Optional[httpx.AsyncClient] = None self._warned_empty = False + self._snapshot_complete = True # Explicit-window state (start/pause/resume/stop). ``_label is None`` # means no window is open. self._label: Optional[str] = None @@ -249,6 +249,7 @@ async def _fetch_one(self, client: httpx.AsyncClient, url: str) -> str: async def _fetch_all(self) -> ParsedSamples: client = await self._get_client() texts = await asyncio.gather(*(self._fetch_one(client, u) for u in self._urls)) + self._snapshot_complete = all(bool(text) for text in texts) merged: ParsedSamples = {} for text in texts: if not text: @@ -268,6 +269,8 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: return None parsed = await self._fetch_all() + if not self._snapshot_complete: + return None if not parsed and not self._warned_empty: logger.warning( "VLLMMetricsScraper: scraped Ray metrics agents but found no " @@ -327,6 +330,8 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f """ snapshot = await self._read_snapshot() if snapshot is None: + self._prev_aggregated = None + self._prev_timestamp = None return {} now = time.monotonic() @@ -337,9 +342,6 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f else: out = self._window_metrics(None, snapshot, None, "vllm/") # gauges only - self.last_window = WindowStatistics.between( - self._prev_aggregated, snapshot, window if self._prev_timestamp is not None else None - ) self._prev_aggregated = snapshot self._prev_timestamp = now return out @@ -396,7 +398,6 @@ async def stop(self) -> Dict[str, float]: self._window_prev = None self._active_since = None self._paused = False - self.last_window = WindowStatistics.between(prev, new_snapshot, window) if new_snapshot is None: return {} return self._window_metrics(prev, new_snapshot, window, f"{label}/") @@ -452,11 +453,6 @@ def delta(name: str) -> Optional[float]: if q_d is not None and h_d is not None and q_d > 0: out[f"{prefix}prefix_cache_hit_rate"] = h_d / q_d - ttft_sum_d = delta(_HIST_TTFT_SUM) - ttft_count_d = delta(_HIST_TTFT_COUNT) - if ttft_sum_d is not None and ttft_count_d is not None and ttft_count_d > 0: - out[f"{prefix}ttft_seconds_avg"] = ttft_sum_d / ttft_count_d - stats = WindowStatistics.between(prev, cur, throughput_window_s) latencies = stats.latency_metrics(prefix) for suffix in ("avg", "p90"): diff --git a/skyrl/train/utils/vllm_window_statistics.py b/skyrl/train/utils/vllm_window_statistics.py index 37569473c1..0fd8b52813 100644 --- a/skyrl/train/utils/vllm_window_statistics.py +++ b/skyrl/train/utils/vllm_window_statistics.py @@ -61,7 +61,9 @@ def between(cls, previous, current, duration): return result def latency_metrics(self, prefix): - """Reduce merged histogram deltas to latency means and quantiles.""" + """Reduce merged histogram deltas to latency means and P90 estimates.""" + if not self.valid: + return {} out = {} for exported, public in ( ("time_to_first_token_seconds", "ttft_seconds"), @@ -77,8 +79,7 @@ def latency_metrics(self, prefix): for name, value in self.deltas.items() if name.startswith(base + "_bucket::") } - for q in (0.5, 0.9): - value = histogram_quantile(q, buckets) - if value is not None: - out[f"{prefix}{public}_p{int(q * 100)}"] = value + value = histogram_quantile(0.9, buckets) + if value is not None: + out[f"{prefix}{public}_p90"] = value return out diff --git a/tests/train/test_vllm_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index c0e2457c8d..520d5b0b79 100644 --- a/tests/train/test_vllm_metrics_scraper.py +++ b/tests/train/test_vllm_metrics_scraper.py @@ -3,8 +3,9 @@ """ import asyncio -from unittest.mock import patch +from unittest.mock import Mock, patch +import httpx import pytest from skyrl.train.utils.vllm_metrics_scraper import ( @@ -698,6 +699,34 @@ async def test_worker_filter_and_merged_latency_buckets(): assert snapshot["ray_vllm_time_to_first_token_seconds_bucket::1"] == 2 +@pytest.mark.asyncio +async def test_failed_node_scrape_skips_partial_totals_and_reestablishes_baseline(): + step = 0 + + def respond(request): + worker = request.url.host + if step == 1 and worker == "b": + return httpx.Response(503) + text = ( + f'ray_vllm_num_requests_running{{WorkerId="{worker}"}} 1\n' + f'ray_vllm_generation_tokens_total{{WorkerId="{worker}"}} {step * 10}\n' + ) + return httpx.Response(200, text=text) + + scraper = VLLMMetricsScraper(urls=["http://a/metrics", "http://b/metrics"], worker_ids=["a", "b"]) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + scraper._client = client + await scraper.sample(generation_time_s=1) + step = 1 + assert await scraper.sample(generation_time_s=1) == {} + step = 2 + recovered = await scraper.sample(generation_time_s=1) + assert "vllm/generation_throughput_tok_s" not in recovered + step = 3 + metrics = await scraper.sample(generation_time_s=1) + assert metrics["vllm/generation_throughput_tok_s"] == 20 + + def test_offload_and_preemption_scalar_reductions(): current = { "ray_vllm_num_preemptions_total": 2, @@ -738,3 +767,33 @@ def test_logged_tpot_uses_request_histogram_and_only_mean_p90(): assert metrics["vllm/ttft_seconds_avg"] == pytest.approx(1.2) assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(1.8) assert not any("itl" in key or "request_tpot" in key or "p50" in key for key in metrics) + + +@pytest.mark.parametrize("error", [None, RuntimeError("identity failed"), TimeoutError("identity timed out")]) +def test_worker_identity_lookup_is_bounded_and_failure_does_not_abort_setup(tmp_path, monkeypatch, error): + from skyrl.train.config import SkyRLTrainConfig + from skyrl.train.entrypoints.main_base import BasePPOExp + + cfg = SkyRLTrainConfig() + cfg.trainer.export_path = str(tmp_path / "export") + cfg.trainer.ckpt_path = str(tmp_path / "checkpoints") + cfg.trainer.fully_async.simulate_training = True + experiment = BasePPOExp.__new__(BasePPOExp) + experiment.cfg = cfg + experiment.tokenizer = Mock() + experiment.train_dataset = None + experiment.eval_dataset = None + experiment.colocate_pg = None + actor = Mock() + actor.get_metrics_worker_id.remote.return_value = "identity-ref" + experiment._server_groups = [Mock(get_actors=Mock(return_value=[actor]))] + trainer = Mock() + experiment.get_trainer = Mock(return_value=trainer) + for method in ("get_tracker", "get_inference_client", "get_generator", "get_trajectory_logger"): + setattr(experiment, method, Mock()) + lookup = Mock(return_value=["owner"], side_effect=error) + monkeypatch.setattr("skyrl.train.entrypoints.main_base.ray.get", lookup) + + assert experiment._setup_trainer() is trainer + lookup.assert_called_once_with(["identity-ref"], timeout=10) + trainer._vllm_metrics_scraper.set_worker_ids.assert_called_once_with([] if error else ["owner"]) diff --git a/tests/train/test_vllm_window_statistics.py b/tests/train/test_vllm_window_statistics.py index d9a53be46f..3b353244f5 100644 --- a/tests/train/test_vllm_window_statistics.py +++ b/tests/train/test_vllm_window_statistics.py @@ -27,28 +27,13 @@ def test_window_omits_missing_and_reset_counters(): assert result.duration_seconds == 3 -def test_request_tpot_and_ttft_are_separate(): - current = {} - for name, count, total in [ - ("time_to_first_token_seconds", 10, 12), - ("request_time_per_output_token_seconds", 4, 2), - ]: - base = "ray_vllm_" + name - current.update( - { - base + "_count": count, - base + "_sum": total, - base + "_bucket::1": count / 2, - base + "_bucket::2": count, - base + "_bucket::+Inf": count, - } - ) - result = WindowStatistics.between(dict.fromkeys(current, 0), current, 5) - metrics = result.latency_metrics("vllm/") - assert metrics["vllm/ttft_seconds_avg"] == 1.2 - assert metrics["vllm/request_tpot_seconds_avg"] == 0.5 - assert metrics["vllm/request_tpot_seconds_p50"] == 1 - assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(1.8) +def test_invalid_histogram_window_omits_latency_metrics(): + base = "ray_vllm_time_to_first_token_seconds" + previous = {base + "_count": 10, base + "_sum": 100, base + "_bucket::1": 5, base + "_bucket::+Inf": 10} + current = {base + "_count": 20, base + "_sum": 20, base + "_bucket::1": 10, base + "_bucket::+Inf": 20} + window = WindowStatistics.between(previous, current, 1) + assert not window.valid + assert window.latency_metrics("vllm/") == {} def test_lazy_histogram_buckets_use_confirmed_empty_baseline(): @@ -63,7 +48,6 @@ def test_lazy_histogram_buckets_use_confirmed_empty_baseline(): } metrics = WindowStatistics.between(previous, current, 1).latency_metrics("vllm/") assert metrics["vllm/ttft_seconds_avg"] == pytest.approx(0.075) - assert metrics["vllm/ttft_seconds_p50"] == pytest.approx(0.1) assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(0.18) # Absence without an explicit empty count is unknown. assert not WindowStatistics.between({}, current, 1).latency_metrics("vllm/") From 83978f6b81a3571f1ad270f2651617480ab44a45 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Thu, 1 Oct 2026 19:30:25 +0000 Subject: [PATCH 5/6] Simplify counter and histogram reductions Signed-off-by: SumanthRH --- skyrl/train/utils/vllm_metrics_scraper.py | 27 ++---- skyrl/train/utils/vllm_window_statistics.py | 95 +++++++++------------ tests/train/test_vllm_window_statistics.py | 19 ++--- 3 files changed, 54 insertions(+), 87 deletions(-) diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 69d3c426a0..08916a9e54 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -20,7 +20,7 @@ from loguru import logger from skyrl.backends.skyrl_train.inference_servers.common import format_http_url -from skyrl.train.utils.vllm_window_statistics import WindowStatistics +from skyrl.train.utils.vllm_window_statistics import latency_metrics # vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`) # AND the `ray_` prefix that Ray's metrics agent adds to every custom metric. @@ -56,17 +56,13 @@ _COUNTER_SPEC_DRAFTS, _COUNTER_SPEC_DRAFT_TOKENS, _COUNTER_SPEC_ACCEPTED_TOKENS, -) -_ADDITIONAL_COUNTERS = ( "ray_vllm_num_preemptions_total", "ray_vllm_external_prefix_cache_queries_total", "ray_vllm_external_prefix_cache_hits_total", - "ray_vllm_kv_offload_store_bytes_total", "ray_vllm_kv_offload_load_bytes_total", "ray_vllm_request_time_per_output_token_seconds_sum", "ray_vllm_request_time_per_output_token_seconds_count", ) -_SUM_METRICS += _ADDITIONAL_COUNTERS _MEAN_METRICS = (_GAUGE_KV_CACHE_USAGE,) ParsedSamples = Dict[Tuple[str, FrozenSet[Tuple[str, str]]], float] @@ -302,23 +298,18 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: for name, owners in schemas.items(): if len({frozenset(bounds) for bounds in owners.values()}) != 1: buckets = {key: value for key, value in buckets.items() if not key.startswith(name + "::")} - # Ray counters skip zero increments. A live engine gauge confirms the - # exporter exists; connector queries/store counters establish support. - if any(name == _GAUGE_NUM_RUNNING for name, _ in parsed): - buckets.setdefault("ray_vllm_num_preemptions_total", 0.0) sums = aggregate(parsed, _SUM_METRICS, how="sum") means = aggregate(parsed, _MEAN_METRICS, how="mean") per_pos = sum_by_position(parsed, _COUNTER_SPEC_ACCEPTED_PER_POS) - if "ray_vllm_num_preemptions_total" in sums: - buckets.pop("ray_vllm_num_preemptions_total", None) + # 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("ray_vllm_num_preemptions_total", 0.0) for hits, queries in ( ("ray_vllm_prefix_cache_hits_total", _COUNTER_PREFIX_QUERIES), ("ray_vllm_external_prefix_cache_hits_total", "ray_vllm_external_prefix_cache_queries_total"), ): if queries in sums: sums.setdefault(hits, 0.0) - if "ray_vllm_kv_offload_store_bytes_total" in sums: - sums.setdefault("ray_vllm_kv_offload_load_bytes_total", 0.0) return {**sums, **means, **per_pos, **buckets} async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, float]: @@ -453,15 +444,7 @@ def delta(name: str) -> Optional[float]: if q_d is not None and h_d is not None and q_d > 0: out[f"{prefix}prefix_cache_hit_rate"] = h_d / q_d - stats = WindowStatistics.between(prev, cur, throughput_window_s) - latencies = stats.latency_metrics(prefix) - for suffix in ("avg", "p90"): - request_key = f"{prefix}request_tpot_seconds_{suffix}" - if request_key in latencies: - out[f"{prefix}tpot_seconds_{suffix}"] = latencies[request_key] - ttft_key = f"{prefix}ttft_seconds_{suffix}" - if ttft_key in latencies: - out[ttft_key] = latencies[ttft_key] + out.update(latency_metrics(prev, cur, prefix)) preemptions = delta("ray_vllm_num_preemptions_total") if preemptions is not None: out[prefix + "num_preemptions"] = preemptions diff --git a/skyrl/train/utils/vllm_window_statistics.py b/skyrl/train/utils/vllm_window_statistics.py index 0fd8b52813..8bb2935022 100644 --- a/skyrl/train/utils/vllm_window_statistics.py +++ b/skyrl/train/utils/vllm_window_statistics.py @@ -1,7 +1,6 @@ """Internal counter and histogram statistics for a vLLM observation window.""" import math -from dataclasses import dataclass, field from typing import Dict, Optional @@ -29,57 +28,47 @@ def histogram_quantile(q: float, buckets: Dict[float, float]) -> Optional[float] return None -@dataclass -class WindowStatistics: - """Keep denominators and histogram buckets out of the public scalar payload.""" - - duration_seconds: float = 0.0 - deltas: Dict[str, float] = field(default_factory=dict) - valid: bool = True +def counter_deltas(previous, current): + """Difference observed counters; return None for missing snapshots or resets.""" + if previous is None or current is None: + return None + deltas = {} + for name, value in current.items(): + if not name.endswith(("_total", "_sum", "_count")) and "::" not in name: + continue + baseline = previous.get(name) + if baseline is None and "_bucket::" in name and previous.get(name.split("_bucket::", 1)[0] + "_count") == 0: + baseline = 0.0 + if baseline is None: + continue + delta = value - baseline + if not math.isfinite(delta) or delta < 0: + return None + deltas[name] = delta + return deltas - @classmethod - def between(cls, previous, current, duration): - """Difference cumulative snapshots; omit missing or decreasing counters.""" - result = cls(duration_seconds=max(duration or 0.0, 0.0), valid=previous is not None and current is not None) - if not result.valid: - return result - for name, value in current.items(): - if not name.endswith(("_total", "_sum", "_count")) and "::" not in name: - continue - baseline = previous.get(name) - if baseline is None and "_bucket::" in name: - count = name.split("_bucket::", 1)[0] + "_count" - if previous.get(count) == 0: - baseline = 0.0 - if baseline is None: - continue - delta = value - baseline - if not math.isfinite(delta) or delta < 0: - result.valid = False - continue - result.deltas[name] = delta - return result - def latency_metrics(self, prefix): - """Reduce merged histogram deltas to latency means and P90 estimates.""" - if not self.valid: - return {} - out = {} - for exported, public in ( - ("time_to_first_token_seconds", "ttft_seconds"), - ("request_time_per_output_token_seconds", "request_tpot_seconds"), - ): - base = f"ray_vllm_{exported}" - count = self.deltas.get(base + "_count", 0) - total = self.deltas.get(base + "_sum") - if count > 0 and total is not None: - out[prefix + public + "_avg"] = total / count - buckets = { - float(name.split("::", 1)[1]): value - for name, value in self.deltas.items() - if name.startswith(base + "_bucket::") - } - value = histogram_quantile(0.9, buckets) - if value is not None: - out[f"{prefix}{public}_p90"] = value - return out +def latency_metrics(previous, current, prefix): + """Reduce merged histogram deltas to latency means and P90 estimates.""" + deltas = counter_deltas(previous, current) + if deltas is None: + return {} + out = {} + for exported, public in ( + ("time_to_first_token_seconds", "ttft_seconds"), + ("request_time_per_output_token_seconds", "tpot_seconds"), + ): + base = f"ray_vllm_{exported}" + count = deltas.get(base + "_count", 0) + total = deltas.get(base + "_sum") + if count > 0 and total is not None: + out[prefix + public + "_avg"] = total / count + buckets = { + float(name.split("::", 1)[1]): value + for name, value in deltas.items() + if name.startswith(base + "_bucket::") + } + value = histogram_quantile(0.9, buckets) + if value is not None: + out[f"{prefix}{public}_p90"] = value + return out diff --git a/tests/train/test_vllm_window_statistics.py b/tests/train/test_vllm_window_statistics.py index 3b353244f5..2ba96ce3d9 100644 --- a/tests/train/test_vllm_window_statistics.py +++ b/tests/train/test_vllm_window_statistics.py @@ -5,8 +5,9 @@ import pytest from skyrl.train.utils.vllm_window_statistics import ( - WindowStatistics, + counter_deltas, histogram_quantile, + latency_metrics, ) @@ -19,21 +20,15 @@ def test_histogram_quantiles(): def test_window_omits_missing_and_reset_counters(): - result = WindowStatistics.between( - {"tokens_total": 10, "bad_total": 20}, {"tokens_total": 25, "bad_total": 1, "new_total": 4}, 3 - ) - assert result.deltas == {"tokens_total": 15} - assert not result.valid - assert result.duration_seconds == 3 + assert counter_deltas({"tokens_total": 10}, {"tokens_total": 25, "new_total": 4}) == {"tokens_total": 15} + assert counter_deltas({"bad_total": 20}, {"bad_total": 1}) is None def test_invalid_histogram_window_omits_latency_metrics(): base = "ray_vllm_time_to_first_token_seconds" previous = {base + "_count": 10, base + "_sum": 100, base + "_bucket::1": 5, base + "_bucket::+Inf": 10} current = {base + "_count": 20, base + "_sum": 20, base + "_bucket::1": 10, base + "_bucket::+Inf": 20} - window = WindowStatistics.between(previous, current, 1) - assert not window.valid - assert window.latency_metrics("vllm/") == {} + assert latency_metrics(previous, current, "vllm/") == {} def test_lazy_histogram_buckets_use_confirmed_empty_baseline(): @@ -46,8 +41,8 @@ def test_lazy_histogram_buckets_use_confirmed_empty_baseline(): base + "_bucket::0.2": 2, base + "_bucket::+Inf": 2, } - metrics = WindowStatistics.between(previous, current, 1).latency_metrics("vllm/") + metrics = latency_metrics(previous, current, "vllm/") assert metrics["vllm/ttft_seconds_avg"] == pytest.approx(0.075) assert metrics["vllm/ttft_seconds_p90"] == pytest.approx(0.18) # Absence without an explicit empty count is unknown. - assert not WindowStatistics.between({}, current, 1).latency_metrics("vllm/") + assert not latency_metrics({}, current, "vllm/") From 1e2064c27d98d94b46ba395b3ed36d1a0c5149e6 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Sun, 4 Oct 2026 00:47:45 +0000 Subject: [PATCH 6/6] Use metric constants and expose the Ray server worker ID Signed-off-by: SumanthRH --- .../inference_servers/vllm_server_actor.py | 4 +- skyrl/train/entrypoints/main_base.py | 2 +- skyrl/train/utils/vllm_metrics_scraper.py | 39 +++++++++++-------- tests/train/test_vllm_metrics_scraper.py | 2 +- 4 files changed, 26 insertions(+), 21 deletions(-) diff --git a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py index 7ca5a444aa..c1b2aa1812 100644 --- a/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py +++ b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py @@ -376,8 +376,8 @@ def _setup_mooncake_port(self, mooncake_server_port: int) -> None: f"host={self._ip}, port={mooncake_server_port}, engine_id={engine_id}" ) - def get_metrics_worker_id(self) -> str: - """Return the Ray worker identity used by this server's metrics logger.""" + def get_ray_worker_id(self) -> str: + """Return the Ray worker ID of the actor process hosting this API server.""" import ray return ray.get_runtime_context().get_worker_id() diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index 6533661cec..45ee325d5d 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -297,7 +297,7 @@ def _setup_trainer(self) -> RayPPOTrainer: 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_metrics_worker_id.remote() for actor in actors], timeout=10) + 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) except Exception as error: trainer._vllm_metrics_scraper.set_worker_ids([]) diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 08916a9e54..92e7ae2986 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -34,8 +34,16 @@ _COUNTER_PREFIX_HITS = "ray_vllm_prefix_cache_hits_total" _COUNTER_PROMPT_TOKENS = "ray_vllm_prompt_tokens_total" _COUNTER_GENERATION_TOKENS = "ray_vllm_generation_tokens_total" +_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_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" +_HIST_TTFT_BUCKET = "ray_vllm_time_to_first_token_seconds_bucket" +_HIST_REQUEST_TPOT_SUM = "ray_vllm_request_time_per_output_token_seconds_sum" +_HIST_REQUEST_TPOT_COUNT = "ray_vllm_request_time_per_output_token_seconds_count" +_HIST_REQUEST_TPOT_BUCKET = "ray_vllm_request_time_per_output_token_seconds_bucket" # Speculative-decoding (MTP draft) counters. The per-position counter additionally carries a # `position` label ("0".."k-1"); it is summed per-position in `sum_by_position` rather than through # `_SUM_METRICS` (which would collapse the label and lose the per-depth breakdown). @@ -56,12 +64,12 @@ _COUNTER_SPEC_DRAFTS, _COUNTER_SPEC_DRAFT_TOKENS, _COUNTER_SPEC_ACCEPTED_TOKENS, - "ray_vllm_num_preemptions_total", - "ray_vllm_external_prefix_cache_queries_total", - "ray_vllm_external_prefix_cache_hits_total", - "ray_vllm_kv_offload_load_bytes_total", - "ray_vllm_request_time_per_output_token_seconds_sum", - "ray_vllm_request_time_per_output_token_seconds_count", + _COUNTER_PREEMPTIONS, + _COUNTER_EXTERNAL_PREFIX_QUERIES, + _COUNTER_EXTERNAL_PREFIX_HITS, + _COUNTER_KV_OFFLOAD_LOAD_BYTES, + _HIST_REQUEST_TPOT_SUM, + _HIST_REQUEST_TPOT_COUNT, ) _MEAN_METRICS = (_GAUGE_KV_CACHE_USAGE,) @@ -282,10 +290,7 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: buckets = {} schemas = {} for (name, labels), value in parsed.items(): - if name not in { - "ray_vllm_time_to_first_token_seconds_bucket", - "ray_vllm_request_time_per_output_token_seconds_bucket", - }: + if name not in {_HIST_TTFT_BUCKET, _HIST_REQUEST_TPOT_BUCKET}: continue label_dict = dict(labels) bound = label_dict.get("le") @@ -303,10 +308,10 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: per_pos = sum_by_position(parsed, _COUNTER_SPEC_ACCEPTED_PER_POS) # 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("ray_vllm_num_preemptions_total", 0.0) + sums.setdefault(_COUNTER_PREEMPTIONS, 0.0) for hits, queries in ( - ("ray_vllm_prefix_cache_hits_total", _COUNTER_PREFIX_QUERIES), - ("ray_vllm_external_prefix_cache_hits_total", "ray_vllm_external_prefix_cache_queries_total"), + (_COUNTER_PREFIX_HITS, _COUNTER_PREFIX_QUERIES), + (_COUNTER_EXTERNAL_PREFIX_HITS, _COUNTER_EXTERNAL_PREFIX_QUERIES), ): if queries in sums: sums.setdefault(hits, 0.0) @@ -445,16 +450,16 @@ def delta(name: str) -> Optional[float]: out[f"{prefix}prefix_cache_hit_rate"] = h_d / q_d out.update(latency_metrics(prev, cur, prefix)) - preemptions = delta("ray_vllm_num_preemptions_total") + preemptions = delta(_COUNTER_PREEMPTIONS) if preemptions is not None: out[prefix + "num_preemptions"] = preemptions if gen_d is not None and gen_d > 0: out[prefix + "preemptions_per_million_tokens"] = preemptions * 1e6 / gen_d - external_q = delta("ray_vllm_external_prefix_cache_queries_total") - external_h = delta("ray_vllm_external_prefix_cache_hits_total") + external_q = delta(_COUNTER_EXTERNAL_PREFIX_QUERIES) + external_h = delta(_COUNTER_EXTERNAL_PREFIX_HITS) if external_q is not None and external_q > 0 and external_h is not None: out[prefix + "external_prefix_cache_hit_rate"] = external_h / external_q - load_bytes = delta("ray_vllm_kv_offload_load_bytes_total") + load_bytes = delta(_COUNTER_KV_OFFLOAD_LOAD_BYTES) if load_bytes is not None and has_window: out[f"{prefix}kv_offload_load_throughput_bytes_s"] = load_bytes / throughput_window_s diff --git a/tests/train/test_vllm_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index 520d5b0b79..9befb25189 100644 --- a/tests/train/test_vllm_metrics_scraper.py +++ b/tests/train/test_vllm_metrics_scraper.py @@ -785,7 +785,7 @@ def test_worker_identity_lookup_is_bounded_and_failure_does_not_abort_setup(tmp_ experiment.eval_dataset = None experiment.colocate_pg = None actor = Mock() - actor.get_metrics_worker_id.remote.return_value = "identity-ref" + actor.get_ray_worker_id.remote.return_value = "identity-ref" experiment._server_groups = [Mock(get_actors=Mock(return_value=[actor]))] trainer = Mock() experiment.get_trainer = Mock(return_value=trainer)