Spaces:
Sleeping
Sleeping
File size: 8,298 Bytes
af4583e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | """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
|