soilkon's picture
sync app + active models
af4583e verified
Raw History Blame Contribute Delete
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)
@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