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