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: 25 additions & 1 deletion skyrl/train/entrypoints/main_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,6 +277,7 @@ def _setup_trainer(self) -> RayPPOTrainer:
# NOTE (sumanthrh): Instantiate tracker before trainer init.
# We have custom validation before this step to give better error messages.
tracker = self.get_tracker()
self.tracker = tracker

inference_engine_client = self.get_inference_client()

Expand Down Expand Up @@ -327,10 +328,25 @@ def _setup_trainer(self) -> RayPPOTrainer:

def run(self):
self.trainer = None
self.tracker = None
status = "failed"
try:
trainer = self._setup_trainer()

# Start the training loop
asyncio.run(trainer.train())
async def train_and_finalize():
run_status = "failed"
try:
await trainer.train()
run_status = "success"
finally:
try:
await trainer.finalize_metrics(run_status)
except Exception as finalization_error:
logger.warning(f"Could not finalize run metrics: {finalization_error}")

asyncio.run(train_and_finalize())
status = "success"
except Exception as e:
# OOMs raised inside actor init (e.g. FSDPPolicyWorkerBase.init_model)
# surface here as RayTaskError. Without this they only land in Ray
Expand All @@ -345,6 +361,14 @@ def run(self):
else:
logger.error(f"Setup failed before tracker was initialized:\n{e}")
raise
finally:
if self.tracker is not None:
try:
self.tracker.run_status = status
self.tracker.update_summary({"run_status": status})
self.tracker.finish(exit_code=0 if status == "success" else 1)
except Exception as finalization_error:
logger.warning(f"Could not finish run tracking: {finalization_error}")


@ray.remote(num_cpus=1)
Expand Down
6 changes: 4 additions & 2 deletions skyrl/train/fully_async_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -482,6 +482,9 @@ async def train(self):
if self._ray_gpu_monitor is not None:
self._ray_gpu_monitor.start()

if self._vllm_metrics_scraper is not None:
await self._vllm_metrics_scraper.sample()

