Download hv_note.py from zeechimp/hv-note: direct link, hf CLI and curl.
- Browser
- Download file 27.8 kB
-
https://huggingface.co/zeechimp/hv-note/resolve/main/hv_note.py
- Command line
-
hf download hf://zeechimp/hv-note/hv_note.py
-
curl -L -o hv_note.py https://huggingface.co/zeechimp/hv-note/resolve/main/hv_note.py
27.8 kB
| """ | |
| hv-note | |
| ======= | |
| Symbolic audio event extraction. No lexicon. No translation layer. | |
| Given a WAV file, output a score: a sequence of (onset_frame, symbol_id) | |
| pairs, where each symbol describes a discrete audio event by | |
| (pitch_band, duration_class, amplitude_class). | |
| The symbol alphabet is fixed and unnamed. Two recordings containing the | |
| same pattern produce the same score. You do not need to know what the | |
| pattern means to match it. | |
| Pure stdlib. No dependencies. | |
| Author: zeechimp | |
| License: Apache-2.0 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import wave | |
| from array import array | |
| from collections import Counter | |
| from dataclasses import asdict, dataclass, fields | |
| from pathlib import Path | |
| from typing import Dict, List, Optional, Tuple | |
| # ============================================================================ | |
| # Pitch band grid (log-spaced, Hz) | |
| # ============================================================================ | |
| def _log_bands(lo: float, hi: float, n: int) -> List[float]: | |
| ratio = (hi / lo) ** (1.0 / (n - 1)) | |
| return [lo * (ratio ** i) for i in range(n)] | |
| PITCH_BANDS: List[float] = _log_bands(80.0, 6000.0, 16) | |
| N_BANDS: int = len(PITCH_BANDS) | |
| N_DUR: int = 4 | |
| N_AMP: int = 3 | |
| N_SYMBOLS: int = N_BANDS * N_DUR * N_AMP # 192 | |
| def encode_symbol(band: int, dur: int, amp: int) -> int: | |
| return band * (N_DUR * N_AMP) + dur * N_AMP + amp | |
| def decode_symbol(sid: int) -> Tuple[int, int, int]: | |
| band = sid // (N_DUR * N_AMP) | |
| rem = sid % (N_DUR * N_AMP) | |
| return band, rem // N_AMP, rem % N_AMP | |
| # ============================================================================ | |
| # Config | |
| # ============================================================================ | |
| class HVNoteConfig: | |
| # Framing | |
| frame_ms: float = 20.0 | |
| hop_ms: float = 10.0 | |
| # Onset detection | |
| onset_floor: float = 0.02 | |
| onset_rise_mult: float = 1.5 | |
| # Event segmentation | |
| min_peak_rms: float = 0.02 | |
| max_event_frames: int = 200 | |
| min_event_frames: int = 1 | |
| # Pitch window | |
| pitch_window_s: float = 0.030 | |
| # Quantization | |
| dur_bounds: Tuple[int, int, int] = (3, 8, 20) | |
| amp_bounds: Tuple[float, float] = (0.3, 0.7) | |
| # Corpus | |
| max_files: int = 0 | |
| glob: str = "*.wav" | |
| version: str = "0.1.4" | |
| # ============================================================================ | |
| # Report dataclasses | |
| # ============================================================================ | |
| class AudioEvent: | |
| start_frame: int | |
| end_frame: int | |
| peak_frame: int | |
| start_s: float | |
| end_s: float | |
| peak_s: float | |
| peak_rms: float | |
| pitch_hz: float | |
| band: int | |
| dur_class: int | |
| amp_class: int | |
| symbol: int | |
| def to_dict(self) -> dict: | |
| return { | |
| "start_frame": self.start_frame, | |
| "end_frame": self.end_frame, | |
| "peak_frame": self.peak_frame, | |
| "start_s": round(self.start_s, 4), | |
| "end_s": round(self.end_s, 4), | |
| "peak_s": round(self.peak_s, 4), | |
| "peak_rms": round(self.peak_rms, 6), | |
| "pitch_hz": round(self.pitch_hz, 2), | |
| "band": self.band, | |
| "dur_class": self.dur_class, | |
| "amp_class": self.amp_class, | |
| "symbol": self.symbol, | |
| } | |
| class ScoreReport: | |
| path: str | |
| sample_rate: int | |
| duration_s: float | |
| n_frames: int | |
| n_samples: int | |
| events: List[AudioEvent] | |
| score: List[Tuple[int, int]] | |
| histogram: Dict[int, int] | |
| entropy_bits: float | |
| def to_dict(self) -> dict: | |
| return { | |
| "path": self.path, | |
| "sample_rate": self.sample_rate, | |
| "duration_s": round(self.duration_s, 4), | |
| "n_frames": self.n_frames, | |
| "n_samples": self.n_samples, | |
| "n_events": len(self.events), | |
| "score": [[int(t), int(s)] for t, s in self.score], | |
| "histogram": {str(k): v for k, v in self.histogram.items()}, | |
| "entropy_bits": round(self.entropy_bits, 4), | |
| "events": [e.to_dict() for e in self.events], | |
| } | |
| class CorpusReport: | |
| root: str | |
| n_files: int | |
| n_events: int | |
| unique_symbols: int | |
| total_entropy_bits: float | |
| symbol_counts: Dict[int, int] | |
| group_counts: Dict[str, int] | |
| group_unique: Dict[str, int] | |
| file_scores: Dict[str, List[int]] | |
| def to_dict(self) -> dict: | |
| return { | |
| "root": self.root, | |
| "n_files": self.n_files, | |
| "n_events": self.n_events, | |
| "unique_symbols": self.unique_symbols, | |
| "total_entropy_bits": round(self.total_entropy_bits, 4), | |
| "symbol_counts": {str(k): v for k, v in self.symbol_counts.items()}, | |
| "group_counts": dict(self.group_counts), | |
| "group_unique": dict(self.group_unique), | |
| "file_scores": {k: list(v) for k, v in self.file_scores.items()}, | |
| } | |
| # ============================================================================ | |
| # Audio I/O & primitives | |
| # ============================================================================ | |
| def read_wav(path: str) -> Tuple[array, int]: | |
| """Return (samples_int16, sample_rate). Left channel only.""" | |
| with wave.open(path, "rb") as w: | |
| n_ch = w.getnchannels() | |
| sw = w.getsampwidth() | |
| sr = w.getframerate() | |
| n = w.getnframes() | |
| raw = w.readframes(n) | |
| if sw != 2: | |
| raise ValueError(f"Only 16-bit PCM WAV supported, got {sw*8}-bit") | |
| samples = array("h") | |
| samples.frombytes(raw) | |
| if n_ch > 1: | |
| samples = samples[::n_ch] | |
| return samples, sr | |
| def rms_norm(chunk: array) -> float: | |
| n = len(chunk) | |
| if n == 0: | |
| return 0.0 | |
| s = 0 | |
| for x in chunk: | |
| s += x * x | |
| return math.sqrt(s / n) / 32768.0 | |
| def goertzel_power(window: array, target_hz: float, sr: int) -> float: | |
| n = len(window) | |
| if n == 0 or target_hz >= sr / 2: | |
| return 0.0 | |
| k = int(0.5 + n * target_hz / sr) | |
| if k < 1: | |
| k = 1 | |
| omega = 2.0 * math.pi * k / n | |
| coeff = 2.0 * math.cos(omega) | |
| s_prev = 0.0 | |
| s_prev2 = 0.0 | |
| for x in window: | |
| s = x + coeff * s_prev - s_prev2 | |
| s_prev2 = s_prev | |
| s_prev = s | |
| return s_prev2 * s_prev2 + s_prev * s_prev - coeff * s_prev * s_prev2 | |
| def entropy_bits(counts: Dict[int, int]) -> float: | |
| total = sum(counts.values()) | |
| if total == 0: | |
| return 0.0 | |
| e = 0.0 | |
| for c in counts.values(): | |
| if c > 0: | |
| p = c / total | |
| e -= p * math.log2(p) | |
| return e | |
| # ============================================================================ | |
| # The model | |
| # ============================================================================ | |
| class HVNote: | |
| """Symbolic audio event extraction.""" | |
| def __init__(self, config: Optional[HVNoteConfig] = None): | |
| self.config = config or HVNoteConfig() | |
| self._obs = 0 | |
| def __repr__(self) -> str: | |
| return ( | |
| f"HVNote(frame_ms={self.config.frame_ms}, " | |
| f"hop_ms={self.config.hop_ms}, " | |
| f"n_symbols={N_SYMBOLS}, " | |
| f"version={self.config.version})" | |
| ) | |
| # ------------------------------------------------------------------ | |
| # Framing | |
| # ------------------------------------------------------------------ | |
| def _frame_rms(self, samples: array, sr: int) -> Tuple[List[float], int, int]: | |
| win = max(1, int(sr * self.config.frame_ms / 1000)) | |
| hop = max(1, int(sr * self.config.hop_ms / 1000)) | |
| rms_vals: List[float] = [] | |
| i = 0 | |
| n = len(samples) | |
| while i + win <= n: | |
| rms_vals.append(rms_norm(samples[i:i + win])) | |
| i += hop | |
| return rms_vals, win, hop | |
| # ------------------------------------------------------------------ | |
| # Event detection | |
| # ------------------------------------------------------------------ | |
| def _detect_events( | |
| self, rms_vals: List[float] | |
| ) -> List[Tuple[int, int, int, float]]: | |
| """ | |
| Onset + extent in a single forward scan. | |
| Onset fires at frame i when rms[i] >= onset_floor AND either: | |
| - rms[i] >= rms[i-1] * onset_rise_mult (sharp rise), or | |
| - rms[i-1] < onset_floor (silence crossing). | |
| The event's extent is the connected region above onset_floor | |
| starting at the onset. The peak is argmax over that region only — | |
| never over a fixed search horizon. This prevents a later, louder | |
| event from hijacking the current one. | |
| """ | |
| n = len(rms_vals) | |
| if n == 0: | |
| return [] | |
| floor = self.config.onset_floor | |
| rise_mult = self.config.onset_rise_mult | |
| events: List[Tuple[int, int, int, float]] = [] | |
| i = 0 | |
| while i < n: | |
| cur = rms_vals[i] | |
| if cur < floor: | |
| i += 1 | |
| continue | |
| prev = rms_vals[i - 1] if i > 0 else 0.0 | |
| is_rise = (i == 0) or (prev > 1e-9 and cur >= prev * rise_mult) | |
| is_crossing = prev < floor | |
| if not (is_rise or is_crossing): | |
| i += 1 | |
| continue | |
| # Extent: stay in the connected region above floor. | |
| max_end = min(n, i + self.config.max_event_frames) | |
| end = i | |
| while end < max_end and rms_vals[end] >= floor: | |
| end += 1 | |
| if end - i < self.config.min_event_frames: | |
| i += 1 | |
| continue | |
| # Peak is argmax over the extent only. | |
| peak_frame = max(range(i, end), key=lambda j: rms_vals[j]) | |
| peak_rms = rms_vals[peak_frame] | |
| if peak_rms < self.config.min_peak_rms: | |
| i += 1 | |
| continue | |
| events.append((i, end, peak_frame, peak_rms)) | |
| # Skip past the whole event. | |
| i = max(i + 1, end) | |
| return events | |
| # ------------------------------------------------------------------ | |
| # Pitch band | |
| # ------------------------------------------------------------------ | |
| def _pitch_band( | |
| self, samples: array, sr: int, peak_frame: int, hop: int | |
| ) -> Tuple[int, float]: | |
| window_len = max(1, int(sr * self.config.pitch_window_s)) | |
| center = peak_frame * hop + (hop // 2) | |
| start = max(0, center - window_len // 2) | |
| end = min(len(samples), start + window_len) | |
| window = samples[start:end] | |
| if len(window) < 8: | |
| return 0, 0.0 | |
| best_band = 0 | |
| best_power = -1.0 | |
| for bi, f in enumerate(PITCH_BANDS): | |
| p = goertzel_power(window, f, sr) | |
| if p > best_power: | |
| best_power = p | |
| best_band = bi | |
| return best_band, PITCH_BANDS[best_band] | |
| # ------------------------------------------------------------------ | |
| # Quantization | |
| # ------------------------------------------------------------------ | |
| def _dur_class(self, length: int) -> int: | |
| b0, b1, b2 = self.config.dur_bounds | |
| if length < b0: | |
| return 0 | |
| if length < b1: | |
| return 1 | |
| if length < b2: | |
| return 2 | |
| return 3 | |
| def _amp_class(self, peak_rms: float, recording_peak: float) -> int: | |
| if recording_peak <= 0: | |
| return 0 | |
| r = peak_rms / recording_peak | |
| a0, a1 = self.config.amp_bounds | |
| if r < a0: | |
| return 0 | |
| if r < a1: | |
| return 1 | |
| return 2 | |
| # ------------------------------------------------------------------ | |
| # Public API | |
| # ------------------------------------------------------------------ | |
| def score_file(self, path: str) -> ScoreReport: | |
| samples, sr = read_wav(path) | |
| return self.score_samples(samples, sr, path=path) | |
| def score_samples( | |
| self, samples: array, sr: int, path: str = "<memory>" | |
| ) -> ScoreReport: | |
| rms_vals, win, hop = self._frame_rms(samples, sr) | |
| n_frames = len(rms_vals) | |
| recording_peak = max(rms_vals) if rms_vals else 0.0 | |
| raw_events = self._detect_events(rms_vals) | |
| events: List[AudioEvent] = [] | |
| for start, end, peak_frame, peak_rms in raw_events: | |
| band, hz = self._pitch_band(samples, sr, peak_frame, hop) | |
| dur = self._dur_class(end - start) | |
| amp = self._amp_class(peak_rms, recording_peak) | |
| sym = encode_symbol(band, dur, amp) | |
| events.append(AudioEvent( | |
| start_frame=start, | |
| end_frame=end, | |
| peak_frame=peak_frame, | |
| start_s=start * self.config.hop_ms / 1000.0, | |
| end_s=end * self.config.hop_ms / 1000.0, | |
| peak_s=peak_frame * self.config.hop_ms / 1000.0, | |
| peak_rms=peak_rms, | |
| pitch_hz=hz, | |
| band=band, | |
| dur_class=dur, | |
| amp_class=amp, | |
| symbol=sym, | |
| )) | |
| score = [(e.start_frame, e.symbol) for e in events] | |
| hist = Counter(e.symbol for e in events) | |
| hist_dict = {int(k): int(v) for k, v in hist.items()} | |
| ent = entropy_bits(hist_dict) | |
| dur = len(samples) / sr if sr > 0 else 0.0 | |
| self._obs += 1 | |
| return ScoreReport( | |
| path=path, | |
| sample_rate=sr, | |
| duration_s=dur, | |
| n_frames=n_frames, | |
| n_samples=len(samples), | |
| events=events, | |
| score=score, | |
| histogram=hist_dict, | |
| entropy_bits=ent, | |
| ) | |
| def score_directory( | |
| self, root: str, max_files: Optional[int] = None | |
| ) -> CorpusReport: | |
| root_path = Path(root) | |
| if not root_path.exists(): | |
| raise FileNotFoundError(root) | |
| files = sorted(root_path.rglob(self.config.glob)) | |
| if not files: | |
| raise FileNotFoundError( | |
| f"No files matching {self.config.glob!r} under {root}" | |
| ) | |
| limit = max_files if max_files is not None else self.config.max_files | |
| if limit and limit > 0: | |
| files = files[:limit] | |
| symbol_counts: Counter = Counter() | |
| group_counts: Counter = Counter() | |
| group_symbols: Dict[str, set] = {} | |
| file_scores: Dict[str, List[int]] = {} | |
| total_events = 0 | |
| for f in files: | |
| try: | |
| report = self.score_file(str(f)) | |
| except Exception as e: | |
| print(f" skip {f.name}: {e}", file=sys.stderr) | |
| continue | |
| file_scores[str(f)] = [s for _, s in report.score] | |
| total_events += len(report.score) | |
| group_key = self._group_key(f.stem) | |
| group_counts[group_key] += 1 | |
| group_symbols.setdefault(group_key, set()) | |
| group_symbols[group_key].update(s for _, s in report.score) | |
| symbol_counts.update(s for _, s in report.score) | |
| symbol_counts_dict = {int(k): int(v) for k, v in symbol_counts.items()} | |
| group_unique = {k: len(v) for k, v in group_symbols.items()} | |
| ent = entropy_bits(symbol_counts_dict) | |
| return CorpusReport( | |
| root=str(root), | |
| n_files=len(file_scores), | |
| n_events=total_events, | |
| unique_symbols=len(symbol_counts), | |
| total_entropy_bits=ent, | |
| symbol_counts=symbol_counts_dict, | |
| group_counts={k: int(v) for k, v in group_counts.items()}, | |
| group_unique=group_unique, | |
| file_scores=file_scores, | |
| ) | |
| def _group_key(stem: str) -> str: | |
| for sep in ("_", "-"): | |
| if sep in stem: | |
| return stem.split(sep)[0] | |
| return stem | |
| # ------------------------------------------------------------------ | |
| # Convenience | |
| # ------------------------------------------------------------------ | |
| def score(self, path: str) -> List[int]: | |
| return [s for _, s in self.score_file(path).score] | |
| def fingerprint(self, path: str, k: int = 8) -> Tuple[int, ...]: | |
| s = self.score(path) | |
| return tuple(s[:k]) | |
| # ------------------------------------------------------------------ | |
| # Rendering | |
| # ------------------------------------------------------------------ | |
| def render_file(self, report: ScoreReport) -> str: | |
| lines: List[str] = [] | |
| bar = "=" * 72 | |
| lines.append(bar) | |
| lines.append("hv-note — symbolic audio score") | |
| lines.append(bar) | |
| lines.append("") | |
| lines.append(f" file : {report.path}") | |
| lines.append(f" duration : {report.duration_s:.3f} s") | |
| lines.append(f" sample rate: {report.sample_rate} Hz") | |
| lines.append(f" frames : {report.n_frames}") | |
| lines.append(f" events : {len(report.events)}") | |
| lines.append(f" unique sym : {len(report.histogram)}") | |
| lines.append(f" entropy : {report.entropy_bits:.3f} bits") | |
| lines.append("") | |
| if not report.events: | |
| lines.append(" (no events detected)") | |
| return "\n".join(lines) | |
| lines.append(" SCORE") | |
| lines.append(" " + "-" * 68) | |
| lines.append( | |
| f" {'frame':>6} {'time':>7} {'sym':>4} " | |
| f"(band, dur, amp) {'pitch':>8} rms" | |
| ) | |
| for e in report.events: | |
| lines.append( | |
| f" {e.start_frame:>6} {e.start_s:>6.3f}s " | |
| f"{e.symbol:>4} " | |
| f"({e.band:>2}, {e.dur_class}, {e.amp_class}) " | |
| f"{e.pitch_hz:>7.1f}Hz {e.peak_rms:.4f}" | |
| ) | |
| lines.append("") | |
| lines.append(" HISTOGRAM") | |
| lines.append(" " + "-" * 68) | |
| for sym, count in sorted( | |
| report.histogram.items(), key=lambda kv: -kv[1] | |
| ): | |
| band, dur, amp = decode_symbol(sym) | |
| lines.append( | |
| f" {sym:>4} ({band:>2}, {dur}, {amp}) {count:>4}" | |
| ) | |
| lines.append("") | |
| if len(report.score) >= 2: | |
| lines.append(" TRANSITIONS (top 8)") | |
| lines.append(" " + "-" * 68) | |
| trans: Counter = Counter() | |
| for i in range(len(report.score) - 1): | |
| a = report.score[i][1] | |
| b = report.score[i + 1][1] | |
| trans[(a, b)] += 1 | |
| for (a, b), c in trans.most_common(8): | |
| lines.append(f" {a:>4} -> {b:<4} {c:>4}") | |
| lines.append("") | |
| return "\n".join(lines) | |
| def render_corpus(self, report: CorpusReport) -> str: | |
| lines: List[str] = [] | |
| bar = "=" * 72 | |
| lines.append(bar) | |
| lines.append("hv-note — corpus of symbolic audio scores") | |
| lines.append(bar) | |
| lines.append("") | |
| lines.append(f" root : {report.root}") | |
| lines.append(f" files : {report.n_files}") | |
| lines.append(f" events : {report.n_events}") | |
| lines.append(f" unique sym : {report.unique_symbols} / {N_SYMBOLS}") | |
| lines.append(f" entropy : {report.total_entropy_bits:.3f} bits") | |
| lines.append("") | |
| if report.symbol_counts: | |
| lines.append(" TOP SYMBOLS") | |
| lines.append(" " + "-" * 68) | |
| lines.append( | |
| f" {'sym':>4} (band, dur, amp) count" | |
| ) | |
| for sym, count in sorted( | |
| report.symbol_counts.items(), key=lambda kv: -kv[1] | |
| )[:15]: | |
| band, dur, amp = decode_symbol(sym) | |
| lines.append( | |
| f" {sym:>4} ({band:>2}, {dur}, {amp}) " | |
| f"{count:>6}" | |
| ) | |
| lines.append("") | |
| if report.group_counts: | |
| lines.append(" PER-GROUP SUMMARY") | |
| lines.append(" " + "-" * 68) | |
| lines.append( | |
| f" {'group':<20} {'files':>6} {'unique_sym':>10}" | |
| ) | |
| for g, count in sorted( | |
| report.group_counts.items(), key=lambda kv: -kv[1] | |
| )[:20]: | |
| uniq = report.group_unique.get(g, 0) | |
| lines.append(f" {g:<20} {count:>6} {uniq:>10}") | |
| lines.append("") | |
| return "\n".join(lines) | |
| # ------------------------------------------------------------------ | |
| # Persistence | |
| # ------------------------------------------------------------------ | |
| def save_pretrained(self, save_dir: str) -> None: | |
| os.makedirs(save_dir, exist_ok=True) | |
| payload = { | |
| "config": asdict(self.config), | |
| "observations": self._obs, | |
| "pitch_bands": list(PITCH_BANDS), | |
| "n_symbols": N_SYMBOLS, | |
| } | |
| with open(os.path.join(save_dir, "config.json"), "w") as f: | |
| json.dump(payload, f, indent=2) | |
| def from_pretrained(cls, save_dir: str) -> "HVNote": | |
| with open(os.path.join(save_dir, "config.json"), "r") as f: | |
| payload = json.load(f) | |
| cfg_dict = payload.get("config", {}) | |
| known = {f.name for f in fields(HVNoteConfig)} | |
| cfg_dict = {k: v for k, v in cfg_dict.items() if k in known} | |
| cfg = HVNoteConfig(**cfg_dict) | |
| obj = cls(config=cfg) | |
| obj._obs = int(payload.get("observations", 0)) | |
| return obj | |
| # ============================================================================ | |
| # Demo | |
| # ============================================================================ | |
| def _write_wav(path: str, samples: List[float], sr: int) -> None: | |
| ints = array( | |
| "h", | |
| (max(-32768, min(32767, int(x * 32767))) for x in samples), | |
| ) | |
| with wave.open(path, "wb") as w: | |
| w.setnchannels(1) | |
| w.setsampwidth(2) | |
| w.setframerate(sr) | |
| w.writeframes(ints.tobytes()) | |
| def _synth_beeps(freqs: List[float], dur_ms: int, gap_ms: int, | |
| sr: int = 16000) -> List[float]: | |
| out: List[float] = [] | |
| for f in freqs: | |
| n = int(sr * dur_ms / 1000) | |
| for i in range(n): | |
| t = i / sr | |
| env = math.sin(math.pi * i / n) ** 0.5 | |
| out.append(0.5 * env * math.sin(2 * math.pi * f * t)) | |
| out.extend([0.0] * int(sr * gap_ms / 1000)) | |
| return out | |
| def _synth_chirp(f0: float, f1: float, dur_ms: int, | |
| sr: int = 16000) -> List[float]: | |
| n = int(sr * dur_ms / 1000) | |
| dur_s = dur_ms / 1000.0 | |
| out: List[float] = [] | |
| for i in range(n): | |
| t = i / sr | |
| phase = 2 * math.pi * (f0 * t + 0.5 * (f1 - f0) * t * t / dur_s) | |
| env = math.sin(math.pi * i / n) | |
| out.append(0.5 * env * math.sin(phase)) | |
| return out | |
| def _synth_claps(dur_ms: int, gap_ms: int, seed: int = 42, | |
| sr: int = 16000) -> List[float]: | |
| import random | |
| rng = random.Random(seed) | |
| out: List[float] = [] | |
| for _ in range(4): | |
| n = int(sr * dur_ms / 1000) | |
| for i in range(n): | |
| env = math.exp(-i / (n * 0.3)) | |
| out.append(0.7 * env * (rng.random() * 2 - 1)) | |
| out.extend([0.0] * int(sr * gap_ms / 1000)) | |
| return out | |
| def _demo(output_dir: str = "./hv_note_output") -> None: | |
| os.makedirs(output_dir, exist_ok=True) | |
| m = HVNote() | |
| samples: List[Tuple[str, List[float], int]] = [ | |
| ("beeps_low", _synth_beeps([220, 220, 220], 150, 100), 16000), | |
| ("beeps_high", _synth_beeps([880, 660, 440], 120, 90), 16000), | |
| ("chirp_up", _synth_chirp(200, 2000, 800), 16000), | |
| ("claps", _synth_claps(60, 180), 16000), | |
| ] | |
| print() | |
| print("=" * 72) | |
| print("hv-note — demo on synthesized audio") | |
| print("=" * 72) | |
| print() | |
| paths: List[str] = [] | |
| for name, sig, sr in samples: | |
| path = os.path.join(output_dir, f"{name}.wav") | |
| _write_wav(path, sig, sr) | |
| paths.append(path) | |
| print(f" wrote {path} ({len(sig)} samples @ {sr} Hz)") | |
| print() | |
| for path in paths: | |
| report = m.score_file(path) | |
| print(m.render_file(report)) | |
| print() | |
| print("=" * 72) | |
| print("Summary") | |
| print("=" * 72) | |
| print(f" {'file':<24} {'events':>6} {'unique':>6} {'entropy':>8}") | |
| print(" " + "-" * 60) | |
| for path in paths: | |
| r = m.score_file(path) | |
| print( | |
| f" {Path(path).name:<24} {len(r.events):>6} " | |
| f"{len(r.histogram):>6} {r.entropy_bits:>7.3f}b" | |
| ) | |
| print() | |
| print("=" * 72) | |
| print("Corpus mode (over the demo folder)") | |
| print("=" * 72) | |
| corpus = m.score_directory(output_dir) | |
| print(m.render_corpus(corpus)) | |
| print("=" * 72) | |
| print("Save / load round trip") | |
| print("=" * 72) | |
| save_path = os.path.join(output_dir, "hv_note_model") | |
| m.save_pretrained(save_path) | |
| m2 = HVNote.from_pretrained(save_path) | |
| print(f" saved to : {save_path}") | |
| print(f" reloaded : {m2!r}") | |
| a = m.score(paths[0]) | |
| b = m2.score(paths[0]) | |
| print(f" score [0] : {a}") | |
| print(f" reloaded : {b}") | |
| print(f" identical : {a == b}") | |
| print() | |
| # ============================================================================ | |
| # CLI | |
| # ============================================================================ | |
| def _cli() -> None: | |
| p = argparse.ArgumentParser( | |
| description="hv-note: symbolic audio event extraction." | |
| ) | |
| p.add_argument("--file", type=str, default="", | |
| help="score a single WAV file") | |
| p.add_argument("--dir", type=str, default="", | |
| help="score every WAV under a directory") | |
| p.add_argument("--limit", type=int, default=0, | |
| help="max files in --dir mode (0 = all)") | |
| p.add_argument("--json", action="store_true", | |
| help="output JSON instead of a rendered report") | |
| p.add_argument("--fingerprint", action="store_true", | |
| help="print only the fingerprint of a single file") | |
| p.add_argument("--save-to", type=str, default="", | |
| help="save the model to this directory") | |
| p.add_argument("--outdir", type=str, default="./hv_note_output") | |
| args = p.parse_args() | |
| m = HVNote() | |
| if args.save_to: | |
| m.save_pretrained(args.save_to) | |
| print(f"saved to {args.save_to}", file=sys.stderr) | |
| if args.fingerprint: | |
| if not args.file: | |
| p.error("--fingerprint requires --file") | |
| fp = m.fingerprint(args.file) | |
| print(json.dumps(list(fp))) | |
| return | |
| if args.file: | |
| report = m.score_file(args.file) | |
| if args.json: | |
| print(json.dumps(report.to_dict(), indent=2)) | |
| else: | |
| print(m.render_file(report)) | |
| return | |
| if args.dir: | |
| corpus = m.score_directory(args.dir, max_files=args.limit or None) | |
| if args.json: | |
| print(json.dumps(corpus.to_dict(), indent=2)) | |
| else: | |
| print(m.render_corpus(corpus)) | |
| return | |
| _demo(args.outdir) | |
| if __name__ == "__main__": | |
| _cli() |