stem-restoration / restoflow /robustdeg.py
soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
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 ----------
@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()