# Eval before training
if self.cfg.trainer.eval_interval > 0 and self.cfg.trainer.eval_before_train:
with self._phase_gauge.timed_phase("eval", self.all_timings):
Expand Down Expand Up @@ -731,8 +734,7 @@ async def _watch_generators_done(tasks=generator_tasks, event=all_generators_don
if self.has_critic:
self.dispatch.finalize_pending_saves("critic")

if self._vllm_metrics_scraper is not None:
await self._vllm_metrics_scraper.aclose()
await self.finalize_metrics("success")
self.tracker.finish()
logger.info("Training done!")

Expand Down
25 changes: 22 additions & 3 deletions skyrl/train/trainer.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
import math
import os
import shutil
Expand Down Expand Up @@ -144,6 +145,9 @@ def __init__(
VLLMMetricsScraper() if cfg.generator.inference_engine.enable_ray_prometheus_stats else None
)

self._metrics_finalized = False
self._resumed_from_checkpoint = False

self._ray_gpu_monitor = RayGpuMonitor() if cfg.trainer.enable_ray_gpu_monitor else None

# trajectory logger is installed after construction if needed
Expand Down Expand Up @@ -232,6 +236,22 @@ def _build_train_dataloader_and_compute_training_steps(self):
if self.cfg.trainer.max_training_steps is not None:
self.total_training_steps = min(self.total_training_steps, self.cfg.trainer.max_training_steps)

async def finalize_metrics(self, status: str) -> None:
"""Finalize observations before the tracker closes, including failed runs."""
if self._metrics_finalized:
return
self._metrics_finalized = True
self.tracker.run_status = status
try:
if self._vllm_metrics_scraper is not None:
summary = await asyncio.wait_for(self._vllm_metrics_scraper.finalize(), timeout=10)
if not self._resumed_from_checkpoint and not self.cfg.generator.inference_engine.enable_pd:
self.tracker.update_summary(summary)
except Exception as e:
logger.warning(f"Could not finalize vLLM metrics: {e}")
finally:
self.tracker.update_summary({"run_status": status})

@torch.no_grad()
async def eval(self, vllm_metrics_scraper: Optional[VLLMMetricsScraper] = None) -> Dict[str, float]:
"""
Expand Down Expand Up @@ -598,9 +618,7 @@ async def train(self):
if self.has_critic:
self.dispatch.finalize_pending_saves("critic")

if self._vllm_metrics_scraper is not None:
await self._vllm_metrics_scraper.aclose()

await self.finalize_metrics("success")
if self._ray_gpu_monitor is not None:
self._ray_gpu_monitor.stop()

Expand Down Expand Up @@ -1819,6 +1837,7 @@ def load_checkpoints(self) -> Tuple[int, str]:
# 1. Load and validate trainer state
with io.open_file(trainer_state_path, "rb") as f:
trainer_state = torch.load(f, map_location="cpu", weights_only=False)
self._resumed_from_checkpoint = True
saved_global_step = trainer_state.get("global_step", global_step)
logger.info("Successfully loaded trainer state")
if saved_global_step != global_step:
Expand Down
37 changes: 34 additions & 3 deletions skyrl/train/utils/tracking.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,22 +75,52 @@ def __init__(
self.logger = ConsoleLogger()

self._exception_logged = False
self._finished = False
self.run_status = "running"
self._vllm_history_metrics = set()

def update_summary(self, data) -> None:
"""Write run summaries without adding another training history step."""
try:
if self.backend == "wandb" and not self._finished and self.logger.run is not None:
self.logger.run.summary.update(data)
except Exception as e:
logger.warning(f"Could not update run summary: {e}")

def log(self, data, step, commit=False):
if self.backend == "wandb":
for name in data:
if name.startswith("vllm/") and name not in self._vllm_history_metrics:
self.logger.define_metric(name, summary="none")
self._vllm_history_metrics.add(name)
self.logger.log(data=data, step=step, commit=commit)
else:
self.logger.log(data=data, step=step)

def finish(self):
def finish(self, exit_code: int = 0):
if self._finished:
return
if self.run_status == "running":
self.run_status = "success" if exit_code == 0 else "failed"
if self.backend == "wandb" and self.logger.run is not None:
# Explicit removals also clear summaries already inferred from history.
for name in self._vllm_history_metrics:
try:
del self.logger.run.summary[name]
except KeyError:
pass
except Exception as e:
logger.warning(f"Could not remove automatic metric summary {name}: {e}")
self.update_summary({"run_status": self.run_status})
self._finished = True
if self.backend == "console":
return
# NOTE (sumanthrh): We use a try-except block here while finishing tracking.
# This is because wandb often errors out with a BrokenPipeError when closing.
# https://github.com/wandb/wandb/issues/6449
try:
if self.backend == "wandb":
self.logger.finish(exit_code=0)
self.logger.finish(exit_code=exit_code)
else:
self.logger.finish()
except Exception as e:
Expand All @@ -113,6 +143,7 @@ def log_exception(self, e: BaseException, step: int = 0) -> None:
if self._exception_logged:
return
self._exception_logged = True
self.run_status = "failed"
tb_str = traceback.format_exc()[-10000:]
logger.error(f"Training failed at step {step} with {type(e).__name__}:\n{tb_str}")
if self.backend == "wandb":
Expand All @@ -128,7 +159,7 @@ def log_exception(self, e: BaseException, step: int = 0) -> None:
# Tables upload asynchronously. Finish the run so the upload
# completes before the caller re-raises and the process dies.
try:
self.finish()
self.finish(exit_code=1)
except Exception as finish_exc:
logger.warning(f"tracker.finish() raised after logging exception: {finish_exc}")
except Exception as log_exc:
Expand Down
24 changes: 24 additions & 0 deletions skyrl/train/utils/vllm_metrics_scraper.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from loguru import logger

from skyrl.backends.skyrl_train.inference_servers.common import format_http_url
from skyrl.train.utils.vllm_run_statistics import RunStatistics
from skyrl.train.utils.vllm_window_statistics import latency_metrics

# vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`)
Expand All @@ -39,6 +40,7 @@
_COUNTER_PREEMPTIONS = "ray_vllm_num_preemptions_total"
_COUNTER_EXTERNAL_PREFIX_QUERIES = "ray_vllm_external_prefix_cache_queries_total"
_COUNTER_EXTERNAL_PREFIX_HITS = "ray_vllm_external_prefix_cache_hits_total"
_COUNTER_KV_OFFLOAD_STORE_BYTES = "ray_vllm_kv_offload_store_bytes_total"
_COUNTER_KV_OFFLOAD_LOAD_BYTES = "ray_vllm_kv_offload_load_bytes_total"
_HIST_TTFT_SUM = "ray_vllm_time_to_first_token_seconds_sum"
_HIST_TTFT_COUNT = "ray_vllm_time_to_first_token_seconds_count"
Expand Down Expand Up @@ -69,6 +71,7 @@
_COUNTER_PREEMPTIONS,
_COUNTER_EXTERNAL_PREFIX_QUERIES,
_COUNTER_EXTERNAL_PREFIX_HITS,
_COUNTER_KV_OFFLOAD_STORE_BYTES,
_COUNTER_KV_OFFLOAD_LOAD_BYTES,
_HIST_REQUEST_TPOT_SUM,
_HIST_REQUEST_TPOT_COUNT,
Expand Down Expand Up @@ -211,6 +214,7 @@ def __init__(
self._urls = urls if urls is not None else discover_ray_metrics_urls()
self._timeout = request_timeout_s
self._worker_ids = None if worker_ids is None else frozenset(worker_ids)
self.run_statistics = RunStatistics()
self._prev_aggregated: Optional[Dict[str, float]] = None
self._prev_timestamp: Optional[float] = None
self._client: Optional[httpx.AsyncClient] = None
Expand Down Expand Up @@ -306,6 +310,8 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]:
counters = engine_counters.setdefault(engine, {})
counters[name] = counters.get(name, 0.0) + value
self._engine_snapshot = engine_counters
if self._worker_ids is not None and {engine[0] for engine in engine_counters} != self._worker_ids:
return None
buckets = {}
schemas = {}
for (name, labels), value in parsed.items():
Expand All @@ -328,14 +334,29 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]:
# Ray counters skip zero increments. A live engine gauge confirms the exporter exists.
if any(name == _GAUGE_NUM_RUNNING for name, _ in parsed):
sums.setdefault(_COUNTER_PREEMPTIONS, 0.0)
sums.setdefault(_COUNTER_PROMPT_TOKENS, 0.0)
sums.setdefault(_COUNTER_GENERATION_TOKENS, 0.0)
for hits, queries in (
(_COUNTER_PREFIX_HITS, _COUNTER_PREFIX_QUERIES),
(_COUNTER_EXTERNAL_PREFIX_HITS, _COUNTER_EXTERNAL_PREFIX_QUERIES),
):
if queries in sums:
sums.setdefault(hits, 0.0)
if _COUNTER_KV_OFFLOAD_STORE_BYTES in sums:
sums.setdefault(_COUNTER_KV_OFFLOAD_LOAD_BYTES, 0.0)
return {**sums, **means, **per_pos, **buckets}

async def finalize(self) -> Dict[str, float]:
"""Attempt one final collection and close the HTTP client."""
try:
if self._label is not None:
await self.stop()
else:
await self.sample()
return self.run_statistics.summary()
Comment thread
cursor[bot] marked this conversation as resolved.
finally:
await self.aclose()

async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, float]:
"""Return ``vllm/...`` scalars for the current step (empty if unavailable).

Expand All @@ -345,6 +366,7 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f
"""
snapshot = await self._read_snapshot()
if snapshot is None:
self.run_statistics.incomplete.add("combined")
self._prev_aggregated = None
self._prev_timestamp = None
self._previous_engines = {}
Expand All @@ -355,6 +377,7 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f
dt = max(now - self._prev_timestamp, 1e-9)
window = generation_time_s if (generation_time_s is not None and generation_time_s > 0) else dt
out = self._window_metrics(self._prev_aggregated, snapshot, window, "vllm/")
self.run_statistics.add("combined", self._prev_aggregated, snapshot, window)
else:
out = self._window_metrics(None, snapshot, None, "vllm/") # gauges only

Expand Down Expand Up @@ -418,6 +441,7 @@ async def stop(self) -> Dict[str, float]:
self._window_prev = None
self._active_since = None
self._paused = False
self.run_statistics.add(label.removeprefix("vllm/"), prev, new_snapshot, window)

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.

P1 Baseline evaluation missing from totals When eval_before_train is enabled, the synchronous trainer runs its initial evaluation without opening a vllm/eval window. This new aggregate records only closed windows, so vllm_correct_aggregate/eval/* excludes that evaluation's requests and tokens even though it is presented as a run-level eval total. Include the initial evaluation in the eval window accounting.

if new_snapshot is None:
return {}
result = self._window_metrics(prev, new_snapshot, window, f"{label}/")
Expand Down
55 changes: 55 additions & 0 deletions skyrl/train/utils/vllm_run_statistics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
"""Weighted summaries of observed, non-overlapping vLLM windows."""

from skyrl.train.utils.vllm_window_statistics import counter_deltas


class RunStatistics:
"""Sum raw counters and durations, keeping sync train and eval separate."""

def __init__(self):
self.windows = {}
self.seconds = {}
self.incomplete = set()

def add(self, scope, previous, current, duration):
"""Omit a scope with a missing/reset window; never average step rates."""
deltas = counter_deltas(previous, current)
if deltas is None:
self.incomplete.add(scope)
return
if scope not in self.windows:
self.windows[scope] = deltas
else:
total = self.windows[scope]
for name in deltas:
if "_bucket::" in name and total.get(name.split("_bucket::", 1)[0] + "_count") == 0:
total.setdefault(name, 0.0)
# Only counters observed over every window have full-run denominators.
self.windows[scope] = {name: value + deltas[name] for name, value in total.items() if name in deltas}
Comment thread
cursor[bot] marked this conversation as resolved.
self.seconds[scope] = self.seconds.get(scope, 0.0) + duration
Comment thread
cursor[bot] marked this conversation as resolved.

def summary(self):
"""Derive rates, request-weighted means and merged histogram P90 values."""
from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper

result = {}
for scope, deltas in self.windows.items():
if scope in self.incomplete:
continue
prefix = f"vllm_correct_aggregate/{scope}/"
metrics = VLLMMetricsScraper._derive(deltas, dict.fromkeys(deltas, 0), self.seconds[scope], prefix)
result.update(
{key: value for key, value in metrics.items() if "draft_num_" not in key and "_pos_" not in key}
)
result[prefix + "measurement_seconds"] = self.seconds[scope]
for counter, public in (
("generation_tokens", "output_tokens_total"),
("prompt_tokens", "prompt_tokens_total"),
("num_preemptions", "preemptions_total"),
("kv_offload_store_bytes", "kv_offload_store_bytes_total"),
("kv_offload_load_bytes", "kv_offload_load_bytes_total"),
):
value = deltas.get(f"ray_vllm_{counter}_total")
if value is not None:
result[prefix + public] = value
return result
Loading
Loading