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