Skip to content
Closed
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
12 changes: 9 additions & 3 deletions skyrl/train/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,7 @@ async def evaluate(
trajectory_logger: Optional[TrajectoryLogger] = None,
tracker: Optional["Tracking"] = None,
vllm_metrics_scraper: Optional["VLLMMetricsScraper"] = None,
generation_activity=None,
) -> Dict[str, float]:
"""Runs generation and evaluation of trajectories.

Expand Down Expand Up @@ -200,9 +201,14 @@ async def evaluate(
gen_start = time.monotonic()
if vllm_metrics_scraper is not None:
vllm_metrics_scraper.resume()
generator_output: GeneratorOutput = await generator.generate(generator_input)
if vllm_metrics_scraper is not None:
vllm_metrics_scraper.pause()
from contextlib import nullcontext

try:
with generation_activity() if generation_activity is not None else nullcontext():
generator_output: GeneratorOutput = await generator.generate(generator_input)
finally:
if vllm_metrics_scraper is not None:
vllm_metrics_scraper.pause()
eval_generate_time += time.monotonic() - gen_start
validate_generator_output(len(generator_input["prompts"]), generator_output, step_wise=step_wise)
generator_outputs.append(generator_output)
Expand Down
24 changes: 15 additions & 9 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_active()

# 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 @@ -642,7 +645,7 @@ async def _watch_generators_done(tasks=generator_tasks, event=all_generators_don
if self._ray_gpu_monitor is not None:
timing_payload.update(self._ray_gpu_monitor.flush())
if self._vllm_metrics_scraper is not None:
timing_payload.update(await self._vllm_metrics_scraper.sample())
timing_payload.update(await self._vllm_metrics_scraper.sample_active())
self.tracker.log(timing_payload, step=self.global_step, commit=True)
self.all_timings = {}
self.global_step += 1
Expand Down Expand Up @@ -732,6 +735,8 @@ async def _watch_generators_done(tasks=generator_tasks, event=all_generators_don
self.dispatch.finalize_pending_saves("critic")

if self._vllm_metrics_scraper is not None:
final_metrics = await self._vllm_metrics_scraper.sample_active()
self.tracker.log(final_metrics, step=self.global_step, commit=False)
await self._vllm_metrics_scraper.aclose()
self.tracker.finish()
logger.info("Training done!")
Expand Down Expand Up @@ -920,14 +925,15 @@ async def _run_generate_for_a_group_loop(self, generation_output_group_buffer: a
global_step_at_start = self.global_step # for staleness control

group_start_time = time.monotonic()
if "disable_tqdm" in inspect.signature(self.generator.generate).parameters:
# A workaround to disable tqdm for the SkyRLGymGenerator.generate method which will
# blast the console with each worker's progress bar.
cur_generator_output: GeneratorOutput = await self.generator.generate(
generator_input, disable_tqdm=True
)
else:
cur_generator_output: GeneratorOutput = await self.generator.generate(generator_input)
with self._generation_activity():
if "disable_tqdm" in inspect.signature(self.generator.generate).parameters:
# A workaround to disable tqdm for the SkyRLGymGenerator.generate method which will
# blast the console with each worker's progress bar.
cur_generator_output: GeneratorOutput = await self.generator.generate(
generator_input, disable_tqdm=True
)
else:
cur_generator_output: GeneratorOutput = await self.generator.generate(generator_input)
group_completion_time_s = time.monotonic() - group_start_time

# 4. Enqueue the completed group and mark accepted to free capacity slot.
Expand Down
21 changes: 17 additions & 4 deletions skyrl/train/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,9 @@ def __init__(
VLLMMetricsScraper() if cfg.generator.inference_engine.enable_ray_prometheus_stats else None
)

if self._vllm_metrics_scraper is not None:
self._vllm_metrics_scraper.publish_activity(cfg.trainer.run_name)

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 +235,12 @@ 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)

def _generation_activity(self):
"""Return an activity context when vLLM metrics are enabled."""
from contextlib import nullcontext

return self._vllm_metrics_scraper.activity.active() if self._vllm_metrics_scraper is not None else nullcontext()

@torch.no_grad()
async def eval(self, vllm_metrics_scraper: Optional[VLLMMetricsScraper] = None) -> Dict[str, float]:
"""
Expand All @@ -257,6 +266,7 @@ async def eval(self, vllm_metrics_scraper: Optional[VLLMMetricsScraper] = None)
trajectory_logger=self.trajectory_logger,
tracker=self.tracker,
vllm_metrics_scraper=vllm_metrics_scraper,
generation_activity=self._generation_activity,
)

async def train(self):
Expand Down Expand Up @@ -349,10 +359,13 @@ async def train(self):
# 1.1. generation phase
if self._vllm_metrics_scraper is not None:
self._vllm_metrics_scraper.resume()
with Timer("generate", self.all_timings):
generator_output: GeneratorOutput = await self.generate(generator_input)
if self._vllm_metrics_scraper is not None:
self._vllm_metrics_scraper.pause()
try:
with Timer("generate", self.all_timings):
with self._generation_activity():
generator_output: GeneratorOutput = await self.generate(generator_input)
finally:
if self._vllm_metrics_scraper is not None:
self._vllm_metrics_scraper.pause()

