File size: 3,040 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
"""Apples-to-apples generator FAD: load saved gen runs and re-run the (now FAD-enabled) eval once
each, on identical val pairs, so per-stem + REMIX latent-FAD is comparable across the capacity ladder
(conv-10M vs xattn-40M vs xattn-63M). Reuses evaluate_gen / evaluate_gen_x.

Run: python -m restoflow.gen_fad_probe --runs gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M --device cuda
"""
from __future__ import annotations
import argparse
from pathlib import Path
import torch
from .config import Cfg, STEMS
from .model import CondFlow, AttnCondFlow, XAttnCondFlow
from . import eval as E
from .gen import index_clips, evaluate_gen, evaluate_gen_x

BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra")


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--runs", default="gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M")
    p.add_argument("--val-cache-roots", default=str(BASE / "demucs_results_val_full"))
    p.add_argument("--device", default="cuda")
    p.add_argument("--eval-pairs", type=int, default=120)
    a = p.parse_args()
    dev = a.device
    cfg = Cfg(device=dev, T=32)
    sa, sr = E.load_same(cfg)
    va_roots = [x for x in a.val_cache_roots.split(",") if x]

    for run in [r for r in a.runs.split(",") if r]:
        ck_path = BASE / "restoflow_runs" / run / "ckpt_best.pt"
        if not ck_path.exists():
            print(f"\n##### {run}: no ckpt_best.pt — skip"); continue
        ck = torch.load(ck_path, map_location="cpu"); ga = ck["args"]
        arch = ga.get("arch", "conv")
        if arch == "xattn":
            m = XAttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS),
                              stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8))
        elif arch == "attn":
            m = AttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS),
                             stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8))
        else:
            m = CondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim)
        m.load_state_dict(ck["model"]); m.eval().to(dev)
        stats = ck["stats"]
        tgt_stems = [s for s in ga.get("target_stems", "bass").split(",") if s]
        sources = tuple(x for x in ga.get("context_sources", "restored,degraded").split(",") if x)
        c = type(cfg)(**{**cfg.__dict__, "sample_steps": ga.get("sample_steps", 40)})
        va = index_clips(va_roots, tgt_stems, sources)
        print(f"\n##### {run}  arch={arch} {ga['hidden']}x{ga['depth']} "
              f"params={sum(p_.numel() for p_ in m.parameters())/1e6:.1f}M  val={len(va)} #####")
        demo = str(BASE / "restoflow_runs" / run / "_fadprobe_demo")
        if arch == "xattn":
            evaluate_gen_x(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), [ga.get("cfg_rescale", 0.0)],
                           0, demo, a.eval_pairs, 0)
        else:
            evaluate_gen(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), 0, demo, a.eval_pairs, 0)


if __name__ == "__main__":
    main()