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