if self.cfg.generator.step_wise_trajectories:
# NOTE: We use instance_ids from `trajectory_ids` here instead of re-using `uids`
Expand Down
47 changes: 47 additions & 0 deletions skyrl/train/utils/generation_activity.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
"""Measure the union of overlapping generation calls."""

import time
from contextlib import contextmanager
from typing import Callable, Optional

from loguru import logger


class GenerationActivity:
"""Track active wall time without summing concurrent call durations."""

def __init__(self, clock: Callable[[], float] = time.monotonic, publish: Optional[Callable[[int], None]] = None):
self._clock = clock
self._publish = publish
self._count = 0
self._since = None
self._total = 0.0

@property
def seconds(self) -> float:
"""Return completed and currently active interval duration."""
return self._total + (self._clock() - self._since if self._since is not None else 0.0)

@contextmanager
def active(self):
"""Mark a call active and always close it on exception or cancellation."""
if self._count == 0:
self._since = self._clock()
self._count += 1
self._emit()
try:
yield
finally:
self._count -= 1
if self._count == 0:
self._total += self._clock() - self._since
self._since = None
self._emit()

def _emit(self):
if self._publish is not None:
try:
self._publish(self._count)
except Exception as error:
self._publish = None
logger.warning(f"Generation activity gauge disabled ({type(error).__name__})")
35 changes: 33 additions & 2 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.generation_activity import GenerationActivity
from skyrl.train.utils.vllm_window_statistics import latency_metrics

# vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`)
Expand Down Expand Up @@ -201,6 +202,8 @@ 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.activity = GenerationActivity()
self._activity_sampled_seconds = 0.0
self._prev_aggregated: Optional[Dict[str, float]] = None
self._prev_timestamp: Optional[float] = None
self._client: Optional[httpx.AsyncClient] = None
Expand Down Expand Up @@ -312,7 +315,28 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]:
sums.setdefault(hits, 0.0)
return {**sums, **means, **per_pos, **buckets}

async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, float]:
def publish_activity(self, run_name: str) -> None:
"""Expose the outstanding-call count alongside existing training phases."""
from ray.util.metrics import Gauge

gauge = Gauge(
"skyrl_generation_active_calls", description="Outstanding generation calls.", tag_keys=("run_name",)
)
gauge.set(0, tags={"run_name": run_name})
self.activity = GenerationActivity(publish=lambda value: gauge.set(value, tags={"run_name": run_name}))
self._activity_sampled_seconds = 0.0

async def sample_active(self) -> Dict[str, float]:
"""Sample overlapping train/eval generation using union active time."""
seconds = self.activity.seconds
delta = seconds - self._activity_sampled_seconds
result = await self.sample(generation_time_s=delta, allow_zero_duration=True)
self._activity_sampled_seconds = seconds
return result

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

``generation_time_s`` is the throughput denominator (engine generation
Expand All @@ -328,7 +352,14 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f
now = time.monotonic()
if self._prev_aggregated is not None and self._prev_timestamp is not None:
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
window = (
generation_time_s
if (
generation_time_s is not None
and (generation_time_s > 0 or (allow_zero_duration and generation_time_s == 0))
)
else dt
)
out = self._window_metrics(self._prev_aggregated, snapshot, window, "vllm/")
else:
out = self._window_metrics(None, snapshot, None, "vllm/") # gauges only
Expand Down
56 changes: 56 additions & 0 deletions tests/train/test_generation_activity.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
"""Tests for generation duration with overlap and failures."""

import pytest

from skyrl.train.utils.generation_activity import GenerationActivity


def test_overlap_counts_union_and_closes_on_error():
now = [0.0]
published = []
activity = GenerationActivity(clock=lambda: now[0], publish=published.append)
with activity.active():
now[0] = 2
with pytest.raises(RuntimeError), activity.active():
now[0] = 5
raise RuntimeError("generation failed")
now[0] = 7
assert activity.seconds == 7
now[0] = 10
with activity.active():
now[0] = 12
assert activity.seconds == 9
assert published == [1, 2, 1, 0, 1, 0]


@pytest.mark.parametrize("failed_count", [0, 1])
def test_failed_gauge_preserves_activity_timing(failed_count):
now = [0.0]

def publish(count):
if count == failed_count:
raise RuntimeError("gauge unavailable")

activity = GenerationActivity(clock=lambda: now[0], publish=publish)
with activity.active():
now[0] = 3
now[0] = 5
assert activity.seconds == 3
with activity.active():
now[0] = 7
assert activity.seconds == 5


@pytest.mark.asyncio
async def test_zero_active_time_does_not_fall_back_to_wall_time():
from unittest.mock import AsyncMock

from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper

scraper = VLLMMetricsScraper(urls=["test"])
scraper._read_snapshot = AsyncMock(
side_effect=[{"ray_vllm_generation_tokens_total": 0}, {"ray_vllm_generation_tokens_total": 10}]
)
await scraper.sample_active()
metrics = await scraper.sample_active()
assert "vllm/generation_throughput_tok_s" not in metrics
Loading