Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions docs/content/docs/checkpointing-logging/vllm-metrics.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -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 |
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The ray.get_runtime_context() object does not have a get_worker_id() method. Calling it will raise an AttributeError, which will silently disable vLLM metrics collection during setup. Instead, use the worker_id property and convert it to a hex string.

Suggested change
return ray.get_runtime_context().get_worker_id()
return ray.get_runtime_context().worker_id.hex()


def get_server_info(self) -> ServerInfo:
"""Get the server's IP and port info."""
return ServerInfo(
Expand Down
12 changes: 12 additions & 0 deletions skyrl/train/entrypoints/main_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
84 changes: 70 additions & 14 deletions skyrl/train/utils/vllm_metrics_scraper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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).
Expand All @@ -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,)

Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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:
Expand All @@ -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 "
Expand All @@ -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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Converting key[1] (which is a frozenset of label tuples) to a dictionary via dict(key[1]) for every single metric sample on every step is inefficient. We can perform a direct generator check over the frozenset tuples to avoid allocating a new dictionary for every key.

Suggested change
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 self._worker_ids is not None:
parsed = {
key: value
for key, value in parsed.items()
if any(k == "WorkerId" and v in self._worker_ids for k, v in key[1])
}

if not parsed:
return None
Comment on lines +286 to +289

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Filtered samples disappear silently If a successful scrape contains samples but none match the fixed actor membership, this filter returns no snapshot without a warning: the empty-sample warning runs before filtering. The run then loses vLLM telemetry, and operators have no diagnostic pointing to the membership mismatch.

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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Creating a dictionary dict(labels) and then iterating over labels again to construct owner is inefficient as it performs multiple passes over the labels. We can extract bound and construct the owner pairs in a single pass over labels.

            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).
Expand All @@ -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()
Expand Down Expand Up @@ -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,
Expand Down
74 changes: 74 additions & 0 deletions skyrl/train/utils/vllm_window_statistics.py
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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Returning None when any single counter resets invalidates all latency metrics (both TTFT and TPOT), even if the reset occurred in an unrelated counter (like speculative decoding or prefix cache queries). Skipping the resetting counter using continue is much more robust, as the downstream latency_metrics and histogram_quantile functions are already designed to handle missing or incomplete counters safely.

Suggested change
if not math.isfinite(delta) or delta < 0:
return None
if not math.isfinite(delta) or delta < 0:
continue

Comment thread
SumanthRH marked this conversation as resolved.
deltas[name] = delta
Comment thread
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
Loading
Loading