Spaces:
Sleeping
Sleeping
Download restoflow/eval.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 8.3 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/eval.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/eval.py
-
curl -L -o eval.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/eval.py
8.3 kB
| """Evaluation: per-stem latent + decoded-audio metrics, and REMIX (sum stems) vs true clean. | |
| Decodes through frozen SAME. Cheap latent metrics (cosine) are the primary signal; | |
| decoded multi-STFT + remix are the honest audio-domain checks. Also dumps demo audio. | |
| """ | |
| from __future__ import annotations | |
| from collections import defaultdict | |
| from pathlib import Path | |
| import warnings | |
| import numpy as np | |
| import torch | |
| import soundfile as sf | |
| from .config import Cfg, STEM_ID | |
| from . import data as D | |
| from . import flow as F | |
| from . import fad | |
| def load_same(cfg: Cfg): | |
| warnings.filterwarnings("ignore") | |
| from stable_audio_tools import get_pretrained_model | |
| sa, sacfg = get_pretrained_model(cfg.model_id) | |
| sa = sa.to(cfg.device).eval() | |
| for p in sa.parameters(): | |
| p.requires_grad_(False) | |
| return sa, int(sacfg.get("sample_rate", 44100)) | |
| def _destd(z_std, stats, stem): # [256,T] standardized -> raw | |
| mu, sd = stats[stem]["mu"], stats[stem]["sd"] | |
| return z_std * sd[:, None].to(z_std.device) + mu[:, None].to(z_std.device) | |
| def _decode(sa, raw_latent): # [256,T] -> stereo [2,S] cpu | |
| a = sa.decode_audio(raw_latent[None].to(next(sa.parameters()).device).float()) | |
| return a[0].clamp(-1, 1).cpu() | |
| def multi_stft(a, b): # mono tensors | |
| tot = 0.0 | |
| for nf in (512, 1024, 2048): | |
| Sa = torch.stft(a, nf, hop_length=nf // 4, return_complex=True).abs() | |
| Sb = torch.stft(b, nf, hop_length=nf // 4, return_complex=True).abs() | |
| tot += (torch.log(Sa + 1e-5) - torch.log(Sb + 1e-5)).abs().mean().item() | |
| return tot / 3 | |
| def _select_pairs(by_pair, budget): | |
| """Stratified pair selection. budget<=0 -> all pairs. Otherwise front-load pairs | |
| that contain rare stems (bass/drums) so every stem gets eval coverage, not just | |
| the head of the list. Order each pair by the rarity of its scarcest stem.""" | |
| pairs = list(by_pair.items()) | |
| if budget is None or budget <= 0 or budget >= len(pairs): | |
| return pairs | |
| freq = defaultdict(int) | |
| for _, its in pairs: | |
| for it in its: | |
| freq[it["stem"]] += 1 | |
| pairs.sort(key=lambda kv: min(freq[it["stem"]] for it in kv[1])) # rarest-stem pairs first | |
| return pairs[:budget] | |
| def evaluate(cfg: Cfg, model, stats, sa, sr, epoch=0): | |
| model.eval() | |
| items = D._index(cfg, cfg.val_cache_roots or cfg.cache_roots, "val") | |
| by_pair = defaultdict(list) | |
| for it in items: | |
| by_pair[it["pair_id"]].append(it) | |
| pairs = _select_pairs(by_pair, cfg.eval_max_pairs) | |
| per = defaultdict(lambda: defaultdict(list)) # stem -> metric -> [vals] | |
| remix = defaultdict(list) | |
| flat = {k: defaultdict(list) for k in ("clean", "in", "out")} # latent-FAD: stem -> [ [256,T] ] | |
| frmx = {"clean": [], "in": [], "out": []} # latent-FAD REMIX (per-pair latent sum) | |
| demo_root = Path(cfg.out_dir) / "demo" / f"epoch_{epoch:04d}" | |
| for pi, (pid, its) in enumerate(pairs): | |
| clean_mix = deg_mix = rest_mix = None | |
| clat = dlat = rlat = None | |
| for it in its: | |
| stem = it["stem"]; sid = torch.tensor([STEM_ID[stem]], device=cfg.device) | |
| raw_c = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) | |
| raw_d = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) | |
| mu, sd = stats[stem]["mu"], stats[stem]["sd"] | |
| std_d = ((raw_d - mu[:, None]) / sd[:, None]).to(cfg.device)[None] | |
| std_c = ((raw_c - mu[:, None]) / sd[:, None]).to(cfg.device)[None] | |
| std_mix = None | |
| if cfg.use_mix: | |
| raw_m = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) | |
| m = stats["__mix__"] | |
| std_mix = ((raw_m - m["mu"][:, None]) / m["sd"][:, None]).to(cfg.device)[None] | |
| if cfg.model_kind in ("det", "attn"): | |
| std_r = model(std_d, sid, std_mix) | |
| elif cfg.model_kind == "flowattn": | |
| std_r = F.sample_mix(model, std_d, std_mix, sid, cfg.sample_steps, cfg.sigma) | |
| else: | |
| std_r = F.sample(model, std_d, sid, cfg.sample_steps, cfg.sigma) | |
| per[stem]["cos_in"].append(F.latent_cos_dist(std_d, std_c).item()) # baseline (degraded) | |
| per[stem]["cos_out"].append(F.latent_cos_dist(std_r, std_c).item()) # restored | |
| per[stem]["relL2_out"].append(F.latent_relL2(std_r, std_c).item()) | |
| per[stem]["energy_out"].append(F.energy_ratio(std_r, std_c).item()) | |
| raw_r = _destd(std_r[0], stats, stem) | |
| rr = raw_r.detach().cpu() # latent-FAD (raw latents, no decode) | |
| flat["clean"][stem].append(raw_c); flat["in"][stem].append(raw_d); flat["out"][stem].append(rr) | |
| clat = raw_c if clat is None else clat + raw_c | |
| dlat = raw_d if dlat is None else dlat + raw_d | |
| rlat = rr if rlat is None else rlat + rr | |
| ac = _decode(sa, raw_c); ad = _decode(sa, raw_d); ar = _decode(sa, raw_r) | |
| per[stem]["stft_in"].append(multi_stft(ad.mean(0), ac.mean(0))) | |
| per[stem]["stft_out"].append(multi_stft(ar.mean(0), ac.mean(0))) | |
| clean_mix = ac if clean_mix is None else clean_mix + ac | |
| deg_mix = ad if deg_mix is None else deg_mix + ad | |
| rest_mix = ar if rest_mix is None else rest_mix + ar | |
| if pi < cfg.demo_pairs: | |
| d = demo_root / pid; d.mkdir(parents=True, exist_ok=True) | |
| sf.write(d / f"{stem}_1_degraded.wav", ad.T.numpy(), sr) | |
| sf.write(d / f"{stem}_2_clean.wav", ac.T.numpy(), sr) | |
| sf.write(d / f"{stem}_3_restored.wav", ar.T.numpy(), sr) | |
| if clat is not None: | |
| frmx["clean"].append(clat); frmx["in"].append(dlat); frmx["out"].append(rlat) | |
| if clean_mix is not None: | |
| remix["stft_in"].append(multi_stft(deg_mix.mean(0), clean_mix.mean(0))) | |
| remix["stft_out"].append(multi_stft(rest_mix.mean(0), clean_mix.mean(0))) | |
| if pi < cfg.demo_pairs: | |
| d = demo_root / pid | |
| sf.write(d / "MIX_1_degraded.wav", deg_mix.T.numpy(), sr) | |
| sf.write(d / "MIX_2_clean.wav", clean_mix.T.numpy(), sr) | |
| sf.write(d / "MIX_3_restored.wav", rest_mix.T.numpy(), sr) | |
| # report | |
| print(f"\n=== eval epoch {epoch} (val pairs={len(pairs)}) ===") | |
| print(f"{'stem':7} {'n':>4} | cos_in->out | relL2 | energy | stft_in->out") | |
| out = {} | |
| for s in cfg.stems: | |
| if not per[s]["cos_out"]: | |
| continue | |
| m = {k: float(np.mean(v)) for k, v in per[s].items()} | |
| out[s] = m | |
| print(f"{s:7} {len(per[s]['cos_out']):>4} | {m['cos_in']:.3f}->{m['cos_out']:.3f} | " | |
| f"{m['relL2_out']:.3f} | {m['energy_out']:.3f} | {m['stft_in']:.3f}->{m['stft_out']:.3f}") | |
| if remix["stft_out"]: | |
| ri, ro = float(np.mean(remix["stft_in"])), float(np.mean(remix["stft_out"])) | |
| out["remix"] = {"stft_in": ri, "stft_out": ro} | |
| print(f"{'REMIX':7} {len(remix['stft_out']):>4} | full-mix multi-STFT degraded={ri:.3f} -> restored={ro:.3f}") | |
| # latent-FAD: distributional distance to clean (no decode) — sees realism/dullness the | |
| # reference-matching multi-STFT cannot. degraded(in) -> restored(out), lower = closer to clean. | |
| try: | |
| print(f"{'stem':7} | latent-FAD in->out (lower=closer)") | |
| for s in cfg.stems: | |
| if not flat["out"][s]: | |
| continue | |
| cl = fad.latent_set(flat["clean"][s]) | |
| fi = fad.fad(cl, fad.latent_set(flat["in"][s])); fo = fad.fad(cl, fad.latent_set(flat["out"][s])) | |
| out.setdefault(s, {}).update(fad_in=fi, fad_out=fo) | |
| print(f"{s:7} | {fi:.3f}->{fo:.3f}") | |
| if frmx["out"]: | |
| clr = fad.latent_set(frmx["clean"]) | |
| fi = fad.fad(clr, fad.latent_set(frmx["in"])); fo = fad.fad(clr, fad.latent_set(frmx["out"])) | |
| out["fad_remix"] = {"in": fi, "out": fo} | |
| print(f"{'FAD-REMIX':12} | {fi:.3f}->{fo:.3f} (distributional; complements REMIX-STFT)") | |
| except Exception as e: | |
| print(f"[eval] latent-FAD skipped: {e}") | |
| print(f"(demo audio -> {demo_root})") | |
| return out | |