"""robustdeg_v1 — realistic lo-fi degradation for restoration training/eval (LOCAL, self-contained). Vendored + trimmed from the dataset generator `create_rebetiko_style_v2.py` (melkor169 stereo_a2sb, READ-ONLY source). We keep ONLY *time-aligned* degradations so the degraded↔clean pair stays sample-aligned (the restorer learns a residual, and LSD/SI-SDR are per-sample) — i.e. we DROP `wow_flutter` (time-warp) and `reverb` (convolution tail). Enabled here (vs the lowpass-only synthetic set): broadband noise=HISS, hum, clicks/crackle, tape saturation, on top of bandwidth loss + piecewise EQ + occasional mono. Operates on [2, T] float32 (our pipeline convention). Deterministic given a seed. CLI (degrade a folder of wavs -> degraded wavs, same basenames): python -m restoflow.robustdeg --in-dir --out-dir [--seed 0] """ from __future__ import annotations import argparse, glob, json, math, os, random from dataclasses import dataclass, asdict from pathlib import Path from typing import Any, Dict, Tuple import numpy as np import soundfile as sf from scipy import signal ROBUSTDEG_VERSION = "robustdeg_v1" # ---------- helpers (faithful to source) ---------- def db_to_lin(db: float) -> float: return 10.0 ** (db / 20.0) def rms(x, eps=1e-12) -> float: return float(np.sqrt(np.mean(np.square(x), dtype=np.float64) + eps)) def as_stereo_nx2(x: np.ndarray) -> np.ndarray: """-> [n,2] (source convention).""" if x.ndim == 1: return np.stack([x, x], axis=-1) if x.ndim == 2 and x.shape[1] == 2: return x if x.ndim == 2 and x.shape[0] == 2: return x.T if x.ndim == 2 and x.shape[1] == 1: return np.concatenate([x, x], axis=1) raise ValueError(f"bad audio shape {x.shape}") def peak_normalize(x, target_peak): p = float(np.max(np.abs(x))) return x if p < 1e-12 else x * (target_peak / p) def match_rms_to_ref(y, ref, eps=1e-9): r = max(rms(ref[:, 0]), eps) / max(rms(y[:, 0]), eps) return (y * r).astype(np.float32), float(r) def resample_poly_stereo(x, sr_in, sr_out): if sr_in == sr_out: return x g = math.gcd(sr_in, sr_out); up, down = sr_out // g, sr_in // g yL = signal.resample_poly(x[:, 0], up, down).astype(np.float32) yR = signal.resample_poly(x[:, 1], up, down).astype(np.float32) n = min(len(yL), len(yR)); return np.stack([yL[:n], yR[:n]], axis=-1) def _butter_lp(sr, cut, order): return signal.butter(order, min(0.999, max(1e-6, cut / (sr / 2.0))), btype="lowpass") def _butter_hp(sr, cut, order): return signal.butter(order, cut, btype="highpass", fs=sr) def _iir(x, b, a): return np.stack([signal.lfilter(b, a, x[:, c]).astype(np.float32) for c in range(2)], axis=-1) def _fftconv(x, ir): return np.stack([signal.fftconvolve(x[:, c], ir, mode="same").astype(np.float32) for c in range(2)], axis=-1) def _colored_noise(n, rng, color): w = np.array([rng.gauss(0, 1) for _ in range(n)], dtype=np.float64) if color == "white": return w.astype(np.float32) W = np.fft.rfft(w); f = np.fft.rfftfreq(n, 1.0); f[0] = f[1] if len(f) > 1 else 1.0 W *= (1.0 / np.sqrt(f)) if color == "pink" else (1.0 / f) y = np.fft.irfft(W, n=n); y /= max(np.std(y), 1e-9); return y.astype(np.float32) # ---------- aligned degradations ---------- def mono_collapse(x): m = 0.5 * (x[:, 0] + x[:, 1]); return np.stack([m, m], axis=-1) def apply_bandwidth_loss(x, sr, rng, lp=(3500., 7000.), order=(3, 5), ds=(16000, 22050, None)): hp = None; u = rng.random() hp = rng.uniform(60., 180.) if u < 0.70 else (rng.uniform(180., 300.) if u < 0.95 else rng.uniform(300., 400.)) lpc = rng.uniform(*lp); od = rng.randint(*order); y = x if hp is not None: b, a = _butter_hp(sr, hp, max(2, od - 1)); y = _iir(y, b, a) b, a = _butter_lp(sr, lpc, od); y = _iir(y, b, a) d = rng.choice(list(ds)) if d is not None and d < sr: y = resample_poly_stereo(resample_poly_stereo(y, sr, d), d, sr) n = min(len(y), len(x)); y = y[:n] # guard resample length drift return y, {"lp_hz": lpc, "hp_hz": hp, "order": od, "downsample_hz": d} def apply_piecewise_filter(x, sr, rng, n_bp=(5, 8), db=(-9., 3.), f_min=40., numtaps=2049): nbp = rng.randint(*n_bp); f_max = min(18000., sr * 0.49) freqs = np.geomspace(f_min, f_max, nbp).astype(np.float64) g = np.array([rng.uniform(*db) for _ in range(nbp)]) g -= np.linspace(0., rng.uniform(6., 18.), nbp) # vintage HF roll-off lf = freqs <= 200.; g[lf] = np.minimum(g[lf], 0.0) - rng.uniform(2., 10.) ff = np.concatenate([[0.], freqs, [sr / 2.]]) gl = db_to_lin(float(g[0])) * np.ones_like(ff); gl[1:-1] = 10 ** (g / 20.); gl[-1] = gl[-2] fir = signal.firwin2(numtaps=numtaps, freq=ff, gain=gl, fs=sr).astype(np.float32) return _fftconv(x, fir), {"n_bp": nbp} def add_broadband_noise(x, sr, rng, snr_db=(8., 30.), colors=("pink", "white", "brown")): n = x.shape[0]; snr = rng.uniform(*snr_db); col = rng.choice(list(colors)) noise = _colored_noise(n, rng, col) noise *= (max(rms(x[:, 0]), 1e-9) / db_to_lin(snr)) / max(rms(noise), 1e-9) return x + noise[:, None], {"snr_db": snr, "color": col} # HISS def add_hum(x, sr, rng, base=50.0, n_harm=(2, 7), hum_db=(-45., -25.)): n = x.shape[0]; t = np.arange(n) / sr; nh = rng.randint(*n_harm); hd = rng.uniform(*hum_db) amp = db_to_lin(hd) * max(rms(x[:, 0]), 1e-6); hum = np.zeros(n) for k in range(1, nh + 1): hum += (1.0 / k) * np.sin(2 * math.pi * base * k * t + rng.uniform(0, 2 * math.pi)) hum = hum / max(np.max(np.abs(hum)), 1e-9) * amp return x + hum[:, None].astype(np.float32), {"base_hz": base, "n_harm": nh, "hum_db": hd} def add_clicks(x, sr, rng, click_rate=(0.2, 1.0), crackle_rate=(8., 35.), click_db=(-28., -18.), crackle_db=(-42., -30.)): n = x.shape[0]; y = x.copy(); seg = max(rms(0.5 * (x[:, 0] + x[:, 1])), 1e-9) def burst(num, lvl, decay_ms): for _ in range(num): idx = rng.randrange(0, n); amp = seg * db_to_lin(rng.uniform(*lvl)); sgn = -1.0 if rng.random() < .5 else 1.0 dn = max(3, int(sr * rng.uniform(*decay_ms) / 1000.)) b = np.array([rng.gauss(0, 1) for _ in range(dn)], dtype=np.float32) * np.exp(-np.linspace(0, 6, dn)) b = np.concatenate([b[:1], np.diff(b)]).astype(np.float32); b = sgn * amp * b / (np.max(np.abs(b)) + 1e-9) e = min(n, idx + dn); b = b[:e - idx] y[idx:e, 0] += b * (1 + rng.uniform(-.12, .12)); y[idx:e, 1] += b * (1 + rng.uniform(-.12, .12)) burst(int(rng.uniform(*click_rate) * n / sr), click_db, (0.6, 3.5)) burst(int(rng.uniform(*crackle_rate) * n / sr), crackle_db, (0.2, 1.2)) return y.astype(np.float32), {"clicks": True} def apply_tape_saturation(x, sr, rng, drive_db=(0., 9.), post_lp=(4500., 12000.)): y = np.tanh(db_to_lin(rng.uniform(*drive_db)) * x).astype(np.float32) b, a = _butter_lp(sr, rng.uniform(*post_lp), 2) return _iir(y, b, a), {"drive_db": True} # ---------- config + chain ---------- @dataclass class RobustDegConfig: final_peak: float = 0.95 p_mono: float = 0.25 p_bandwidth_loss: float = 0.90 p_piecewise_filter: float = 0.75 p_noise: float = 0.60 # HISS (was 0.0 in synthetic_out) p_hum: float = 0.40 # (was 0.0) p_clicks: float = 0.50 # (was 0.0) p_tape_saturation: float = 0.30 # (was 0.0) # wow_flutter / reverb deliberately EXCLUDED (break time-alignment) def degrade(x_2T: np.ndarray, sr: int, rng: random.Random, cfg: RobustDegConfig) -> Tuple[np.ndarray, Dict[str, Any]]: """[2,T] clean -> ([2,T] degraded, meta). Aligned-only chain; length preserved.""" x = as_stereo_nx2(np.asarray(x_2T, np.float32)); T0 = x.shape[0] x = peak_normalize(x, cfg.final_peak); chain = [] if rng.random() < cfg.p_mono: x = mono_collapse(x); chain.append("mono") if rng.random() < cfg.p_bandwidth_loss: x, m = apply_bandwidth_loss(x, sr, rng); chain.append(("bandwidth", m)) if rng.random() < cfg.p_piecewise_filter: x, m = apply_piecewise_filter(x, sr, rng); chain.append(("piecewise", m)) if rng.random() < cfg.p_noise: x, m = add_broadband_noise(x, sr, rng); chain.append(("noise", m)) if rng.random() < cfg.p_hum: x, m = add_hum(x, sr, rng); chain.append(("hum", m)) if rng.random() < cfg.p_clicks: x, m = add_clicks(x, sr, rng); chain.append("clicks") if rng.random() < cfg.p_tape_saturation: x, m = apply_tape_saturation(x, sr, rng); chain.append("tape") x = x[:T0] # enforce exact alignment length if x.shape[0] < T0: x = np.pad(x, ((0, T0 - x.shape[0]), (0, 0))) x = peak_normalize(x, cfg.final_peak).astype(np.float32) return x.T, {"version": ROBUSTDEG_VERSION, "chain": chain} # back to [2,T] def degrade_file(in_path, out_path, cfg, seed): a, sr = sf.read(str(in_path), dtype="float32"); a = a.T if a.ndim == 2 else np.stack([a, a]) deg, meta = degrade(a, sr, random.Random(seed), cfg) Path(out_path).parent.mkdir(parents=True, exist_ok=True) sf.write(str(out_path), deg.T, sr, subtype="FLOAT") return meta def main(): p = argparse.ArgumentParser() p.add_argument("--in-dir", required=True); p.add_argument("--out-dir", required=True) p.add_argument("--seed", type=int, default=0) a = p.parse_args() cfg = RobustDegConfig() files = sorted(glob.glob(os.path.join(a.in_dir, "*.wav"))) print(f"[robustdeg] {ROBUSTDEG_VERSION} cfg={asdict(cfg)}") print(f"[robustdeg] {len(files)} files -> {a.out_dir}") for i, f in enumerate(files): degrade_file(f, os.path.join(a.out_dir, os.path.basename(f)), cfg, a.seed + i) if (i + 1) % 50 == 0: print(f" {i+1}/{len(files)}", flush=True) print("ROBUSTDEG_DONE") if __name__ == "__main__": main()