"""Run FlashSR (2025, 1-step diffusion-distill SR) over a folder of degraded wavs -> 48kHz restored wavs. RUNS IN THE ISOLATED VENV (/tmp/audiosr_env, which also has diffusers + the editable FlashSR repo). Invoked by eval_full.py via subprocess (its deps clash with the SAME-L env). FlashSR processes fixed 245760-sample (5.12s @ 48kHz) stereo windows -> chunk, pad tail, crop back. Weights live at /tmp/flashsr/ModelWeights (downloaded from HF jakeoneijk/FlashSR_weights). """ import argparse, glob, os import numpy as np import soundfile as sf import torch import librosa from FlashSR.FlashSR import FlashSR SEG = 245760 # 5.12s @ 48kHz, the model's fixed input length def _stereo48(wav, sr): w = np.asarray(wav) w = w.T if (w.ndim == 2 and w.shape[0] > w.shape[1]) else w # -> [ch,T] if w.ndim == 1: w = np.stack([w, w]) if w.shape[0] == 1: w = np.repeat(w, 2, 0) if sr != 48000: w = np.stack([librosa.resample(w[c], orig_sr=sr, target_sr=48000) for c in range(2)]) return w.astype(np.float32) def process_file(model, f, dev): w = _stereo48(*sf.read(f, dtype="float32")); T = w.shape[1]; outs = [] for st in range(0, max(1, T), SEG): seg = w[:, st:st + SEG]; true_len = seg.shape[1] if true_len < SEG: seg = np.pad(seg, ((0, 0), (0, SEG - true_len))) x = torch.from_numpy(seg).float().to(dev) with torch.no_grad(): y = model(x, lowpass_input=False) outs.append(np.asarray(y.cpu())[:, :true_len]) if torch.cuda.is_available(): torch.cuda.empty_cache() return np.concatenate(outs, axis=1) # [2,T] 48kHz def main(): ap = argparse.ArgumentParser() ap.add_argument("--in-dir", required=True); ap.add_argument("--out-dir", required=True) ap.add_argument("--weights", default="/tmp/flashsr/ModelWeights"); ap.add_argument("--device", default="cuda") a = ap.parse_args() os.makedirs(a.out_dir, exist_ok=True) model = FlashSR(f"{a.weights}/student_ldm.pth", f"{a.weights}/sr_vocoder.pth", f"{a.weights}/vae.pth").to(a.device) files = sorted(glob.glob(os.path.join(a.in_dir, "*.wav"))) print(f"[flashsr] {len(files)} files") for i, f in enumerate(files): try: sf.write(os.path.join(a.out_dir, os.path.basename(f)), process_file(model, f, a.device).T, 48000) except Exception as e: print(f"[flashsr] FAIL {os.path.basename(f)}: {type(e).__name__}: {e}") if torch.cuda.is_available(): torch.cuda.empty_cache() if (i + 1) % 5 == 0: print(f"[flashsr] {i+1}/{len(files)}") print("FLASHSR_DONE") if __name__ == "__main__": main()