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/backends/skyrl_train/inference_servers/vllm_server_actor.py b/skyrl/backends/skyrl_train/inference_servers/vllm_server_actor.py index 525951284f..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,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_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() + 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..45ee325d5d 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -292,6 +292,18 @@ 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": + 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) + 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 9178c47715..92e7ae2986 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 latency_metrics # vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`) # AND the `ray_` prefix that Ray's metrics agent adds to every custom metric. @@ -33,10 +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_ITL_SUM = "ray_vllm_inter_token_latency_seconds_sum" -_HIST_ITL_COUNT = "ray_vllm_inter_token_latency_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). @@ -54,11 +61,15 @@ _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, + _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,) @@ -193,13 +204,16 @@ 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._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 @@ -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) @@ -235,6 +253,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: @@ -254,6 +273,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 " @@ -262,10 +283,39 @@ 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 {_HIST_TTFT_BUCKET, _HIST_REQUEST_TPOT_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 + "::")} 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} + # 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) + 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) + 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). @@ -276,6 +326,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() @@ -397,15 +449,19 @@ 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 - - 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.update(latency_metrics(prev, cur, prefix)) + 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(_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(_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 # 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..8bb2935022 --- /dev/null +++ b/skyrl/train/utils/vllm_window_statistics.py @@ -0,0 +1,74 @@ +"""Internal counter and histogram statistics for a vLLM observation window.""" + +import math +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 + + +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 + + +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_metrics_scraper.py b/tests/train/test_vllm_metrics_scraper.py index 4aaf4be6af..9befb25189 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 ( @@ -241,7 +242,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 +364,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 @@ -677,3 +680,120 @@ 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 + + +@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, + "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 "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) + + +@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_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) + 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 new file mode 100644 index 0000000000..2ba96ce3d9 --- /dev/null +++ b/tests/train/test_vllm_window_statistics.py @@ -0,0 +1,48 @@ +"""Tests for vLLM counter windows and merged latency histograms.""" + +import math + +import pytest + +from skyrl.train.utils.vllm_window_statistics import ( + counter_deltas, + histogram_quantile, + latency_metrics, +) + + +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(): + 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} + assert latency_metrics(previous, current, "vllm/") == {} + + +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 = 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 latency_metrics({}, current, "vllm/")