Spaces:
Sleeping
Sleeping
Download restoflow/robustdeg.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 9.88 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/robustdeg.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/robustdeg.py
-
curl -L -o robustdeg.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/robustdeg.py
9.88 kB
| """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 <clean_wavs> --out-dir <degraded_wavs> [--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 ---------- | |
| 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() | |