From 0ee7e11775328f9c4d39ba39cbb87d7bd7ba6b75 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Sun, 4 Oct 2026 02:39:57 +0000 Subject: [PATCH 1/4] [metrics 5/5] Scope prefill and decode metrics and run summaries Signed-off-by: SumanthRH --- .../checkpointing-logging/vllm-metrics.mdx | 33 +++- skyrl/train/entrypoints/main_base.py | 15 +- skyrl/train/trainer.py | 4 +- skyrl/train/utils/vllm_metrics_scraper.py | 58 +++++++ tests/train/test_metrics_lifecycle.py | 11 +- tests/train/test_vllm_pd_metrics.py | 155 ++++++++++++++++++ 6 files changed, 265 insertions(+), 11 deletions(-) create mode 100644 tests/train/test_vllm_pd_metrics.py diff --git a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx index 980cef0c8b..d29c3bbded 100644 --- a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx +++ b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx @@ -97,10 +97,35 @@ Loading a checkpoint disables run aggregates while preserving step metrics. `LATEST` with no checkpoint still counts as a fresh run. A hard kill cannot publish final summaries or status. -Prefill/decode deployments keep step collection and raw Ray metrics. Generic -run summaries are omitted: both roles can report the same prompt and request, -and prefill can emit a placeholder output token. These need separate accounting -before their totals can be used for run comparisons. +### Prefill/decode metrics + +For SkyRL-managed prefill/decode deployments, metrics are collected separately +from each role's server actors. Summary keys add a role after the scope: +`vllm_correct_aggregate/eval/prefill/prompt_tokens_total` and +`vllm_correct_aggregate/eval/decode/output_tokens_total`, for example. Step keys +use the same role segment, such as `vllm/eval/prefill/ttft_seconds_avg`. + +Each role uses the same metrics and reductions as a regular deployment, +including running/waiting requests, KV usage, throughput and latency. Engine +imbalance is computed within each role. Separate queue charts can show whether +requests are waiting at prefill or decode. Fully async step keys use +`vllm/prefill/...` and `vllm/decode/...`; sync keys also include `train` or `eval`. +Scheduler gauges remain step snapshots and are not added to run summaries. + +These are engine observations: both roles can count the same prompt, and +prefill can report a placeholder output token. SkyRL keeps those observations +in their role scopes without adding the two roles' totals or merging their +latency histograms. Missing series and undefined ratios remain omitted. + +Role TTFT measures time within that engine's request. Prefill and decode TTFT +describe different parts of serving and cannot be added to obtain client TTFT. +End-to-end latency includes the proxy and transfer path and needs separate +client instrumentation. + +Both roles retain the fresh-run, baseline and finalization rules above. A +missing worker or reset omits the affected role's summary. External PD +deployments without known prefill/decode worker membership keep raw/step +collection but omit run summaries. ## Mark runs in Grafana diff --git a/skyrl/train/entrypoints/main_base.py b/skyrl/train/entrypoints/main_base.py index 8fb7259427..765b714c07 100644 --- a/skyrl/train/entrypoints/main_base.py +++ b/skyrl/train/entrypoints/main_base.py @@ -294,12 +294,23 @@ def _setup_trainer(self) -> RayPPOTrainer: 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 [])) + enable_pd = self.cfg.generator.inference_engine.enable_pd + groups = ( + (self._prefill_server_groups or []) + (self._decode_server_groups or []) + if enable_pd + else self._server_groups + ) 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) + if enable_pd: + num_prefill = sum(len(group.get_actors()) for group in (self._prefill_server_groups or [])) + trainer._vllm_metrics_scraper.set_worker_roles( + {"prefill": worker_ids[:num_prefill], "decode": worker_ids[num_prefill:]} + ) + else: + trainer._vllm_metrics_scraper.set_worker_ids(worker_ids) except Exception as error: trainer._vllm_metrics_scraper.set_worker_ids([]) logger.warning( diff --git a/skyrl/train/trainer.py b/skyrl/train/trainer.py index 49a874d9d5..c94e7280d3 100644 --- a/skyrl/train/trainer.py +++ b/skyrl/train/trainer.py @@ -245,7 +245,9 @@ async def finalize_metrics(self, status: str) -> None: 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: + if not self._resumed_from_checkpoint and ( + not self.cfg.generator.inference_engine.enable_pd or self._vllm_metrics_scraper.has_worker_roles + ): self.tracker.update_summary(summary) except Exception as e: logger.warning(f"Could not finalize vLLM metrics: {e}") diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index a2a5c3693a..3338026475 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -214,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._role_scrapers: Dict[str, VLLMMetricsScraper] = {} self.run_statistics = RunStatistics() self._prev_aggregated: Optional[Dict[str, float]] = None self._prev_timestamp: Optional[float] = None @@ -239,6 +240,34 @@ def __init__( 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) + self._role_scrapers = {} + + @property + def has_worker_roles(self) -> bool: + """Whether both PD roles have known, separate worker membership.""" + return bool(self._role_scrapers) + + def set_worker_roles(self, worker_ids_by_role: Dict[str, Iterable[str]]) -> None: + """Collect prefill and decode independently using their launched workers.""" + groups = {role: frozenset(ids) for role, ids in worker_ids_by_role.items()} + if set(groups) != {"prefill", "decode"} or not all(groups.values()): + raise ValueError("PD metrics require nonempty prefill and decode worker groups") + if not groups["prefill"].isdisjoint(groups["decode"]): + raise ValueError("Prefill and decode metric workers must be disjoint") + self._worker_ids = groups["prefill"] | groups["decode"] + self._role_scrapers = { + role: VLLMMetricsScraper(urls=self._urls, request_timeout_s=self._timeout, worker_ids=ids) + for role, ids in groups.items() + } + + def _role_metrics(self, results: List[Dict[str, float]]) -> Dict[str, float]: + """Keep the existing metrics under separate prefill and decode scopes.""" + out = {} + for role, metrics in zip(self._role_scrapers, results): + for key, value in metrics.items(): + prefix, name = key.rsplit("/", 1) + out[f"{prefix}/{role}/{name}"] = value + return out async def _get_client(self) -> httpx.AsyncClient: if self._client is None: @@ -246,6 +275,7 @@ async def _get_client(self) -> httpx.AsyncClient: return self._client async def aclose(self) -> None: + await asyncio.gather(*(scraper.aclose() for scraper in self._role_scrapers.values())) if self._client is not None: await self._client.aclose() self._client = None @@ -349,6 +379,10 @@ async def _read_snapshot(self) -> Optional[Dict[str, float]]: async def finalize(self) -> Dict[str, float]: """Attempt one final collection and close the HTTP client.""" try: + if self._role_scrapers: + return self._role_metrics( + await asyncio.gather(*(scraper.finalize() for scraper in self._role_scrapers.values())) + ) if self._label is not None: await self.stop() else: @@ -364,6 +398,10 @@ async def sample(self, generation_time_s: Optional[float] = None) -> Dict[str, f time since the previous call); ``None`` falls back to the wall-clock interval (fully-async overlap). """ + if self._role_scrapers: + return self._role_metrics( + await asyncio.gather(*(scraper.sample(generation_time_s) for scraper in self._role_scrapers.values())) + ) snapshot = await self._read_snapshot() if snapshot is None: self.run_statistics.incomplete.add("combined") @@ -398,6 +436,14 @@ async def start(self, label: str) -> None: """ if self._label is not None: raise ValueError(f"`start({label!r})` called while window {self._label!r} is still open") + if self._role_scrapers: + await asyncio.gather(*(scraper.start(label) for scraper in self._role_scrapers.values())) + # Both roles begin timing after their baselines are ready. + started = time.monotonic() + for scraper in self._role_scrapers.values(): + scraper._active_since = started + self._label = label + return self._window_prev = await self._read_snapshot() self._window_engines = self._engine_snapshot self._label = label @@ -407,6 +453,10 @@ async def start(self, label: str) -> None: def pause(self) -> None: """Stop accumulating active time until the next :meth:`resume`.""" + if self._role_scrapers: + for scraper in self._role_scrapers.values(): + scraper.pause() + return if self._label is None: raise ValueError("`pause` called without an open window") if self._paused: @@ -417,6 +467,10 @@ def pause(self) -> None: def resume(self) -> None: """Resume accumulating active time after a :meth:`pause`.""" + if self._role_scrapers: + for scraper in self._role_scrapers.values(): + scraper.resume() + return if self._label is None: raise ValueError("`resume` called without an open window") if not self._paused: @@ -433,6 +487,10 @@ async def stop(self) -> Dict[str, float]: """ if self._label is None: raise ValueError("`stop` called without an open window") + if self._role_scrapers: + result = self._role_metrics(await asyncio.gather(*(s.stop() for s in self._role_scrapers.values()))) + self._label = None + return result new_snapshot = await self._read_snapshot() if not self._paused: self._window_time_s += time.monotonic() - self._active_since diff --git a/tests/train/test_metrics_lifecycle.py b/tests/train/test_metrics_lifecycle.py index ef237ce371..22b405a100 100644 --- a/tests/train/test_metrics_lifecycle.py +++ b/tests/train/test_metrics_lifecycle.py @@ -49,21 +49,24 @@ def setup(): @pytest.mark.asyncio -@pytest.mark.parametrize("resumed,enable_pd", [(False, False), (True, False), (False, True)]) -async def test_trainer_finalizes_once_and_omits_resumed_or_pd_aggregates(resumed, enable_pd): +@pytest.mark.parametrize( + "resumed,enable_pd,known_roles", + [(False, False, False), (True, False, False), (False, True, False), (False, True, True), (True, True, True)], +) +async def test_trainer_finalizes_once_and_omits_resumed_or_unknown_pd_aggregates(resumed, enable_pd, known_roles): trainer = RayPPOTrainer.__new__(RayPPOTrainer) trainer._metrics_finalized = False trainer.tracker = Tracking("test", "test", backend="console") trainer._resumed_from_checkpoint = resumed trainer.cfg = SimpleNamespace(generator=SimpleNamespace(inference_engine=SimpleNamespace(enable_pd=enable_pd))) trainer.tracker.update_summary = Mock() - scraper = SimpleNamespace(finalize=AsyncMock(return_value={"tokens": 10})) + scraper = SimpleNamespace(finalize=AsyncMock(return_value={"tokens": 10}), has_worker_roles=known_roles) trainer._vllm_metrics_scraper = scraper await trainer.finalize_metrics("failed") await trainer.finalize_metrics("failed") scraper.finalize.assert_awaited_once() assert trainer.tracker.run_status == "failed" - if resumed or enable_pd: + if resumed or (enable_pd and not known_roles): trainer.tracker.update_summary.assert_called_once_with({"run_status": "failed"}) else: assert trainer.tracker.update_summary.call_args_list[0].args == ({"tokens": 10},) diff --git a/tests/train/test_vllm_pd_metrics.py b/tests/train/test_vllm_pd_metrics.py new file mode 100644 index 0000000000..c97261b1fa --- /dev/null +++ b/tests/train/test_vllm_pd_metrics.py @@ -0,0 +1,155 @@ +"""Role-specific collection for prefill/decode inference servers.""" + +from unittest.mock import Mock + +import httpx +import pytest + +from skyrl.train.utils.vllm_metrics_scraper import VLLMMetricsScraper + + +def _exports(phase, missing_worker=None): + lines = [] + for worker, prompt, output, latency in [ + ("p0", 60, 1, 0.1), + ("p1", 40, 1, 0.1), + ("d0", 60, 30, 0.3), + ("d1", 40, 70, 0.3), + ("unrelated", 10000, 10000, 100), + ]: + if worker == missing_worker: + continue + labels = f'WorkerId="{worker}",engine="0"' + for name, value in { + "num_requests_running": 2 if worker.startswith("p") else 3, + "num_requests_waiting": 5 if worker.startswith("p") else 1, + "generation_tokens_total": 1000 + phase * output, + "prompt_tokens_total": 1000 + phase * prompt, + "prefix_cache_queries_total": 1000 + phase * prompt, + "prefix_cache_hits_total": 100 + phase * prompt / 2, + }.items(): + lines.append(f"ray_vllm_{name}{{{labels}}} {value}") + for name, mean, bound in [ + ("time_to_first_token_seconds", latency, 0.2 if worker.startswith("p") else 1), + ("request_time_per_output_token_seconds", 0.03, 0.05), + ]: + count = 1 + phase + lines.extend( + [ + f"ray_vllm_{name}_sum{{{labels}}} {mean * count}", + f"ray_vllm_{name}_count{{{labels}}} {count}", + f'ray_vllm_{name}_bucket{{{labels},le="{bound}"}} {count}', + f'ray_vllm_{name}_bucket{{{labels},le="+Inf"}} {count}', + ] + ) + return "\n".join(lines) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", [False, True]) +async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync): + clock = {"time": 0, "phase": 0} + monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["time"]) + scraper = VLLMMetricsScraper(urls=["http://test/metrics"]) + scraper.set_worker_roles({"prefill": ["p0", "p1"], "decode": ["d0", "d1"]}) + + def transport(request): + if sync and clock["phase"] == 0 and request.headers["role"] == "decode": + clock["time"] += 3 + return httpx.Response(200, text=_exports(clock["phase"])) + + for role, child in scraper._role_scrapers.items(): + child._client = httpx.AsyncClient(headers={"role": role}, transport=httpx.MockTransport(transport)) + try: + if not sync: + await scraper.sample() + for phase, seconds in [(1, 2), (2, 8)]: + if sync: + await scraper.start("vllm/train") + scraper.pause() + clock["time"] += 5 + scraper.resume() + clock["phase"] = phase + clock["time"] += seconds + step = await scraper.stop() if sync else await scraper.sample() + prefix = "vllm/train/" if sync else "vllm/" + assert step[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(100 / seconds) + assert step[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(100 / seconds) + assert step[prefix + "prefill/prompt_throughput_cv"] == pytest.approx(0.2) + assert step[prefix + "decode/generation_throughput_cv"] == pytest.approx(0.4) + assert step[prefix + "prefill/num_requests_waiting"] == 10 + assert step[prefix + "decode/num_requests_waiting"] == 2 + # The sync window is closed; finalization still closes both role clients. + summary = await scraper.finalize() + finally: + await scraper.aclose() + scope = "train" if sync else "combined" + prefix = f"vllm_correct_aggregate/{scope}/" + assert summary[prefix + "prefill/prompt_tokens_total"] == 200 + assert summary[prefix + "decode/output_tokens_total"] == 200 + assert summary[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(20) + assert summary[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(20) + assert summary[prefix + "prefill/ttft_seconds_avg"] == pytest.approx(0.1) + assert summary[prefix + "decode/ttft_seconds_avg"] == pytest.approx(0.3) + assert summary[prefix + "prefill/ttft_seconds_p90"] == pytest.approx(0.18) + assert summary[prefix + "decode/ttft_seconds_p90"] == pytest.approx(0.9) + assert summary[prefix + "decode/tpot_seconds_avg"] == pytest.approx(0.03) + assert summary[prefix + "prefill/prefix_cache_hit_rate"] == pytest.approx(0.5) + assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.5) + assert summary[prefix + "prefill/output_tokens_total"] == 4 + assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.03) + assert summary[prefix + "decode/prompt_tokens_total"] == 200 + assert not any(key.rsplit("/", 2)[-2] not in {"prefill", "decode"} for key in summary) + assert all(child._client is None for child in scraper._role_scrapers.values()) + + +@pytest.mark.asyncio +async def test_missing_prefill_worker_omits_only_prefill_summary(monkeypatch): + clock = {"phase": 0} + monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["phase"]) + scraper = VLLMMetricsScraper(urls=["http://test/metrics"]) + scraper.set_worker_roles({"prefill": ["p0", "p1"], "decode": ["d0", "d1"]}) + for child in scraper._role_scrapers.values(): + child._client = httpx.AsyncClient( + transport=httpx.MockTransport( + lambda request: httpx.Response( + 200, text=_exports(clock["phase"], "p1" if clock["phase"] == 2 else None) + ) + ) + ) + await scraper.sample() + clock["phase"] = 1 + await scraper.sample() + clock["phase"] = 2 + summary = await scraper.finalize() + assert not any("/prefill/" in key for key in summary) + assert summary["vllm_correct_aggregate/combined/decode/output_tokens_total"] == 200 + + +def test_setup_uses_server_groups_to_assign_worker_roles(tmp_path, monkeypatch): + from skyrl.train.config import SkyRLTrainConfig + from skyrl.train.entrypoints.main_base import BasePPOExp + + cfg = SkyRLTrainConfig() + cfg.generator.inference_engine.enable_pd = True + cfg.trainer.fully_async.simulate_training = True + cfg.trainer.export_path = str(tmp_path / "export") + cfg.trainer.ckpt_path = str(tmp_path / "checkpoints") + exp = BasePPOExp.__new__(BasePPOExp) + exp.cfg = cfg + exp.tokenizer = Mock() + exp.train_dataset = exp.eval_dataset = exp.colocate_pg = None + actors = [Mock() for _ in range(3)] + exp._server_groups = None + exp._prefill_server_groups = [Mock(get_actors=Mock(return_value=actors[:2]))] + exp._decode_server_groups = [Mock(get_actors=Mock(return_value=actors[2:]))] + trainer = Mock() + exp.get_trainer = Mock(return_value=trainer) + for method in ("get_tracker", "get_inference_client", "get_generator", "get_trajectory_logger"): + setattr(exp, method, Mock()) + lookup = Mock(return_value=["p0", "p1", "d0"]) + monkeypatch.setattr("skyrl.train.entrypoints.main_base.ray.get", lookup) + assert exp._setup_trainer() is trainer + lookup.assert_called_once_with([a.get_ray_worker_id.remote.return_value for a in actors], timeout=10) + trainer._vllm_metrics_scraper.set_worker_roles.assert_called_once_with({"prefill": ["p0", "p1"], "decode": ["d0"]}) + trainer._vllm_metrics_scraper.set_worker_ids.assert_not_called() From 9742eb23c8e365c25ea1eaaacf9ba3286fb38e93 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Tue, 6 Oct 2026 23:04:15 +0000 Subject: [PATCH 2/4] Expose selected common metrics for prefill/decode comparisons Select prompt/cache metrics from prefill and output/TPOT metrics from decode. Approximate common TTFT with the sum of the role means and omit its P90 while preserving all role metrics. Signed-off-by: SumanthRH --- .../checkpointing-logging/vllm-metrics.mdx | 39 +++++++++++-- skyrl/train/utils/vllm_metrics_scraper.py | 26 ++++++++- tests/train/test_vllm_pd_metrics.py | 55 +++++++++++++++---- 3 files changed, 105 insertions(+), 15 deletions(-) diff --git a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx index d29c3bbded..e414870741 100644 --- a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx +++ b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx @@ -117,10 +117,41 @@ prefill can report a placeholder output token. SkyRL keeps those observations in their role scopes without adding the two roles' totals or merging their latency histograms. Missing series and undefined ratios remain omitted. -Role TTFT measures time within that engine's request. Prefill and decode TTFT -describe different parts of serving and cannot be added to obtain client TTFT. -End-to-end latency includes the proxy and transfer path and needs separate -client instrumentation. +Selected metrics also use the common keys from regular serving, such as +`vllm/train/prompt_throughput_tok_s` and +`vllm_correct_aggregate/train/output_tokens_total`. The same mapping applies +to sync `train`/`eval` steps, fully async `vllm/*` steps and all run-summary +scopes. Every role-scoped value remains available. + +| Common metric suffix | PD source | +| --- | --- | +| `prompt_throughput_tok_s`, `prompt_tokens_total` | Prefill | +| `generation_throughput_tok_s`, `output_tokens_total` | Decode | +| `ttft_seconds_avg` | Prefill mean + decode mean | +| `tpot_seconds_avg`, `tpot_seconds_p90` | Decode | +| `prefix_cache_hit_rate`, `external_prefix_cache_hit_rate` | Prefill | +| `kv_cache_usage_perc` | Prefill | +| `prompt_throughput_cv`, `prompt_throughput_cv_num_engines` | Prefill | +| `generation_throughput_cv`, `generation_throughput_cv_num_engines` | Decode | + +Role TTFT measures engine arrival through first output, including that engine's +queue. In sequential PD serving, the sum of the two means approximates the +engine portions of TTFT when both cover the same requests. It excludes gaps +between engine observations, such as proxy/network overhead. Async sampling +can include different request cohorts in each role. This approximation is not +client-observed TTFT; measuring that requires client instrumentation. +There is no common PD `ttft_seconds_p90`: role quantiles cannot be added. +Both role-specific TTFT P90s remain available. Request TPOT comes from decode's +first-to-last output interval, excluding the initial wait for KV transfer. + +The common KV-usage value is the mean prefill-engine usage, a convention for +comparisons rather than total cluster pressure. Decode can exhaust its cache +independently; use both role metrics to investigate capacity and preemptions. +Queues, preemptions, offload metrics, speculative-decoding diagnostics and +measurement durations stay role-scoped. Throughput retains its source role's +existing denominator. Gauges and engine CV remain step metrics only. +Missing selected-role values omit the common key; the common TTFT mean +requires both role means. Both roles retain the fresh-run, baseline and finalization rules above. A missing worker or reset omits the affected role's summary. External PD diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index 3338026475..d6b15617ca 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -78,6 +78,23 @@ ) _MEAN_METRICS = (_GAUGE_KV_CACHE_USAGE,) +# Common PD keys select one role's existing reduction. TTFT mean uses both roles below. +_PD_METRIC_ROLES = { + "prompt_throughput_tok_s": "prefill", + "prompt_tokens_total": "prefill", + "generation_throughput_tok_s": "decode", + "output_tokens_total": "decode", + "tpot_seconds_avg": "decode", + "tpot_seconds_p90": "decode", + "prefix_cache_hit_rate": "prefill", + "external_prefix_cache_hit_rate": "prefill", + "kv_cache_usage_perc": "prefill", + "prompt_throughput_cv": "prefill", + "prompt_throughput_cv_num_engines": "prefill", + "generation_throughput_cv": "decode", + "generation_throughput_cv_num_engines": "decode", +} + ParsedSamples = Dict[Tuple[str, FrozenSet[Tuple[str, str]]], float] @@ -261,12 +278,19 @@ def set_worker_roles(self, worker_ids_by_role: Dict[str, Iterable[str]]) -> None } def _role_metrics(self, results: List[Dict[str, float]]) -> Dict[str, float]: - """Keep the existing metrics under separate prefill and decode scopes.""" + """Keep role scopes and add selected common keys for PD comparisons.""" out = {} for role, metrics in zip(self._role_scrapers, results): for key, value in metrics.items(): prefix, name = key.rsplit("/", 1) out[f"{prefix}/{role}/{name}"] = value + if _PD_METRIC_ROLES.get(name) == role: + out[key] = value + # Sequential PD engine means approximate TTFT for matching request cohorts. + # Their P90s cannot be added; keep those only in the role scopes. + for key, value in results[0].items(): + if key.endswith("/ttft_seconds_avg") and key in results[1]: + out[key] = value + results[1][key] return out async def _get_client(self) -> httpx.AsyncClient: diff --git a/tests/train/test_vllm_pd_metrics.py b/tests/train/test_vllm_pd_metrics.py index c97261b1fa..8352ca0b1d 100644 --- a/tests/train/test_vllm_pd_metrics.py +++ b/tests/train/test_vllm_pd_metrics.py @@ -23,15 +23,22 @@ def _exports(phase, missing_worker=None): for name, value in { "num_requests_running": 2 if worker.startswith("p") else 3, "num_requests_waiting": 5 if worker.startswith("p") else 1, + "kv_cache_usage_perc": 0.2 if worker.startswith("p") else 0.7, "generation_tokens_total": 1000 + phase * output, "prompt_tokens_total": 1000 + phase * prompt, "prefix_cache_queries_total": 1000 + phase * prompt, - "prefix_cache_hits_total": 100 + phase * prompt / 2, + "prefix_cache_hits_total": 100 + phase * prompt * (0.5 if worker.startswith("p") else 0.25), + "external_prefix_cache_queries_total": 1000 + phase * prompt, + "external_prefix_cache_hits_total": 100 + phase * prompt * (0.2 if worker.startswith("p") else 0.75), }.items(): lines.append(f"ray_vllm_{name}{{{labels}}} {value}") for name, mean, bound in [ ("time_to_first_token_seconds", latency, 0.2 if worker.startswith("p") else 1), - ("request_time_per_output_token_seconds", 0.03, 0.05), + ( + "request_time_per_output_token_seconds", + 0.01 if worker.startswith("p") else 0.03, + 0.02 if worker.startswith("p") else 0.05, + ), ]: count = 1 + phase lines.extend( @@ -46,8 +53,8 @@ def _exports(phase, missing_worker=None): @pytest.mark.asyncio -@pytest.mark.parametrize("sync", [False, True]) -async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync): +@pytest.mark.parametrize("sync,label", [(False, "vllm"), (True, "vllm/train"), (True, "vllm/eval")]) +async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync, label): clock = {"time": 0, "phase": 0} monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["time"]) scraper = VLLMMetricsScraper(urls=["http://test/metrics"]) @@ -65,25 +72,39 @@ def transport(request): await scraper.sample() for phase, seconds in [(1, 2), (2, 8)]: if sync: - await scraper.start("vllm/train") + await scraper.start(label) scraper.pause() clock["time"] += 5 scraper.resume() clock["phase"] = phase clock["time"] += seconds step = await scraper.stop() if sync else await scraper.sample() - prefix = "vllm/train/" if sync else "vllm/" + prefix = label + "/" assert step[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(100 / seconds) assert step[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(100 / seconds) assert step[prefix + "prefill/prompt_throughput_cv"] == pytest.approx(0.2) assert step[prefix + "decode/generation_throughput_cv"] == pytest.approx(0.4) assert step[prefix + "prefill/num_requests_waiting"] == 10 assert step[prefix + "decode/num_requests_waiting"] == 2 + assert step[prefix + "prompt_throughput_tok_s"] == step[prefix + "prefill/prompt_throughput_tok_s"] + assert step[prefix + "generation_throughput_tok_s"] == step[prefix + "decode/generation_throughput_tok_s"] + assert step[prefix + "ttft_seconds_avg"] == pytest.approx(0.4) + assert prefix + "ttft_seconds_p90" not in step + assert step[prefix + "tpot_seconds_avg"] == pytest.approx(0.03) + assert step[prefix + "tpot_seconds_p90"] == pytest.approx(0.045) + assert step[prefix + "prefix_cache_hit_rate"] == pytest.approx(0.5) + assert step[prefix + "external_prefix_cache_hit_rate"] == pytest.approx(0.2) + assert step[prefix + "kv_cache_usage_perc"] == pytest.approx(0.2) + assert step[prefix + "prompt_throughput_cv"] == pytest.approx(0.2) + assert step[prefix + "generation_throughput_cv"] == pytest.approx(0.4) + assert step[prefix + "prompt_throughput_cv_num_engines"] == 2 + assert step[prefix + "generation_throughput_cv_num_engines"] == 2 + assert prefix + "num_requests_waiting" not in step # The sync window is closed; finalization still closes both role clients. summary = await scraper.finalize() finally: await scraper.aclose() - scope = "train" if sync else "combined" + scope = label.split("/")[-1] if sync else "combined" prefix = f"vllm_correct_aggregate/{scope}/" assert summary[prefix + "prefill/prompt_tokens_total"] == 200 assert summary[prefix + "decode/output_tokens_total"] == 200 @@ -95,11 +116,21 @@ def transport(request): assert summary[prefix + "decode/ttft_seconds_p90"] == pytest.approx(0.9) assert summary[prefix + "decode/tpot_seconds_avg"] == pytest.approx(0.03) assert summary[prefix + "prefill/prefix_cache_hit_rate"] == pytest.approx(0.5) - assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.5) + assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.25) assert summary[prefix + "prefill/output_tokens_total"] == 4 - assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.03) + assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.01) assert summary[prefix + "decode/prompt_tokens_total"] == 200 - assert not any(key.rsplit("/", 2)[-2] not in {"prefill", "decode"} for key in summary) + assert summary[prefix + "prompt_tokens_total"] == 200 + assert summary[prefix + "output_tokens_total"] == 200 + assert summary[prefix + "prompt_throughput_tok_s"] == pytest.approx(20) + assert summary[prefix + "generation_throughput_tok_s"] == pytest.approx(20) + assert summary[prefix + "ttft_seconds_avg"] == pytest.approx(0.4) + assert prefix + "ttft_seconds_p90" not in summary + assert summary[prefix + "tpot_seconds_avg"] == pytest.approx(0.03) + assert summary[prefix + "tpot_seconds_p90"] == pytest.approx(0.045) + assert summary[prefix + "prefix_cache_hit_rate"] == pytest.approx(0.5) + assert summary[prefix + "external_prefix_cache_hit_rate"] == pytest.approx(0.2) + assert prefix + "measurement_seconds" not in summary assert all(child._client is None for child in scraper._role_scrapers.values()) @@ -124,6 +155,10 @@ async def test_missing_prefill_worker_omits_only_prefill_summary(monkeypatch): summary = await scraper.finalize() assert not any("/prefill/" in key for key in summary) assert summary["vllm_correct_aggregate/combined/decode/output_tokens_total"] == 200 + assert summary["vllm_correct_aggregate/combined/output_tokens_total"] == 200 + assert "vllm_correct_aggregate/combined/prompt_tokens_total" not in summary + assert "vllm_correct_aggregate/combined/ttft_seconds_avg" not in summary + assert summary["vllm_correct_aggregate/combined/tpot_seconds_p90"] == pytest.approx(0.045) def test_setup_uses_server_groups_to_assign_worker_roles(tmp_path, monkeypatch): From b40db2d0ab93d343379c8447edb5e2ce35c863e7 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Tue, 6 Oct 2026 23:55:11 +0000 Subject: [PATCH 3/4] Revert common prefill/decode metric keys Keep PD step metrics and run summaries under their prefill/decode roles. Reverts 9742eb23c8e365c25ea1eaaacf9ba3286fb38e93. Signed-off-by: SumanthRH --- .../checkpointing-logging/vllm-metrics.mdx | 39 ++----------- skyrl/train/utils/vllm_metrics_scraper.py | 26 +-------- tests/train/test_vllm_pd_metrics.py | 55 ++++--------------- 3 files changed, 15 insertions(+), 105 deletions(-) diff --git a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx index e414870741..d29c3bbded 100644 --- a/docs/content/docs/checkpointing-logging/vllm-metrics.mdx +++ b/docs/content/docs/checkpointing-logging/vllm-metrics.mdx @@ -117,41 +117,10 @@ prefill can report a placeholder output token. SkyRL keeps those observations in their role scopes without adding the two roles' totals or merging their latency histograms. Missing series and undefined ratios remain omitted. -Selected metrics also use the common keys from regular serving, such as -`vllm/train/prompt_throughput_tok_s` and -`vllm_correct_aggregate/train/output_tokens_total`. The same mapping applies -to sync `train`/`eval` steps, fully async `vllm/*` steps and all run-summary -scopes. Every role-scoped value remains available. - -| Common metric suffix | PD source | -| --- | --- | -| `prompt_throughput_tok_s`, `prompt_tokens_total` | Prefill | -| `generation_throughput_tok_s`, `output_tokens_total` | Decode | -| `ttft_seconds_avg` | Prefill mean + decode mean | -| `tpot_seconds_avg`, `tpot_seconds_p90` | Decode | -| `prefix_cache_hit_rate`, `external_prefix_cache_hit_rate` | Prefill | -| `kv_cache_usage_perc` | Prefill | -| `prompt_throughput_cv`, `prompt_throughput_cv_num_engines` | Prefill | -| `generation_throughput_cv`, `generation_throughput_cv_num_engines` | Decode | - -Role TTFT measures engine arrival through first output, including that engine's -queue. In sequential PD serving, the sum of the two means approximates the -engine portions of TTFT when both cover the same requests. It excludes gaps -between engine observations, such as proxy/network overhead. Async sampling -can include different request cohorts in each role. This approximation is not -client-observed TTFT; measuring that requires client instrumentation. -There is no common PD `ttft_seconds_p90`: role quantiles cannot be added. -Both role-specific TTFT P90s remain available. Request TPOT comes from decode's -first-to-last output interval, excluding the initial wait for KV transfer. - -The common KV-usage value is the mean prefill-engine usage, a convention for -comparisons rather than total cluster pressure. Decode can exhaust its cache -independently; use both role metrics to investigate capacity and preemptions. -Queues, preemptions, offload metrics, speculative-decoding diagnostics and -measurement durations stay role-scoped. Throughput retains its source role's -existing denominator. Gauges and engine CV remain step metrics only. -Missing selected-role values omit the common key; the common TTFT mean -requires both role means. +Role TTFT measures time within that engine's request. Prefill and decode TTFT +describe different parts of serving and cannot be added to obtain client TTFT. +End-to-end latency includes the proxy and transfer path and needs separate +client instrumentation. Both roles retain the fresh-run, baseline and finalization rules above. A missing worker or reset omits the affected role's summary. External PD diff --git a/skyrl/train/utils/vllm_metrics_scraper.py b/skyrl/train/utils/vllm_metrics_scraper.py index d6b15617ca..3338026475 100644 --- a/skyrl/train/utils/vllm_metrics_scraper.py +++ b/skyrl/train/utils/vllm_metrics_scraper.py @@ -78,23 +78,6 @@ ) _MEAN_METRICS = (_GAUGE_KV_CACHE_USAGE,) -# Common PD keys select one role's existing reduction. TTFT mean uses both roles below. -_PD_METRIC_ROLES = { - "prompt_throughput_tok_s": "prefill", - "prompt_tokens_total": "prefill", - "generation_throughput_tok_s": "decode", - "output_tokens_total": "decode", - "tpot_seconds_avg": "decode", - "tpot_seconds_p90": "decode", - "prefix_cache_hit_rate": "prefill", - "external_prefix_cache_hit_rate": "prefill", - "kv_cache_usage_perc": "prefill", - "prompt_throughput_cv": "prefill", - "prompt_throughput_cv_num_engines": "prefill", - "generation_throughput_cv": "decode", - "generation_throughput_cv_num_engines": "decode", -} - ParsedSamples = Dict[Tuple[str, FrozenSet[Tuple[str, str]]], float] @@ -278,19 +261,12 @@ def set_worker_roles(self, worker_ids_by_role: Dict[str, Iterable[str]]) -> None } def _role_metrics(self, results: List[Dict[str, float]]) -> Dict[str, float]: - """Keep role scopes and add selected common keys for PD comparisons.""" + """Keep the existing metrics under separate prefill and decode scopes.""" out = {} for role, metrics in zip(self._role_scrapers, results): for key, value in metrics.items(): prefix, name = key.rsplit("/", 1) out[f"{prefix}/{role}/{name}"] = value - if _PD_METRIC_ROLES.get(name) == role: - out[key] = value - # Sequential PD engine means approximate TTFT for matching request cohorts. - # Their P90s cannot be added; keep those only in the role scopes. - for key, value in results[0].items(): - if key.endswith("/ttft_seconds_avg") and key in results[1]: - out[key] = value + results[1][key] return out async def _get_client(self) -> httpx.AsyncClient: diff --git a/tests/train/test_vllm_pd_metrics.py b/tests/train/test_vllm_pd_metrics.py index 8352ca0b1d..c97261b1fa 100644 --- a/tests/train/test_vllm_pd_metrics.py +++ b/tests/train/test_vllm_pd_metrics.py @@ -23,22 +23,15 @@ def _exports(phase, missing_worker=None): for name, value in { "num_requests_running": 2 if worker.startswith("p") else 3, "num_requests_waiting": 5 if worker.startswith("p") else 1, - "kv_cache_usage_perc": 0.2 if worker.startswith("p") else 0.7, "generation_tokens_total": 1000 + phase * output, "prompt_tokens_total": 1000 + phase * prompt, "prefix_cache_queries_total": 1000 + phase * prompt, - "prefix_cache_hits_total": 100 + phase * prompt * (0.5 if worker.startswith("p") else 0.25), - "external_prefix_cache_queries_total": 1000 + phase * prompt, - "external_prefix_cache_hits_total": 100 + phase * prompt * (0.2 if worker.startswith("p") else 0.75), + "prefix_cache_hits_total": 100 + phase * prompt / 2, }.items(): lines.append(f"ray_vllm_{name}{{{labels}}} {value}") for name, mean, bound in [ ("time_to_first_token_seconds", latency, 0.2 if worker.startswith("p") else 1), - ( - "request_time_per_output_token_seconds", - 0.01 if worker.startswith("p") else 0.03, - 0.02 if worker.startswith("p") else 0.05, - ), + ("request_time_per_output_token_seconds", 0.03, 0.05), ]: count = 1 + phase lines.extend( @@ -53,8 +46,8 @@ def _exports(phase, missing_worker=None): @pytest.mark.asyncio -@pytest.mark.parametrize("sync,label", [(False, "vllm"), (True, "vllm/train"), (True, "vllm/eval")]) -async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync, label): +@pytest.mark.parametrize("sync", [False, True]) +async def test_pd_windows_keep_role_counters_histograms_and_reused_baselines_separate(monkeypatch, sync): clock = {"time": 0, "phase": 0} monkeypatch.setattr("skyrl.train.utils.vllm_metrics_scraper.time.monotonic", lambda: clock["time"]) scraper = VLLMMetricsScraper(urls=["http://test/metrics"]) @@ -72,39 +65,25 @@ def transport(request): await scraper.sample() for phase, seconds in [(1, 2), (2, 8)]: if sync: - await scraper.start(label) + await scraper.start("vllm/train") scraper.pause() clock["time"] += 5 scraper.resume() clock["phase"] = phase clock["time"] += seconds step = await scraper.stop() if sync else await scraper.sample() - prefix = label + "/" + prefix = "vllm/train/" if sync else "vllm/" assert step[prefix + "prefill/prompt_throughput_tok_s"] == pytest.approx(100 / seconds) assert step[prefix + "decode/generation_throughput_tok_s"] == pytest.approx(100 / seconds) assert step[prefix + "prefill/prompt_throughput_cv"] == pytest.approx(0.2) assert step[prefix + "decode/generation_throughput_cv"] == pytest.approx(0.4) assert step[prefix + "prefill/num_requests_waiting"] == 10 assert step[prefix + "decode/num_requests_waiting"] == 2 - assert step[prefix + "prompt_throughput_tok_s"] == step[prefix + "prefill/prompt_throughput_tok_s"] - assert step[prefix + "generation_throughput_tok_s"] == step[prefix + "decode/generation_throughput_tok_s"] - assert step[prefix + "ttft_seconds_avg"] == pytest.approx(0.4) - assert prefix + "ttft_seconds_p90" not in step - assert step[prefix + "tpot_seconds_avg"] == pytest.approx(0.03) - assert step[prefix + "tpot_seconds_p90"] == pytest.approx(0.045) - assert step[prefix + "prefix_cache_hit_rate"] == pytest.approx(0.5) - assert step[prefix + "external_prefix_cache_hit_rate"] == pytest.approx(0.2) - assert step[prefix + "kv_cache_usage_perc"] == pytest.approx(0.2) - assert step[prefix + "prompt_throughput_cv"] == pytest.approx(0.2) - assert step[prefix + "generation_throughput_cv"] == pytest.approx(0.4) - assert step[prefix + "prompt_throughput_cv_num_engines"] == 2 - assert step[prefix + "generation_throughput_cv_num_engines"] == 2 - assert prefix + "num_requests_waiting" not in step # The sync window is closed; finalization still closes both role clients. summary = await scraper.finalize() finally: await scraper.aclose() - scope = label.split("/")[-1] if sync else "combined" + scope = "train" if sync else "combined" prefix = f"vllm_correct_aggregate/{scope}/" assert summary[prefix + "prefill/prompt_tokens_total"] == 200 assert summary[prefix + "decode/output_tokens_total"] == 200 @@ -116,21 +95,11 @@ def transport(request): assert summary[prefix + "decode/ttft_seconds_p90"] == pytest.approx(0.9) assert summary[prefix + "decode/tpot_seconds_avg"] == pytest.approx(0.03) assert summary[prefix + "prefill/prefix_cache_hit_rate"] == pytest.approx(0.5) - assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.25) + assert summary[prefix + "decode/prefix_cache_hit_rate"] == pytest.approx(0.5) assert summary[prefix + "prefill/output_tokens_total"] == 4 - assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.01) + assert summary[prefix + "prefill/tpot_seconds_avg"] == pytest.approx(0.03) assert summary[prefix + "decode/prompt_tokens_total"] == 200 - assert summary[prefix + "prompt_tokens_total"] == 200 - assert summary[prefix + "output_tokens_total"] == 200 - assert summary[prefix + "prompt_throughput_tok_s"] == pytest.approx(20) - assert summary[prefix + "generation_throughput_tok_s"] == pytest.approx(20) - assert summary[prefix + "ttft_seconds_avg"] == pytest.approx(0.4) - assert prefix + "ttft_seconds_p90" not in summary - assert summary[prefix + "tpot_seconds_avg"] == pytest.approx(0.03) - assert summary[prefix + "tpot_seconds_p90"] == pytest.approx(0.045) - assert summary[prefix + "prefix_cache_hit_rate"] == pytest.approx(0.5) - assert summary[prefix + "external_prefix_cache_hit_rate"] == pytest.approx(0.2) - assert prefix + "measurement_seconds" not in summary + assert not any(key.rsplit("/", 2)[-2] not in {"prefill", "decode"} for key in summary) assert all(child._client is None for child in scraper._role_scrapers.values()) @@ -155,10 +124,6 @@ async def test_missing_prefill_worker_omits_only_prefill_summary(monkeypatch): summary = await scraper.finalize() assert not any("/prefill/" in key for key in summary) assert summary["vllm_correct_aggregate/combined/decode/output_tokens_total"] == 200 - assert summary["vllm_correct_aggregate/combined/output_tokens_total"] == 200 - assert "vllm_correct_aggregate/combined/prompt_tokens_total" not in summary - assert "vllm_correct_aggregate/combined/ttft_seconds_avg" not in summary - assert summary["vllm_correct_aggregate/combined/tpot_seconds_p90"] == pytest.approx(0.045) def test_setup_uses_server_groups_to_assign_worker_roles(tmp_path, monkeypatch): From 53d51cee1046b1a4fdfe76eae462946387edebc9 Mon Sep 17 00:00:00 2001 From: SumanthRH Date: Wed, 7 Oct 2026 01:48:01 +0000 Subject: [PATCH 4/4] Fix chat template CPU test across midnight Signed-off-by: SumanthRH --- tests/train/generators/chat_templating_test_constants.py | 6 ++---- .../test_skyrl_gym_generator_chat_templating.py | 8 +++++++- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/tests/train/generators/chat_templating_test_constants.py b/tests/train/generators/chat_templating_test_constants.py index 55767e7538..dc27cc7bf1 100644 --- a/tests/train/generators/chat_templating_test_constants.py +++ b/tests/train/generators/chat_templating_test_constants.py @@ -3,8 +3,6 @@ tests/train/generators/test_skyrl_gym_generator_chat_templating.py::test_skyrl_gym_generator_chat_templating_exact """ -from datetime import date - # Produced by expected_str = tokenizer.apply_chat_template(expected_chat_history, tokenize=False) # where expected_chat_history is: @@ -35,10 +33,10 @@ def get_expected_chat_history(mock_response_text: str): b<|im_end|> """ -LLAMA3_2_EXPECTED_STR = f"""<|begin_of_text|><|start_header_id|>system<|end_header_id|> +LLAMA3_2_EXPECTED_STR = """<|begin_of_text|><|start_header_id|>system<|end_header_id|> Cutting Knowledge Date: December 2023 -Today Date: {date.today().strftime("%d %b %Y")} +Today Date: 26 Jul 2024 <|eot_id|><|start_header_id|>user<|end_header_id|> diff --git a/tests/train/generators/test_skyrl_gym_generator_chat_templating.py b/tests/train/generators/test_skyrl_gym_generator_chat_templating.py index 6b210f4092..6ffc428021 100644 --- a/tests/train/generators/test_skyrl_gym_generator_chat_templating.py +++ b/tests/train/generators/test_skyrl_gym_generator_chat_templating.py @@ -2,6 +2,7 @@ uv run --extra dev --isolated pytest tests/train/generators/test_skyrl_gym_generator_chat_templating.py """ +from datetime import datetime from pathlib import Path from typing import Any, Dict from unittest.mock import AsyncMock, MagicMock @@ -134,7 +135,7 @@ def _make_input_batch(prompt, extras): "qwen3-custom_chat_template_builtin", ], ) -async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenization_codepath, expected_str): +async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenization_codepath, expected_str, monkeypatch): """ Tests the behavior of chat templating for various models in multi-turn conversation. @@ -144,6 +145,11 @@ async def test_skyrl_gym_generator_chat_templating_exact(model_name, tokenizatio We hardcode the expected string in the constants file, so it is easier to check. But we also double check that those expected strings are correct by applying the chat template on the expected chat history. """ + # Keep the template's clock consistent with the fixed expected date across midnight. + template_datetime = MagicMock() + template_datetime.now.return_value = datetime(2024, 7, 26) + monkeypatch.setattr("transformers.utils.chat_template_utils.datetime", template_datetime) + # 1. Preparations to mock the generation. _register_test_env_if_needed() # Register only when needed tokenizer = AutoTokenizer.from_pretrained(model_name)