Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions ais_bench/benchmark/cli/argument_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
191 changes: 191 additions & 0 deletions ais_bench/benchmark/cli/workers.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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],
Expand Down
31 changes: 31 additions & 0 deletions ais_bench/benchmark/spec_decode/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
60 changes: 60 additions & 0 deletions ais_bench/benchmark/spec_decode/calculator.py
Original file line number Diff line number Diff line change
@@ -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,
}
37 changes: 37 additions & 0 deletions ais_bench/benchmark/spec_decode/fetcher.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading