"""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) @torch.inference_mode() 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