| |
| """Reproducible batch-1 CPU latency benchmark for PT or ONNX artifacts.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import json |
| import platform |
| import resource |
| import statistics |
| import sys |
| import time |
| from collections.abc import Callable |
| from contextlib import suppress |
| from pathlib import Path |
| from typing import Any |
|
|
| REPOSITORY_ROOT = Path(__file__).resolve().parents[1] |
| SOURCE_ROOT = REPOSITORY_ROOT / "src" |
| if str(SOURCE_ROOT) not in sys.path: |
| sys.path.insert(0, str(SOURCE_ROOT)) |
|
|
|
|
| def _portable_path(path: Path) -> str: |
| try: |
| return path.resolve().relative_to(REPOSITORY_ROOT).as_posix() |
| except ValueError: |
| return path.name |
|
|
|
|
| def _sha256(path: Path) -> str: |
| digest = hashlib.sha256() |
| with path.open("rb") as handle: |
| for block in iter(lambda: handle.read(1024 * 1024), b""): |
| digest.update(block) |
| return digest.hexdigest() |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--model", required=True, help="checkpoint .pt or exported .onnx") |
| parser.add_argument("--metadata", help="model_metadata.json for ONNX") |
| parser.add_argument("--output", default="artifacts/benchmarks/cpu.json") |
| parser.add_argument("--threads", type=int, default=1) |
| parser.add_argument("--batch-size", type=int, default=1) |
| parser.add_argument("--frames", type=int, help="override generated input frames") |
| parser.add_argument( |
| "--include-frontend", |
| action="store_true", |
| help="for ONNX, include waveform-to-log-mel preprocessing", |
| ) |
| parser.add_argument( |
| "--audio-seconds", |
| type=float, |
| help="generated audio duration for --include-frontend (default: model maximum)", |
| ) |
| parser.add_argument("--warmup", type=int, default=20) |
| parser.add_argument("--iterations", type=int, default=200) |
| return parser.parse_args() |
|
|
|
|
| def _percentile(values: list[float], quantile: float) -> float: |
| ordered = sorted(values) |
| position = (len(ordered) - 1) * quantile |
| lower = int(position) |
| upper = min(lower + 1, len(ordered) - 1) |
| fraction = position - lower |
| return ordered[lower] * (1.0 - fraction) + ordered[upper] * fraction |
|
|
|
|
| def _timed_loop( |
| inference: Callable[[], Any], warmup: int, iterations: int |
| ) -> tuple[float, list[float]]: |
| start = time.perf_counter_ns() |
| inference() |
| cold_ms = (time.perf_counter_ns() - start) / 1e6 |
| for _ in range(warmup): |
| inference() |
| latencies: list[float] = [] |
| for _ in range(iterations): |
| start = time.perf_counter_ns() |
| inference() |
| latencies.append((time.perf_counter_ns() - start) / 1e6) |
| return cold_ms, latencies |
|
|
|
|
| def _metadata_for_onnx(model_path: Path, explicit: str | None) -> dict[str, Any]: |
| path = Path(explicit) if explicit else model_path.parent / "model_metadata.json" |
| if not path.is_file(): |
| raise SystemExit(f"metadata not found: {path}") |
| loaded = json.loads(path.read_text(encoding="utf-8")) |
| if not isinstance(loaded, dict): |
| raise SystemExit("metadata must be a JSON object") |
| return loaded |
|
|
|
|
| def _benchmark_onnx( |
| model_path: Path, args: argparse.Namespace |
| ) -> tuple[dict[str, Any], float, list[float]]: |
| try: |
| import numpy as np |
| import onnxruntime as ort |
| except ImportError as exc: |
| raise SystemExit("ONNX benchmarking requires numpy and onnxruntime") from exc |
| metadata = _metadata_for_onnx(model_path, args.metadata) |
| frontend = metadata.get("frontend", metadata) |
| n_mels = int(frontend["n_mels"]) |
| frames = args.frames or int( |
| round( |
| float(frontend["max_seconds"]) |
| * int(frontend["sample_rate"]) |
| / int(frontend["hop_length"]) |
| ) |
| ) |
| rng = np.random.default_rng(17) |
| if args.include_frontend: |
| from turn_detection.runtime.predictor import OnnxEndpointPredictor |
|
|
| seconds = float(args.audio_seconds or frontend["max_seconds"]) |
| if seconds <= 0: |
| raise SystemExit("--audio-seconds must be positive") |
| sample_rate = int(frontend["sample_rate"]) |
| audio = rng.standard_normal(round(seconds * sample_rate), dtype=np.float32) * 0.05 |
| load_start = time.perf_counter_ns() |
| predictor = OnnxEndpointPredictor( |
| model_path, |
| args.metadata, |
| intra_op_threads=args.threads, |
| ) |
| load_ms = (time.perf_counter_ns() - load_start) / 1e6 |
|
|
| def infer_audio() -> Any: |
| return predictor.predict(audio, sample_rate) |
|
|
| cold_ms, latencies = _timed_loop(infer_audio, args.warmup, args.iterations) |
| return ( |
| { |
| "runtime": "onnxruntime", |
| "frames": frames, |
| "audio_seconds": seconds, |
| "load_ms": load_ms, |
| "scope": "end_to_end_waveform_to_probability", |
| }, |
| cold_ms, |
| latencies, |
| ) |
|
|
| features = rng.standard_normal((args.batch_size, n_mels, frames), dtype=np.float32) |
| mask = np.ones((args.batch_size, frames), dtype=np.float32) |
| options = ort.SessionOptions() |
| options.intra_op_num_threads = args.threads |
| options.inter_op_num_threads = 1 |
| options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL |
| options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL |
| load_start = time.perf_counter_ns() |
| session = ort.InferenceSession( |
| str(model_path), sess_options=options, providers=["CPUExecutionProvider"] |
| ) |
| load_ms = (time.perf_counter_ns() - load_start) / 1e6 |
|
|
| def infer() -> Any: |
| return session.run( |
| [metadata.get("endpoint_output_name") or "endpoint_probability"], |
| { |
| metadata.get("input_features_name", "log_mel"): features, |
| metadata.get("frame_mask_name", "frame_mask"): mask, |
| }, |
| ) |
|
|
| cold_ms, latencies = _timed_loop(infer, args.warmup, args.iterations) |
| return ( |
| { |
| "runtime": "onnxruntime", |
| "frames": frames, |
| "load_ms": load_ms, |
| "scope": "neural_model_only_log_mel_input", |
| }, |
| cold_ms, |
| latencies, |
| ) |
|
|
|
|
| def _benchmark_torch( |
| model_path: Path, args: argparse.Namespace |
| ) -> tuple[dict[str, Any], float, list[float]]: |
| try: |
| import torch |
| except ImportError as exc: |
| raise SystemExit("checkpoint benchmarking requires PyTorch") from exc |
| from turn_detection.models import load_model_checkpoint |
|
|
| torch.set_num_threads(args.threads) |
| with suppress(RuntimeError): |
| torch.set_num_interop_threads(1) |
| load_start = time.perf_counter_ns() |
| model, checkpoint = load_model_checkpoint(model_path, map_location="cpu") |
| model.eval() |
| load_ms = (time.perf_counter_ns() - load_start) / 1e6 |
| model_config = checkpoint["model_config"] |
| metadata = checkpoint.get("metadata", {}) |
| feature_config = metadata.get("feature_config", {}) |
| n_mels = int(model_config.get("n_mels", feature_config.get("n_mels", 80))) |
| frames = args.frames or int( |
| round( |
| float(metadata.get("max_seconds", 8.0)) |
| * int(feature_config.get("sample_rate", 16_000)) |
| / int(feature_config.get("hop_length", 160)) |
| ) |
| ) |
| generator = torch.Generator().manual_seed(17) |
| features = torch.randn( |
| (args.batch_size, n_mels, frames), generator=generator, dtype=torch.float32 |
| ) |
| mask = torch.ones((args.batch_size, frames), dtype=torch.bool) |
|
|
| def infer() -> Any: |
| with torch.inference_mode(): |
| return torch.sigmoid(model(features, mask).endpoint_logits) |
|
|
| cold_ms, latencies = _timed_loop(infer, args.warmup, args.iterations) |
| parameter_count = sum(parameter.numel() for parameter in model.parameters()) |
| return ( |
| { |
| "runtime": f"pytorch-{torch.__version__}", |
| "frames": frames, |
| "load_ms": load_ms, |
| "parameters": parameter_count, |
| }, |
| cold_ms, |
| latencies, |
| ) |
|
|
|
|
| def _peak_rss_mb() -> float: |
| value = float(resource.getrusage(resource.RUSAGE_SELF).ru_maxrss) |
| |
| return value / (1024.0**2) if platform.system() == "Darwin" else value / 1024.0 |
|
|
|
|
| def main() -> int: |
| args = parse_args() |
| if args.threads < 1 or args.batch_size < 1 or args.iterations < 1 or args.warmup < 0: |
| raise SystemExit("threads, batch-size, iterations must be positive; warmup non-negative") |
| if args.include_frontend and Path(args.model).suffix.lower() != ".onnx": |
| raise SystemExit("--include-frontend currently requires an ONNX model") |
| if args.include_frontend and args.batch_size != 1: |
| raise SystemExit("--include-frontend requires --batch-size 1") |
| model_path = Path(args.model) |
| if not model_path.is_absolute(): |
| model_path = REPOSITORY_ROOT / model_path |
| if model_path.suffix.lower() == ".onnx": |
| runtime, cold_ms, latencies = _benchmark_onnx(model_path, args) |
| else: |
| runtime, cold_ms, latencies = _benchmark_torch(model_path, args) |
|
|
| report = { |
| "artifact": _portable_path(model_path), |
| "artifact_bytes": model_path.stat().st_size, |
| "artifact_sha256": _sha256(model_path), |
| "cpu": platform.processor() or platform.machine(), |
| "platform": platform.platform(), |
| "python": platform.python_version(), |
| "threads": args.threads, |
| "batch_size": args.batch_size, |
| "warmup_iterations": args.warmup, |
| "measured_iterations": args.iterations, |
| **runtime, |
| "cold_first_inference_ms": cold_ms, |
| "warm_latency_ms": { |
| "mean": statistics.fmean(latencies), |
| "p50": _percentile(latencies, 0.50), |
| "p90": _percentile(latencies, 0.90), |
| "p95": _percentile(latencies, 0.95), |
| "p99": _percentile(latencies, 0.99), |
| "min": min(latencies), |
| "max": max(latencies), |
| }, |
| "examples_per_second": args.batch_size * 1000.0 / statistics.fmean(latencies), |
| "peak_rss_mb": _peak_rss_mb(), |
| } |
| output_path = Path(args.output) |
| if not output_path.is_absolute(): |
| output_path = REPOSITORY_ROOT / output_path |
| output_path.parent.mkdir(parents=True, exist_ok=True) |
| output_path.write_text( |
| json.dumps(report, indent=2, sort_keys=True, allow_nan=False), encoding="utf-8" |
| ) |
| print(json.dumps(report, indent=2)) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|