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