-
Notifications
You must be signed in to change notification settings - Fork 448
[metrics 1/5] Add scoped vLLM offload and latency window metrics #2364
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
49686b2
a1ebe27
314eac2
ff9197a
83978f6
1e2064c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -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} | ||||||||||||||||||
|
Comment on lines
+286
to
+287
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Converting
Suggested change
|
||||||||||||||||||
| if not parsed: | ||||||||||||||||||
| return None | ||||||||||||||||||
|
Comment on lines
+286
to
+289
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
Knowledge Base Used: Trainer execution and evaluation |
||||||||||||||||||
| 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") | ||||||||||||||||||
|
Comment on lines
+295
to
+299
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Creating a dictionary bound = None
owner_pairs = []
for k, v in labels:
if k == "le":
bound = v
else:
owner_pairs.append((k, v))
if bound is None:
continue
owner = frozenset(owner_pairs) |
||||||||||||||||||
| 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, | ||||||||||||||||||
|
|
||||||||||||||||||
| Original file line number | Diff line number | Diff line change | ||||||||
|---|---|---|---|---|---|---|---|---|---|---|
| @@ -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 | ||||||||||
|
Comment on lines
+45
to
+46
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Returning
Suggested change
SumanthRH marked this conversation as resolved.
|
||||||||||
| deltas[name] = delta | ||||||||||
|
SumanthRH marked this conversation as resolved.
|
||||||||||
| 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 | ||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The
ray.get_runtime_context()object does not have aget_worker_id()method. Calling it will raise anAttributeError, which will silently disable vLLM metrics collection during setup. Instead, use theworker_idproperty and convert it to a hex string.