drum-sample-extractor / scripts /benchmark_subprocesses.py
ChatGPT
feat: add run history and online clustering
b8fa9bf
Raw
History Blame Contribute Delete
3.32 kB
#!/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())