"""Incrementally FILL latents/restored/{stem}.pt by running the v4b restorer over cached degraded stem latents (+ degraded_mix). Latent->latent, no decode -> fast. Skips existing. These restored latents are the realistic CONTEXT for the generator: at inference we never have clean stems, only v4b restorations. Run: python -m restoflow.fill_restored --root demucs_results_all """ from __future__ import annotations import argparse, csv, os from pathlib import Path import torch from .config import Cfg, STEM_ID, STEMS from .model import DetRestorer, AttnRestorer def main(): ap = argparse.ArgumentParser() ap.add_argument("--root", required=True) ap.add_argument("--ckpt", default="restoflow_runs/restorer_attn_distvar_w003/ckpt_best.pt") ap.add_argument("--stats", default="restoflow_runs/restorer_attn_distvar_w003/norm_stats.pt") ap.add_argument("--device", default="cuda") ap.add_argument("--overwrite", action="store_true", help="regenerate even if restored/{stem}.pt exists") a = ap.parse_args() c = Cfg(); dev = a.device BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") root = BASE / a.root ck = torch.load(BASE / a.ckpt, map_location="cpu"); cfg = ck["cfg"] saved = {k: v for k, v in cfg.items() if k in ("latent_dim", "hidden", "depth", "stem_emb_dim", "use_mix")} if cfg.get("model_kind") == "attn": # advramp/cotrain/w003 are AttnRestorer; v4b was DetRestorer model = AttnRestorer(saved.get("latent_dim", 256), saved["hidden"], saved["depth"], n_stems=len(STEMS), stem_emb=saved["stem_emb_dim"], use_mix=saved["use_mix"], heads=cfg.get("heads", 4)) else: model = DetRestorer(saved.get("latent_dim", 256), saved["hidden"], saved["depth"], n_stems=len(STEMS), stem_emb=saved["stem_emb_dim"], use_mix=saved["use_mix"]) model.load_state_dict(ck["model"]); model.eval().to(dev) stats = torch.load(BASE / a.stats, map_location="cpu") mmu, msd = stats["__mix__"]["mu"], stats["__mix__"]["sd"] print(f"[fill] v4b loaded; root={root}") rows = list(csv.DictReader(open(root / "metadata.csv"))) made = skipped = nomix = 0 with torch.inference_mode(): for r in rows: stem = r.get("stem") if stem not in STEMS or r.get("deg_present") != "1" or not r.get("deg_latent_path"): continue pair = r["pair_id"] out = root / pair / "latents" / "restored" / f"{stem}.pt" if out.exists() and not a.overwrite: skipped += 1; continue mix_p = root / pair / "latents" / "degraded_mix.pt" if not mix_p.exists(): nomix += 1; continue deg = torch.load(root / r["deg_latent_path"], map_location="cpu").float() mix = torch.load(mix_p, map_location="cpu").float() mu, sd = stats[stem]["mu"], stats[stem]["sd"] std_d = ((deg - mu[:, None]) / sd[:, None]).to(dev)[None] std_m = ((mix - mmu[:, None]) / msd[:, None]).to(dev)[None] sid = torch.tensor([STEM_ID[stem]], device=dev) std_r = model(std_d, sid, std_m)[0].cpu() raw_r = std_r * sd[:, None] + mu[:, None] out.parent.mkdir(parents=True, exist_ok=True) torch.save(raw_r.half(), out); made += 1 if made % 2000 == 0: print(f" made={made} skipped={skipped}") print(f"[fill] done. made={made} skipped(existing)={skipped} no_mix={nomix}") if __name__ == "__main__": main()