#!/usr/bin/env python3 """Benchmark significant sample-extraction subprocesses using synthetic fixtures. This intentionally defaults to `stem=all` so the DSP stages can be measured without Demucs download/runtime noise. Use `--include-demucs` with a real input file if you want to benchmark stem separation on the current machine. """ from __future__ import annotations import argparse import json import statistics import sys import tempfile from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) import soundfile as sf from pipeline_runner import PipelineParams, run_extraction_pipeline from sample_extractor import cache_clear from synth_generator import generate_test_song def run_case(pattern: str, bars: int, bpm: float, run_index: int, clustering_mode: str) -> dict: tmp = Path(tempfile.mkdtemp(prefix="dse-bench-")) song = generate_test_song(pattern_name=pattern, bars=bars, bpm=bpm, add_bass=False, seed=42 + run_index) src = tmp / f"{pattern}-{bars}bars.wav" sf.write(src, song.drums_only, song.sr) cache_clear() params = PipelineParams(stem="all", clustering_mode=clustering_mode, target_min=4, target_max=12, synthesize=True) result = run_extraction_pipeline(src, tmp / "out", params) return { "pattern": pattern, "bars": bars, "bpm": bpm, "run_index": run_index, "clustering_mode": clustering_mode, "audio_duration_sec": result.audio_duration_sec, "total_duration_sec": result.duration_sec, "realtime_factor": result.realtime_factor, "hit_count": result.hit_count, "cluster_count": result.cluster_count, "stages": result.stages, } def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--runs", type=int, default=2) parser.add_argument("--bars", type=int, default=4) parser.add_argument("--bpm", type=float, default=120.0) parser.add_argument("--output", default="docs/benchmark-subprocesses.json") parser.add_argument("--clustering-mode", choices=["batch_quality", "online_preview"], default="batch_quality") args = parser.parse_args() # Warm imports/JIT and discard the result. run_case("rock", 1, args.bpm, -1, args.clustering_mode) rows = [] for run_index in range(args.runs): for pattern in ["rock", "funk", "halftime"]: rows.append(run_case(pattern, args.bars, args.bpm, run_index, args.clustering_mode)) stage_keys = [stage["key"] for stage in rows[0]["stages"]] summary = [] for key in stage_keys: values = [next(stage for stage in row["stages"] if stage["key"] == key)["duration_sec"] for row in rows] summary.append({ "stage": key, "mean_sec": round(statistics.mean(values), 6), "median_sec": round(statistics.median(values), 6), "min_sec": round(min(values), 6), "max_sec": round(max(values), 6), }) payload = {"clustering_mode": args.clustering_mode, "runs": rows, "summary": summary} out = Path(args.output) out.parent.mkdir(parents=True, exist_ok=True) out.write_text(json.dumps(payload, indent=2), encoding="utf-8") print(json.dumps(payload, indent=2)) return 0 if __name__ == "__main__": raise SystemExit(main())