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