Spaces:
Sleeping
Sleeping
Download restoflow/gen_fad_probe.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 3.04 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/gen_fad_probe.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/gen_fad_probe.py
-
curl -L -o gen_fad_probe.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/gen_fad_probe.py
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() | |