From cfe476225edc1c8d2dfce44f52edd3faf8ac7f5d Mon Sep 17 00:00:00 2001 From: Codex Date: Wed, 30 Sep 2026 16:22:29 +0000 Subject: [PATCH 1/2] feat: measure union generation activity in sync and async training Signed-off-by: Codex Signed-off-by: SumanthRH --- skyrl/train/evaluate.py | 12 +++++-- skyrl/train/fully_async_trainer.py | 24 ++++++++----- skyrl/train/trainer.py | 21 +++++++++--- skyrl/train/utils/generation_activity.py | 41 +++++++++++++++++++++++ skyrl/train/utils/vllm_metrics_scraper.py | 35 +++++++++++++++++-- tests/train/test_generation_activity.py | 38 +++++++++++++++++++++ 6 files changed, 153 insertions(+), 18 deletions(-) create mode 100644 skyrl/train/utils/generation_activity.py create mode 100644 tests/train/test_generation_activity.py diff --git a/skyrl/train/evaluate.py b/skyrl/train/evaluate.py index 6c5fb6467d..d65cf43c77 100644 --- a/skyrl/train/evaluate.py +++ b/skyrl/train/evaluate.py @@ -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. @@ -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) diff --git a/skyrl/train/fully_async_trainer.py b/skyrl/train/fully_async_trainer.py index 457562d4ba..1a813033dc 100644 --- a/skyrl/train/fully_async_trainer.py +++ b/skyrl/train/fully_async_trainer.py @@ -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): @@ -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 @@ -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!") @@ -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. diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index b3ad7023fd..95517e23c3 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -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 @@ -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]: """ @@ -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): @@ -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` diff --git a/skyrl/train/utils/generation_activity.py b/skyrl/train/utils/generation_activity.py new file mode 100644 index 0000000000..c7fc281b8d --- /dev/null +++ b/skyrl/train/utils/generation_activity.py @@ -0,0 +1,41 @@ +"""Measure the union of overlapping generation calls.""" + +import time +from contextlib import contextmanager +from typing import Callable, Optional + + +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: + self._publish(self._count) diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 08916a9e54..bd4a50f6dc 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.generation_activity import GenerationActivity from skyrl.train.utils.vllm_window_statistics import latency_metrics # vLLM metric base names after RayPrometheusStatLogger sanitization (`:` -> `_`) @@ -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 @@ -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 @@ -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 diff --git a/tests/train/test_generation_activity.py b/tests/train/test_generation_activity.py new file mode 100644 index 0000000000..5c018449f7 --- /dev/null +++ b/tests/train/test_generation_activity.py @@ -0,0 +1,38 @@ +"""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.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 From a99c0686df656e124700f1e6c8827948bbf6328a Mon Sep 17 00:00:00 2001 From: Codex Date: Thu, 1 Oct 2026 13:51:11 +0000 Subject: [PATCH 2/2] Keep generation timing intact when activity gauge fails Signed-off-by: Codex Signed-off-by: SumanthRH --- skyrl/train/utils/generation_activity.py | 8 +++++++- tests/train/test_generation_activity.py | 18 ++++++++++++++++++ 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/skyrl/train/utils/generation_activity.py b/skyrl/train/utils/generation_activity.py index c7fc281b8d..aa7c1430d4 100644 --- a/skyrl/train/utils/generation_activity.py +++ b/skyrl/train/utils/generation_activity.py @@ -4,6 +4,8 @@ from contextlib import contextmanager from typing import Callable, Optional +from loguru import logger + class GenerationActivity: """Track active wall time without summing concurrent call durations.""" @@ -38,4 +40,8 @@ def active(self): def _emit(self): if self._publish is not None: - self._publish(self._count) + try: + self._publish(self._count) + except Exception as error: + self._publish = None + logger.warning(f"Generation activity gauge disabled ({type(error).__name__})") diff --git a/tests/train/test_generation_activity.py b/tests/train/test_generation_activity.py index 5c018449f7..627b0a89f4 100644 --- a/tests/train/test_generation_activity.py +++ b/tests/train/test_generation_activity.py @@ -23,6 +23,24 @@ def test_overlap_counts_union_and_closes_on_error(): 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