stem-restoration / restoflow /gen_fad_probe.py
soilkon's picture
sync app + active models
af4583e verified
Raw History Blame Contribute Delete
3.04 kB
"""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()