diff --git a/ais_bench/benchmark/cli/argument_parser.py b/ais_bench/benchmark/cli/argument_parser.py index 38a02e7b..c7d535bb 100644 --- a/ais_bench/benchmark/cli/argument_parser.py +++ b/ais_bench/benchmark/cli/argument_parser.py @@ -157,6 +157,12 @@ def _perf_parser(self): type=validate_pressure_time, default=DEFAULT_PRESSURE_TIME ) + parser.add_argument( + '--spec-decode', + help='Enable speculative decoding metrics collection. Only effective in --mode perf.', + action='store_true', + default=False, + ) def _custom_dataset_parser(self): """These args are all for the quick construction of custom datasets.""" diff --git a/ais_bench/benchmark/cli/workers.py b/ais_bench/benchmark/cli/workers.py index 66c533e2..3db84c14 100644 --- a/ais_bench/benchmark/cli/workers.py +++ b/ais_bench/benchmark/cli/workers.py @@ -1,9 +1,14 @@ +import glob import os import os.path as osp import copy import shutil +import asyncio +import json from abc import ABC, abstractmethod from collections import defaultdict +from dataclasses import dataclass, field +from typing import Any, NamedTuple from mmengine.config import ConfigDict @@ -24,6 +29,23 @@ logger = AISLogger() +class _URLEntry(NamedTuple): + """A single URL's before/after snapshot state.""" + before: Any = None # SpecDecodeSnapshot | None + error: str | None = None + + +@dataclass +class _SpecDecodeContext: + """Carries state between before/after spec decode Prometheus snapshots. + + Each unique metrics URL gets its own entry so that spec decode metrics + from different servers are collected and reported independently. + """ + enabled: bool = False + entries: dict[str, _URLEntry] = field(default_factory=dict) + + class BaseWorker(ABC): def __init__(self, args) -> None: self.args = args @@ -102,8 +124,13 @@ def do_work(self, cfg: ConfigDict): logger.info("Merging datasets with the same model and inferencer...") tasks = self._merge_datasets(tasks) + spec_ctx = self._spec_decode_before_snapshot(cfg) + runner = RUNNERS.build(cfg.infer.runner) runner(tasks) + + self._spec_decode_after_snapshot(cfg, spec_ctx) + logger.info("Inference tasks completed.") def _merge_datasets(self, tasks): @@ -134,6 +161,131 @@ def _update_tasks_cfg(self, tasks, cfg: ConfigDict): cfg.attack.dataset = task.datasets[0][0].abbr task.attack = cfg.attack + # ------------------------------------------------------------------ + # Speculative Decoding — per-URL before/after snapshot methods + # ------------------------------------------------------------------ + + def _spec_decode_before_snapshot(self, cfg: ConfigDict) -> _SpecDecodeContext: + """Fetch before-snapshots for every unique metrics URL.""" + cli_args = cfg.get("cli_args", {}) + ctx = _SpecDecodeContext() + + ctx.enabled = ( + cli_args.get("spec_decode", False) + and cli_args.get("mode") == "perf" + ) + if cli_args.get("spec_decode") and cli_args.get("mode") != "perf": + logger.warning( + "--spec-decode is only effective in --mode perf. " + "Ignoring spec decode for current mode '%s'.", + cli_args.get("mode"), + ) + + if not ctx.enabled: + return ctx + + from ais_bench.benchmark.spec_decode.urls import resolve_metrics_urls + + urls = resolve_metrics_urls(cfg.get("models", [])) + if not urls: + ctx.enabled = False + logger.info("Spec decode before-snapshot failed: no metrics URLs found.") + return ctx + + from ais_bench.benchmark.spec_decode.fetcher import ( + fetch_spec_decode_metrics_with_error, + ) + for url in urls: + snapshot, error = asyncio.run( + fetch_spec_decode_metrics_with_error(url) + ) + ctx.entries[url] = _URLEntry(before=snapshot, error=error) + if error: + logger.info( + "Spec decode [%s] before-snapshot failed: %s", url, error + ) + else: + logger.info( + "Spec decode [%s] before-snapshot captured successfully.", url + ) + return ctx + + def _spec_decode_after_snapshot( + self, cfg: ConfigDict, ctx: _SpecDecodeContext + ) -> None: + """Fetch after-snapshots, compute deltas, save per-URL results.""" + if not ctx.enabled: + return + + for url, entry in ctx.entries.items(): + try: + self._process_spec_decode_url(cfg, url, entry) + except Exception: + logger.warning( + "Spec decode [%s] after-snapshot failed unexpectedly, skipping", + url, exc_info=True, + ) + + def _process_spec_decode_url( + self, cfg: ConfigDict, url: str, entry: _URLEntry + ) -> None: + """After-snapshot → compute delta → save for a single URL.""" + from ais_bench.benchmark.spec_decode.fetcher import ( + fetch_spec_decode_metrics_with_error, + ) + from ais_bench.benchmark.spec_decode.calculator import ( + compute_spec_decode_stats, + ) + from ais_bench.benchmark.spec_decode.reporter import ( + save_spec_decode_result, + ) + + after_snapshot, after_error = asyncio.run( + fetch_spec_decode_metrics_with_error(url) + ) + error = self._merge_spec_decode_errors(entry.error, after_error) + + spec_stats = None + if entry.before is not None and after_snapshot is not None: + spec_stats = compute_spec_decode_stats(entry.before, after_snapshot) + if spec_stats is None: + error = "No spec decode activity detected during benchmark window" + + self._log_spec_decode_result(url, spec_stats, error) + + save_spec_decode_result( + spec_stats, error, cfg["work_dir"], url, + before_snapshot=entry.before, + after_snapshot=after_snapshot, + ) + + @staticmethod + def _merge_spec_decode_errors(before_error: str | None, after_error: str | None) -> str | None: + """Merge before/after error messages, preserving both when possible.""" + if not after_error: + return before_error + if before_error: + return f"{before_error}; {after_error}" + return after_error + + @staticmethod + def _log_spec_decode_result(url: str, spec_stats: dict | None, error: str | None) -> None: + """Log the outcome of spec decode collection for a single URL.""" + if spec_stats: + logger.info( + "Spec decode [%s] collected: acceptance_rate=%.2f%%, " + "acceptance_length=%.2f", + url, + spec_stats["acceptance_rate"], + spec_stats["acceptance_length"], + ) + else: + logger.info( + "Spec decode [%s] unavailable: %s", + url, error or "unknown reason", + ) + + class JudgeInfer(BaseWorker): def __init__(self, args) -> None: @@ -468,6 +620,45 @@ def do_work(self, cfg: ConfigDict) -> int: logger.info("Summarizing performance results...") summarizer.summarize() + # ========== Speculative Decoding Results ========== + if cfg.get("cli_args", {}).get("spec_decode", False): + self._output_spec_decode_results(cfg) + + @staticmethod + def _output_spec_decode_results(cfg: ConfigDict) -> None: + """Read all per-URL spec_decode_*.json files and print results.""" + from ais_bench.benchmark.spec_decode.reporter import ( + format_spec_decode_console, + format_spec_decode_na, + ) + + pattern = osp.join(cfg["work_dir"], "performances", "spec_decode_*.json") + spec_files = sorted(glob.glob(pattern)) + + if not spec_files: + logger.warning( + "Spec decode enabled but no result files found matching %s", + pattern, + ) + return + + for spec_file in spec_files: + try: + with open(spec_file, "r", encoding="utf-8") as f: + result = json.load(f) + except Exception: + logger.warning( + "Failed to read spec decode result file %s, skipping", + spec_file, exc_info=True, + ) + continue + + url = result.get("url", "") + if result.get("status") == "ok" and result.get("data"): + print(format_spec_decode_console(result["data"], url)) + else: + print(format_spec_decode_na(url, result.get("error"))) + WORK_FLOW = dict( all=[Infer, JudgeInfer, Eval, AccViz], diff --git a/ais_bench/benchmark/spec_decode/__init__.py b/ais_bench/benchmark/spec_decode/__init__.py new file mode 100644 index 00000000..74a6e3f2 --- /dev/null +++ b/ais_bench/benchmark/spec_decode/__init__.py @@ -0,0 +1,31 @@ +# Speculative Decoding Metrics Collection +# +# This module provides client-side collection of speculative decoding performance +# metrics from vLLM-compatible inference servers via the Prometheus /metrics endpoint. +# +# Architecture: +# snapshot.py - Data model + Prometheus text format parsing +# fetcher.py - Async HTTP fetching of /metrics +# calculator.py - Before/after delta computation + derived metrics +# reporter.py - Console and file output formatting + +from ais_bench.benchmark.spec_decode.snapshot import SpecDecodeSnapshot, parse_spec_decode_metrics +from ais_bench.benchmark.spec_decode.fetcher import ( + fetch_spec_decode_metrics_with_error, +) +from ais_bench.benchmark.spec_decode.calculator import compute_spec_decode_stats +from ais_bench.benchmark.spec_decode.reporter import ( + format_spec_decode_console, + format_spec_decode_na, + save_spec_decode_result, +) + +__all__ = [ + "SpecDecodeSnapshot", + "parse_spec_decode_metrics", + "fetch_spec_decode_metrics_with_error", + "compute_spec_decode_stats", + "format_spec_decode_console", + "format_spec_decode_na", + "save_spec_decode_result", +] diff --git a/ais_bench/benchmark/spec_decode/calculator.py b/ais_bench/benchmark/spec_decode/calculator.py new file mode 100644 index 00000000..d37cde6a --- /dev/null +++ b/ais_bench/benchmark/spec_decode/calculator.py @@ -0,0 +1,60 @@ +from ais_bench.benchmark.spec_decode.snapshot import SpecDecodeSnapshot + + +def compute_spec_decode_stats( + before: SpecDecodeSnapshot, + after: SpecDecodeSnapshot, +) -> dict | None: + """Compute derived spec decode metrics from before/after snapshots. + + All metrics are based on the delta between the two snapshots, isolating + only the activity that occurred during the benchmark window. + + Args: + before: Snapshot taken before the benchmark inference. + after: Snapshot taken after the benchmark inference. + + Returns: + A dict of derived metrics: + { + "num_drafts": int, # draft-and-verify cycles + "draft_tokens": int, # total candidate tokens proposed + "accepted_tokens": int, # total tokens accepted + "acceptance_rate": float, # percentage (0-100) + "acceptance_length": float, # avg tokens per forward pass + "per_position_acceptance_rates": {int: float}, # position → acceptance rate + } + Returns None if delta_draft_tokens <= 0 (no spec decode activity). + """ + delta_drafts = after.num_drafts - before.num_drafts + delta_draft_tokens = after.num_draft_tokens - before.num_draft_tokens + delta_accepted = after.num_accepted_tokens - before.num_accepted_tokens + + if delta_draft_tokens <= 0: + return None + + per_pos_rates: dict[int, float] = {} + if delta_drafts > 0: + all_positions = sorted( + set(before.accepted_per_pos.keys()) | set(after.accepted_per_pos.keys()) + ) + for pos in all_positions: + before_val = before.accepted_per_pos.get(pos, 0) + after_val = after.accepted_per_pos.get(pos, before_val) + delta_pos = max(after_val - before_val, 0) + per_pos_rates[pos] = delta_pos / delta_drafts + + acceptance_rate = (delta_accepted / delta_draft_tokens) * 100 + + acceptance_length = ( + 1 + (delta_accepted / delta_drafts) if delta_drafts > 0 else 0.0 + ) + + return { + "num_drafts": delta_drafts, + "draft_tokens": delta_draft_tokens, + "accepted_tokens": delta_accepted, + "acceptance_rate": acceptance_rate, + "acceptance_length": acceptance_length, + "per_position_acceptance_rates": per_pos_rates, + } diff --git a/ais_bench/benchmark/spec_decode/fetcher.py b/ais_bench/benchmark/spec_decode/fetcher.py new file mode 100644 index 00000000..df1b29d4 --- /dev/null +++ b/ais_bench/benchmark/spec_decode/fetcher.py @@ -0,0 +1,37 @@ +import aiohttp + +from ais_bench.benchmark.spec_decode.snapshot import SpecDecodeSnapshot, parse_spec_decode_metrics +from ais_bench.benchmark.utils.logging.logger import AISLogger + +_METRICS_FETCH_TIMEOUT = aiohttp.ClientTimeout(total=5) + +logger = AISLogger() + + +async def fetch_spec_decode_metrics_with_error( + metrics_url: str, +) -> tuple[SpecDecodeSnapshot | None, str | None]: + """GET {metrics_url} and return (snapshot, error_message) tuple. + + Returns (None, reason) on any failure so callers can distinguish + "server not enabled" from "network error" for N/A display. + """ + try: + async with aiohttp.ClientSession(timeout=_METRICS_FETCH_TIMEOUT, trust_env=True) as s: + async with s.get(metrics_url) as response: + if response.status != 200: + msg = f"Metrics endpoint returned HTTP {response.status}" + logger.debug("%s for %s", msg, metrics_url) + return None, msg + text = await response.text() + + snapshot = parse_spec_decode_metrics(text) + if snapshot is None: + msg = "No spec decode metrics found on server" + logger.debug("%s (%s)", msg, metrics_url) + return None, msg + return snapshot, None + except Exception as e: + msg = f"Failed to fetch metrics from {metrics_url}: {e}" + logger.debug(msg) + return None, msg diff --git a/ais_bench/benchmark/spec_decode/reporter.py b/ais_bench/benchmark/spec_decode/reporter.py new file mode 100644 index 00000000..acb5c00f --- /dev/null +++ b/ais_bench/benchmark/spec_decode/reporter.py @@ -0,0 +1,131 @@ +import json +import os +import os.path as osp +import re + +from ais_bench.benchmark.utils.logging.logger import AISLogger + +logger = AISLogger() + +_COL_WIDTH = 45 + + +def _make_title(url: str = "") -> str: + """Build the console title line with optional URL label.""" + base = "========== Speculative Decoding Metrics" + label = f" [{_url_label(url)}]" if url else "" + return f"{base}{label} ==========" + + +def format_spec_decode_console(stats: dict, url: str = "") -> str: + """Format spec decode statistics as a human-readable table.""" + title = _make_title(url) + sep = "=" * (len(title) - 2) + lines = [ + "", + sep, + title, + sep, + _fmt_line("Acceptance rate (%)", f"{stats['acceptance_rate']:.2f}"), + _fmt_line("Acceptance length", f"{stats['acceptance_length']:.2f}"), + _fmt_line("Drafts", str(stats["num_drafts"])), + _fmt_line("Draft tokens", str(stats["draft_tokens"])), + _fmt_line("Accepted tokens", str(stats["accepted_tokens"])), + _fmt_line( + "Per-position acceptance rates", + _format_per_pos_rates(stats.get("per_position_acceptance_rates", [])), + ), + ] + return "\n".join(lines) + + +def format_spec_decode_na(url: str = "", error_message: str | None = None) -> str: + """Format a N/A spec decode block with optional error reason.""" + title = _make_title(url) + sep = "=" * (len(title) - 2) + lines = [ + "", + sep, + title, + sep, + _fmt_line("Status", "N/A"), + ] + if error_message: + lines.append(_fmt_line("Reason", error_message)) + return "\n".join(lines) + + +def _fmt_line(label: str, value: str) -> str: + """Format a single label-value line with consistent alignment.""" + return f"{label:<{_COL_WIDTH}} {value}" + + +def _format_per_pos_rates(rates: dict) -> str: + """Format per-position acceptance rates dict.""" + formatted = [f"{pos}: {rate:.4f}" for pos, rate in sorted(rates.items())] + return "{" + ", ".join(formatted) + "}" + + +def _url_label(url: str) -> str: + """Extract a human-readable label from a metrics URL. + + "http://10.0.0.1:8080/metrics" → "10.0.0.1:8080" + """ + return re.sub(r"^https?://", "", url).rstrip("/").replace("/metrics", "") + + +def _url_to_key(url: str) -> str: + """Convert a metrics URL to a filesystem-safe key. + + "http://10.0.0.1:8080/metrics" → "10.0.0.1_8080" + """ + label = _url_label(url) + return re.sub(r"[^a-zA-Z0-9._-]", "_", label) + + +def save_spec_decode_result( + spec_stats: dict | None, + spec_error: str | None, + work_dir: str, + url: str, + before_snapshot=None, + after_snapshot=None, +) -> None: + """Save per-URL spec decode results to a JSON file. + + Creates performances/spec_decode_{url_key}.json so that each server + gets its own independent spec decode statistics file. Raw + before/after Prometheus counter snapshots are included alongside the + computed derived metrics for debugging and traceability. + + Args: + spec_stats: The computed spec decode statistics dict, or None. + spec_error: Error description if collection failed, or None. + work_dir: The benchmark work directory. + url: The metrics URL this data belongs to. + before_snapshot: Raw SpecDecodeSnapshot from before the benchmark. + after_snapshot: Raw SpecDecodeSnapshot from after the benchmark. + """ + output_dir = osp.join(work_dir, "performances") + os.makedirs(output_dir, exist_ok=True) + + url_key = _url_to_key(url) + result = { + "status": "ok" if spec_stats is not None else "na", + "url": url, + "error": spec_error, + "data": spec_stats, + "raw": { + "before": before_snapshot.to_dict() if before_snapshot else None, + "after": after_snapshot.to_dict() if after_snapshot else None, + }, + } + + output_path = osp.join(output_dir, f"spec_decode_{url_key}.json") + with open(output_path, "w", encoding="utf-8") as f: + json.dump(result, f, indent=2) + + logger.debug( + "Spec decode result saved to %s (url=%s, status=%s)", + output_path, url, result["status"], + ) diff --git a/ais_bench/benchmark/spec_decode/snapshot.py b/ais_bench/benchmark/spec_decode/snapshot.py new file mode 100644 index 00000000..2204d58e --- /dev/null +++ b/ais_bench/benchmark/spec_decode/snapshot.py @@ -0,0 +1,88 @@ +import contextlib +from dataclasses import dataclass, field + + +@dataclass +class SpecDecodeSnapshot: + """A single snapshot of spec decode counters from GET /metrics. + + All counters are monotonically increasing, process-lifetime cumulative values. + Use before/after delta to isolate activity during a benchmark window. + """ + + num_drafts: int = 0 + num_draft_tokens: int = 0 + num_accepted_tokens: int = 0 + accepted_per_pos: dict[int, int] = field(default_factory=dict) + + def to_dict(self) -> dict: + """Serialize to a plain dict for JSON output.""" + return { + "num_drafts": self.num_drafts, + "num_draft_tokens": self.num_draft_tokens, + "num_accepted_tokens": self.num_accepted_tokens, + "accepted_per_pos": dict(sorted(self.accepted_per_pos.items())), + } + + +def parse_spec_decode_metrics(text: str) -> "SpecDecodeSnapshot | None": + """Parse Prometheus text format, extracting spec decode counters. + + Args: + text: Raw response body from GET /metrics. + + Returns: + A SpecDecodeSnapshot if any spec decode metrics were found, or None + if the server does not have speculative decoding enabled. + """ + snapshot = SpecDecodeSnapshot() + found_any = False + + for line in text.split("\n"): + line = line.strip() + if not line or line.startswith("#"): + continue + if not line.startswith("vllm:spec_decode"): + continue + + parts = line.split(None, 1) + if len(parts) < 2: + continue + + metric_name = parts[0].split("{")[0] + if not metric_name.endswith("_total"): + continue + + with contextlib.suppress(ValueError): + val = int(float(parts[-1])) + found_any = True + + if "num_drafts" in metric_name and "num_draft_tokens" not in metric_name: + snapshot.num_drafts += val + elif "num_draft_tokens" in metric_name: + snapshot.num_draft_tokens += val + elif "num_accepted_tokens_per_pos" in metric_name: + pos = _extract_position(line) + if pos is not None: + snapshot.accepted_per_pos[pos] = ( + snapshot.accepted_per_pos.get(pos, 0) + val + ) + elif "num_accepted_tokens" in metric_name: + snapshot.num_accepted_tokens += val + + return snapshot if found_any else None + + +def _extract_position(line: str) -> int | None: + """Extract position=N from a Prometheus metric label string. + + Example: + vllm:spec_decode_num_accepted_tokens_per_pos_total{...,position="3"} 7710.0 + → returns 3 + """ + marker = 'position="' + if marker not in line: + return None + start = line.index(marker) + len(marker) + end = line.index('"', start) + return int(line[start:end]) diff --git a/ais_bench/benchmark/spec_decode/urls.py b/ais_bench/benchmark/spec_decode/urls.py new file mode 100644 index 00000000..01254904 --- /dev/null +++ b/ais_bench/benchmark/spec_decode/urls.py @@ -0,0 +1,41 @@ +import ipaddress + + +def resolve_metrics_urls(models: list[dict]) -> list[str]: + """Return deduplicated metrics URLs from a list of model configs. + + Args: + models: List of model configuration dicts, each may contain + ``host_ip`` (default ``"localhost"``) and ``host_port`` + (default ``8080``). + + Returns: + Deduplicated list of metrics URL strings such as + ``"http://10.0.0.1:8080/metrics"``. IPv6 addresses are + automatically wrapped in brackets (``[::1]``). Hostnames are + used as-is. If no models are provided, an empty list is returned. + """ + urls: list[str] = [] + seen: set[str] = set() + + for model in models: + host = model.get("host_ip", "localhost") + port = model.get("host_port", 8080) + host = _normalize_host(host) + url = f"http://{host}:{port}/metrics" + if url not in seen: + seen.add(url) + urls.append(url) + + return urls + + +def _normalize_host(host: str) -> str: + """Wrap IPv6 literals in brackets; leave IPv4 and hostnames unchanged.""" + try: + ip = ipaddress.ip_address(host) + if isinstance(ip, ipaddress.IPv6Address): + return f"[{ip}]" + except ValueError: + pass + return host diff --git a/tests/UT/spec_decode/test_calculator.py b/tests/UT/spec_decode/test_calculator.py new file mode 100644 index 00000000..a9f8517f --- /dev/null +++ b/tests/UT/spec_decode/test_calculator.py @@ -0,0 +1,91 @@ +"""Unit tests for spec_decode.calculator — delta and derived metrics.""" + +import math +import pytest +from ais_bench.benchmark.spec_decode.snapshot import SpecDecodeSnapshot +from ais_bench.benchmark.spec_decode.calculator import compute_spec_decode_stats + + +def _snap(num_drafts=0, num_draft_tokens=0, num_accepted=0, per_pos=None): + """Shorthand factory for SpecDecodeSnapshot in tests.""" + return SpecDecodeSnapshot( + num_drafts=num_drafts, + num_draft_tokens=num_draft_tokens, + num_accepted_tokens=num_accepted, + accepted_per_pos=per_pos or {}, + ) + + +class TestComputeSpecDecodeStats: + """Tests for compute_spec_decode_stats – the core delta computation.""" + + # ------------------------------------------------------------------ + # Normal case + # ------------------------------------------------------------------ + def test_compute_normal(self): + """Standard before/after pair → correct derived metrics.""" + before = _snap( + num_drafts=100, + num_draft_tokens=500, + num_accepted=300, + per_pos={0: 100, 1: 90, 2: 75, 3: 50, 4: 30}, + ) + after = _snap( + num_drafts=15520, # delta = 15420 + num_draft_tokens=77600, # delta = 77100 + num_accepted=50415, # delta = 50115 + per_pos={0: 15520, 1: 13968, 2: 11640, 3: 7760, 4: 4656}, + ) + + stats = compute_spec_decode_stats(before, after) + assert stats is not None + + # Raw deltas + assert stats["num_drafts"] == 15420 + assert stats["draft_tokens"] == 77100 + assert stats["accepted_tokens"] == 50115 + + # Derived: acceptance_rate = (50115 / 77100) * 100 + expected_rate = (50115 / 77100) * 100 + assert math.isclose(stats["acceptance_rate"], expected_rate, rel_tol=1e-9) + + # Derived: acceptance_length = 1 + (50115 / 15420) + expected_length = 1 + (50115 / 15420) + assert math.isclose(stats["acceptance_length"], expected_length, rel_tol=1e-9) + + # Per-position rates: delta_pos / delta_drafts + per_pos = stats["per_position_acceptance_rates"] + assert len(per_pos) == 5 + assert math.isclose(per_pos[0], 15420 / 15420, rel_tol=1e-9) + assert math.isclose(per_pos[1], (13968 - 90) / 15420, rel_tol=1e-9) + assert math.isclose(per_pos[4], (4656 - 30) / 15420, rel_tol=1e-9) + + # ------------------------------------------------------------------ + # No activity + # ------------------------------------------------------------------ + def test_compute_no_activity(self): + """Delta draft tokens = 0 → returns None (no spec decode activity).""" + before = _snap(num_drafts=10, num_draft_tokens=50, num_accepted=30) + after = _snap(num_drafts=10, num_draft_tokens=50, num_accepted=30) + assert compute_spec_decode_stats(before, after) is None + + # ------------------------------------------------------------------ + # First-time enable: before has no per_pos data + # ------------------------------------------------------------------ + def test_compute_empty_positions_in_before(self): + """Before has no per_pos data → positions are built from scratch.""" + before = _snap(num_drafts=0, num_draft_tokens=0, num_accepted=0) + after = _snap( + num_drafts=100, + num_draft_tokens=500, + num_accepted=300, + per_pos={0: 100, 1: 90}, + ) + + stats = compute_spec_decode_stats(before, after) + assert stats is not None + + per_pos = stats["per_position_acceptance_rates"] + assert len(per_pos) == 2 + assert math.isclose(per_pos[0], 1.0, rel_tol=1e-9) + assert math.isclose(per_pos[1], 0.9, rel_tol=1e-9) diff --git a/tests/UT/spec_decode/test_reporter.py b/tests/UT/spec_decode/test_reporter.py new file mode 100644 index 00000000..5518cced --- /dev/null +++ b/tests/UT/spec_decode/test_reporter.py @@ -0,0 +1,138 @@ +"""Unit tests for spec_decode.reporter — console and file output formatting.""" + +import json +import os +import tempfile +import pytest +from ais_bench.benchmark.spec_decode.reporter import ( + format_spec_decode_console, + format_spec_decode_na, + save_spec_decode_result, +) +from ais_bench.benchmark.spec_decode.snapshot import SpecDecodeSnapshot + + +# Sample stats dict (as returned by compute_spec_decode_stats) +_SAMPLE_STATS = { + "num_drafts": 15420, + "draft_tokens": 77100, + "accepted_tokens": 50115, + "acceptance_rate": 65.0, + "acceptance_length": 4.25, + "per_position_acceptance_rates": {0: 1.0, 1: 0.9, 2: 0.75, 3: 0.5, 4: 0.3}, +} + +_SAMPLE_URL = "http://10.0.0.1:8080/metrics" + + +class TestFormatSpecDecodeConsole: + """Tests for format_spec_decode_console.""" + + def test_format_with_url(self): + """Output includes the URL label and all metric values.""" + output = format_spec_decode_console(_SAMPLE_STATS, _SAMPLE_URL) + + # URL label should appear as [host:port] + assert "[10.0.0.1:8080]" in output + + # Key metric values should be present + assert "Acceptance rate (%)" in output + assert "65.00" in output + assert "Acceptance length" in output + assert "4.25" in output + assert "Drafts" in output + assert "15420" in output + assert "Draft tokens" in output + assert "77100" in output + assert "Accepted tokens" in output + assert "50115" in output + assert "Per-position acceptance rates" in output + assert "{0: 1.0000, 1: 0.9000, 2: 0.7500, 3: 0.5000, 4: 0.3000}" in output + + def test_format_without_url(self): + """Output without URL should NOT contain the URL bracket label.""" + output = format_spec_decode_console(_SAMPLE_STATS) + # The title line should not have [host:port] + assert "Metrics [" not in output + + +class TestFormatSpecDecodeNA: + """Tests for format_spec_decode_na.""" + + def test_format_na_with_url_and_error(self): + """N/A output includes URL label, N/A status, and reason.""" + output = format_spec_decode_na(_SAMPLE_URL, "Connection refused") + + assert "[10.0.0.1:8080]" in output + assert "N/A" in output + assert "Reason" in output + assert "Connection refused" in output + + def test_format_na_without_error(self): + """N/A output without error shows only status.""" + output = format_spec_decode_na(_SAMPLE_URL) + assert "N/A" in output + assert "Reason" not in output + + +class TestSaveSpecDecodeResult: + """Tests for save_spec_decode_result – file I/O.""" + + def test_save_ok_result(self): + """Save a successful result with raw snapshots → file content check.""" + with tempfile.TemporaryDirectory() as tmpdir: + before = SpecDecodeSnapshot( + num_drafts=10, num_draft_tokens=50, num_accepted_tokens=30, + accepted_per_pos={0: 10, 1: 9}, + ) + after = SpecDecodeSnapshot( + num_drafts=15430, num_draft_tokens=77150, num_accepted_tokens=50145, + accepted_per_pos={0: 15430, 1: 13887}, + ) + save_spec_decode_result( + _SAMPLE_STATS, None, tmpdir, _SAMPLE_URL, + before_snapshot=before, after_snapshot=after, + ) + + # Verify file exists with correct naming + files = os.listdir(os.path.join(tmpdir, "performances")) + assert len(files) == 1 + assert files[0].startswith("spec_decode_") + assert files[0].endswith(".json") + + # Verify content + filepath = os.path.join(tmpdir, "performances", files[0]) + with open(filepath, "r", encoding="utf-8") as f: + result = json.load(f) + + assert result["status"] == "ok" + assert result["url"] == _SAMPLE_URL + assert result["error"] is None + # Note: dict with int keys becomes str keys after JSON round-trip + assert result["data"]["num_drafts"] == _SAMPLE_STATS["num_drafts"] + assert result["data"]["acceptance_rate"] == _SAMPLE_STATS["acceptance_rate"] + assert result["data"]["acceptance_length"] == _SAMPLE_STATS["acceptance_length"] + + # Raw snapshots + raw = result["raw"] + assert raw["before"]["num_drafts"] == 10 + assert raw["before"]["num_draft_tokens"] == 50 + assert raw["after"]["num_drafts"] == 15430 + assert raw["after"]["num_draft_tokens"] == 77150 + assert raw["before"]["accepted_per_pos"] == {"0": 10, "1": 9} + + def test_save_na_result(self): + """Save a failed result → file content shows N/A.""" + with tempfile.TemporaryDirectory() as tmpdir: + error_msg = "No spec decode metrics found on server" + save_spec_decode_result(None, error_msg, tmpdir, _SAMPLE_URL) + + files = os.listdir(os.path.join(tmpdir, "performances")) + filepath = os.path.join(tmpdir, "performances", files[0]) + with open(filepath, "r", encoding="utf-8") as f: + result = json.load(f) + + assert result["status"] == "na" + assert result["url"] == _SAMPLE_URL + assert result["error"] == error_msg + assert result["data"] is None diff --git a/tests/UT/spec_decode/test_snapshot.py b/tests/UT/spec_decode/test_snapshot.py new file mode 100644 index 00000000..035e5b47 --- /dev/null +++ b/tests/UT/spec_decode/test_snapshot.py @@ -0,0 +1,87 @@ +"""Unit tests for spec_decode.snapshot — Prometheus text format parser.""" + +import pytest +from ais_bench.benchmark.spec_decode.snapshot import ( + SpecDecodeSnapshot, + parse_spec_decode_metrics, +) + + +# --------------------------------------------------------------------------- +# Standard Prometheus text fixture containing spec decode counters. +# Covers: HELP/TYPE lines (should be skipped), scientific notation (1.542e4), +# multiple "position" labels, and a non-spec-decode metric (should be +# ignored). +# --------------------------------------------------------------------------- +_VALID_METRICS_TEXT = """\ +# HELP vllm:spec_decode_num_drafts_total Number of spec decoding drafts. +# TYPE vllm:spec_decode_num_drafts_total counter +vllm:spec_decode_num_drafts_total{engine="0"} 15420.0 + +# HELP vllm:spec_decode_num_draft_tokens_total Number of draft tokens. +# TYPE vllm:spec_decode_num_draft_tokens_total counter +vllm:spec_decode_num_draft_tokens_total{engine="0"} 7.71e4 + +# HELP vllm:spec_decode_num_accepted_tokens_total Number of accepted tokens. +# TYPE vllm:spec_decode_num_accepted_tokens_total counter +vllm:spec_decode_num_accepted_tokens_total{engine="0"} 50115.0 + +# HELP vllm:spec_decode_num_accepted_tokens_per_pos_total Accepted tokens per draft position. +# TYPE vllm:spec_decode_num_accepted_tokens_per_pos_total counter +vllm:spec_decode_num_accepted_tokens_per_pos_total{engine="0",position="0"} 15420.0 +vllm:spec_decode_num_accepted_tokens_per_pos_total{engine="0",position="1"} 13878.0 +vllm:spec_decode_num_accepted_tokens_per_pos_total{engine="0",position="2"} 1.1565e4 + +# Some other metric that should be ignored +vllm:request_success_total{engine="0"} 9999.0 +""" + +# Metrics text that contains NO spec decode lines — +# simulating a server where speculative decoding is not enabled. +_NO_SPEC_DECODE_TEXT = """\ +# HELP vllm:request_success_total Number of successful requests. +# TYPE vllm:request_success_total counter +vllm:request_success_total{engine="0"} 500.0 + +# HELP vllm:num_requests_running Number of requests running. +# TYPE vllm:num_requests_running gauge +vllm:num_requests_running{engine="0"} 3.0 +""" + + +class TestParseSpecDecodeMetrics: + """Tests for parse_spec_decode_metrics – the core Prometheus parser.""" + + def test_parse_valid_metrics(self): + """Parse standard Prometheus text → correct SpecDecodeSnapshot.""" + snapshot = parse_spec_decode_metrics(_VALID_METRICS_TEXT) + + assert snapshot is not None, "Should return a snapshot when metrics are present" + + # Basic counters + assert snapshot.num_drafts == 15420 + assert snapshot.num_draft_tokens == 77100 # 7.71e4 + assert snapshot.num_accepted_tokens == 50115 + + # Per-position acceptance (3 positions in fixture) + assert len(snapshot.accepted_per_pos) == 3 + assert snapshot.accepted_per_pos[0] == 15420 + assert snapshot.accepted_per_pos[1] == 13878 + assert snapshot.accepted_per_pos[2] == 11565 # 1.1565e4 + + def test_parse_no_spec_decode(self): + """Metrics text without spec decode lines → returns None.""" + snapshot = parse_spec_decode_metrics(_NO_SPEC_DECODE_TEXT) + assert snapshot is None + + +class TestSpecDecodeSnapshot: + """Tests for the SpecDecodeSnapshot dataclass defaults.""" + + def test_defaults(self): + """All fields should default to zero / empty dict.""" + s = SpecDecodeSnapshot() + assert s.num_drafts == 0 + assert s.num_draft_tokens == 0 + assert s.num_accepted_tokens == 0 + assert s.accepted_per_pos == {} diff --git a/tests/UT/spec_decode/test_urls.py b/tests/UT/spec_decode/test_urls.py new file mode 100644 index 00000000..75dcf2ae --- /dev/null +++ b/tests/UT/spec_decode/test_urls.py @@ -0,0 +1,74 @@ +"""Unit tests for spec_decode.urls — metrics URL resolution and IPv6 handling.""" + +from ais_bench.benchmark.spec_decode.urls import resolve_metrics_urls + + +class TestResolveMetricsUrls: + """Tests for resolve_metrics_urls – the URL extraction and dedup logic.""" + + # ------------------------------------------------------------------ + # IPv4 – should NOT get brackets + # ------------------------------------------------------------------ + def test_ipv4_no_brackets(self): + models = [{"host_ip": "10.0.0.1", "host_port": 8080}] + urls = resolve_metrics_urls(models) + assert urls == ["http://10.0.0.1:8080/metrics"] + + # ------------------------------------------------------------------ + # IPv6 – should be wrapped in brackets + # ------------------------------------------------------------------ + def test_ipv6_full_address(self): + models = [{"host_ip": "2001:db8::1", "host_port": 8080}] + urls = resolve_metrics_urls(models) + assert urls == ["http://[2001:db8::1]:8080/metrics"] + + def test_ipv6_loopback(self): + models = [{"host_ip": "::1", "host_port": 9090}] + urls = resolve_metrics_urls(models) + assert urls == ["http://[::1]:9090/metrics"] + + def test_ipv6_zero_compressed(self): + models = [{"host_ip": "::", "host_port": 8080}] + urls = resolve_metrics_urls(models) + assert urls == ["http://[::]:8080/metrics"] + + # ------------------------------------------------------------------ + # Hostname – should pass through as-is + # ------------------------------------------------------------------ + def test_hostname_no_brackets(self): + models = [{"host_ip": "my-inference-server", "host_port": 8080}] + urls = resolve_metrics_urls(models) + assert urls == ["http://my-inference-server:8080/metrics"] + + def test_localhost_no_brackets(self): + models = [{"host_ip": "localhost", "host_port": 8080}] + urls = resolve_metrics_urls(models) + assert urls == ["http://localhost:8080/metrics"] + + # ------------------------------------------------------------------ + # Deduplication – same URL across multiple models → only one entry + # ------------------------------------------------------------------ + def test_duplicate_urls_dedup(self): + models = [ + {"host_ip": "10.0.0.1", "host_port": 8080}, + {"host_ip": "10.0.0.1", "host_port": 8080}, + {"host_ip": "10.0.0.2", "host_port": 9090}, + ] + urls = resolve_metrics_urls(models) + assert urls == [ + "http://10.0.0.1:8080/metrics", + "http://10.0.0.2:9090/metrics", + ] + + # ------------------------------------------------------------------ + # Defaults + # ------------------------------------------------------------------ + def test_default_host_and_port(self): + """No host_ip/host_port → falls back to localhost:8080.""" + models = [{}] + urls = resolve_metrics_urls(models) + assert urls == ["http://localhost:8080/metrics"] + + def test_empty_models(self): + """Empty model list → empty URL list.""" + assert resolve_metrics_urls([]) == []