Spaces:
Sleeping
Sleeping
Download restoflow/refine_residual.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/refine_residual.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/refine_residual.py
-
curl -L -o refine_residual.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/refine_residual.py
10.6 kB
| """SUPERVISED residual refiner. Diagnosis (probe_oracle_resid): the frozen advramp restorer's per-stem | |
| output is OVER-SHARP/artifact-y vs clean; the oracle residual (clean - base) reaches clean exactly, but | |
| the glue GAN's rflow head learned an INERT/artifact residual (latent FM/mmd were satisfied trivially). | |
| Fix: train a deterministic refiner R on the FROZEN advramp base to predict clean, supervised by | |
| L = w_lat * MSE(R(base), clean_std) # anchor to the proven oracle target | |
| + w_stft * MultiResSTFT(decode(R(base)), decode(clean)) # DECODED-domain perceptual shaping | |
| The decoded STFT term is the point: latent-MSE alone is what advramp used -> over-sharp; the perceptual | |
| term is what makes the correction MUSICAL instead of artifact (inference-gain on rflow added artifacts). | |
| R is AttnRestorer (identity-init -> starts == base). Frozen SAME decoder passes grad to R's output. | |
| Smoke: python -m restoflow.refine_residual --smoke --device cpu | |
| """ | |
| from __future__ import annotations | |
| import argparse, math, random, time | |
| from pathlib import Path | |
| import torch, torch.nn as nn, torch.nn.functional as F | |
| from torch.utils.data import DataLoader | |
| from torch.utils.checkpoint import checkpoint | |
| from .config import Cfg, STEMS, STEM_ID | |
| from . import eval as E | |
| from .model import AttnRestorer | |
| from .remix_gan import PairData, EMA | |
| BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") | |
| RESTORE_IDS = [STEM_ID["other"], STEM_ID["vocals"], STEM_ID["drums"]] # {0,1,2}; bass is generation-handled | |
| def _std(x, mu, sd): return (x - mu[:, None]) / sd[:, None] | |
| def _destd(x, mu, sd): return x * sd[:, None] + mu[:, None] | |
| def mrstft(a, b, eps=1e-5): | |
| """Differentiable multi-res STFT loss between [B,2,S] audio: spectral-convergence + log-mag L1.""" | |
| a = a.mean(1); b = b.mean(1) # mono | |
| tot = 0.0 | |
| for nf in (512, 1024, 2048): | |
| win = torch.hann_window(nf, device=a.device) | |
| Sa = torch.stft(a, nf, hop_length=nf // 4, window=win, return_complex=True).abs() | |
| Sb = torch.stft(b, nf, hop_length=nf // 4, window=win, return_complex=True).abs() | |
| sc = torch.linalg.norm(Sb - Sa) / (torch.linalg.norm(Sb) + eps) | |
| mag = (torch.log(Sa + eps) - torch.log(Sb + eps)).abs().mean() | |
| tot = tot + sc + mag | |
| return tot / 3.0 | |
| def load_frozen_restorer(run, dev): | |
| ck = torch.load(Path(run) / "ckpt_best.pt", map_location="cpu"); c = ck["cfg"] | |
| m = AttnRestorer(256, c["hidden"], c["depth"], n_stems=len(STEMS), stem_emb=c["stem_emb_dim"], | |
| heads=c.get("heads", 8), use_mix=c["use_mix"]) | |
| m.load_state_dict(ck["model"]); m.eval().to(dev) | |
| for p in m.parameters(): p.requires_grad_(False) | |
| stats = torch.load(Path(run) / "norm_stats.pt", map_location="cpu") | |
| return m, c, stats | |
| def main(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--restorer", default=str(BASE / "restoflow_runs/restorer_attn_advramp")) | |
| p.add_argument("--cache-roots", default=str(BASE / "demucs_results_all")) | |
| p.add_argument("--out-dir", default=str(BASE / "restoflow_runs/refine_v1")) | |
| p.add_argument("--epochs", type=int, default=12) | |
| p.add_argument("--batch", type=int, default=16) | |
| p.add_argument("--decode-bs", type=int, default=8, help="items decoded per step for the STFT loss (memory bound)") | |
| p.add_argument("--w-lat", type=float, default=1.0) | |
| p.add_argument("--w-stft", type=float, default=1.0) | |
| p.add_argument("--lr", type=float, default=2e-4) | |
| p.add_argument("--clip", type=float, default=1.0, help="grad-norm clip (0 disables)") | |
| p.add_argument("--ema", type=float, default=0.999) | |
| p.add_argument("--hidden", type=int, default=512) | |
| p.add_argument("--depth", type=int, default=8) | |
| p.add_argument("--heads", type=int, default=8) | |
| p.add_argument("--no-mix", action="store_true", help="refiner ignores mix conditioning") | |
| p.add_argument("--num-workers", type=int, default=6) | |
| p.add_argument("--val-pairs", type=int, default=64) | |
| p.add_argument("--device", default="cuda:0") | |
| p.add_argument("--smoke", action="store_true") | |
| a = p.parse_args() | |
| dev = a.device if torch.cuda.is_available() or a.device == "cpu" else "cpu" | |
| out = Path(a.out_dir); out.mkdir(parents=True, exist_ok=True) | |
| cfg = Cfg(); cfg.device = dev | |
| use_mix = not a.no_mix | |
| rest, rc, stats = load_frozen_restorer(a.restorer, dev) | |
| sa, sr = E.load_same(cfg) # frozen SAME-L; grad flows through | |
| refiner = AttnRestorer(256, a.hidden, a.depth, n_stems=len(STEMS), stem_emb=rc["stem_emb_dim"], | |
| heads=a.heads, use_mix=use_mix).to(dev) | |
| opt = torch.optim.AdamW(refiner.parameters(), lr=a.lr, betas=(0.9, 0.99), weight_decay=1e-4) | |
| ema = EMA(list(refiner.parameters()), a.ema) if a.ema > 0 else None | |
| ds = PairData([x for x in a.cache_roots.split(",") if x], cfg, max_pairs=40 if a.smoke else None) | |
| n_val = min(a.val_pairs, len(ds) // 5) if not a.smoke else 4 | |
| val_idx = set(range(n_val)); tr_idx = [i for i in range(len(ds)) if i not in val_idx] | |
| print(f"[refine] {len(ds)} pairs ({len(tr_idx)} train / {n_val} val) use_mix={use_mix} w_lat={a.w_lat} w_stft={a.w_stft}", flush=True) | |
| loader = DataLoader(torch.utils.data.Subset(ds, tr_idx), batch_size=a.batch, shuffle=True, | |
| drop_last=True, num_workers=a.num_workers) | |
| val_loader = DataLoader(torch.utils.data.Subset(ds, list(val_idx)), batch_size=a.batch, num_workers=2) | |
| mu = {s: stats[s]["mu"].to(dev) for s in STEMS}; sd = {s: stats[s]["sd"].to(dev) for s in STEMS} | |
| mm = stats.get("__mix__"); mmu = mm["mu"].to(dev) if mm else None; msd = mm["sd"].to(dev) if mm else None | |
| def gather(batch): | |
| """yield (sid, base_std, clean_std, mix_std) per restoration stem present in the batch.""" | |
| clean = batch["clean"].to(dev); deg = batch["deg"].to(dev); mix = batch["mix"].to(dev) | |
| present = batch["present"].to(dev); role = batch["role"].to(dev) | |
| mix_std_full = _std(mix, mmu, msd) # backbone ALWAYS needs mix (advramp use_mix=True) | |
| for si in RESTORE_IDS: | |
| s = STEMS[si]; m = (role[:, si] == 1) & present[:, si] | |
| if m.sum() == 0: continue | |
| sid = torch.full((int(m.sum()),), si, device=dev, dtype=torch.long) | |
| d_std = _std(deg[m, si], mu[s], sd[s]) | |
| c_std = _std(clean[m, si], mu[s], sd[s]) | |
| mxb = mix_std_full[m] # mix for the frozen backbone (required) | |
| with torch.no_grad(): | |
| base = rest(d_std, sid, mxb) # frozen base, std-space | |
| mx = mxb if use_mix else None # mix for the refiner (optional) | |
| yield si, s, sid, base, c_std, mx | |
| def step_loss(batch, train=True): | |
| L_lat = torch.zeros((), device=dev); n_lat = 0 | |
| dec_pred = []; dec_clean = [] | |
| for si, s, sid, base, c_std, mx in gather(batch): | |
| pred = refiner(base, sid, mx) | |
| L_lat = L_lat + F.mse_loss(pred, c_std); n_lat += 1 | |
| # queue raw latents for decoded STFT | |
| dec_pred.append((_destd(pred, mu[s], sd[s]))) | |
| dec_clean.append((_destd(c_std, mu[s], sd[s]))) | |
| if n_lat == 0: return None | |
| L_lat = L_lat / n_lat | |
| if a.w_stft == 0: # skip decode entirely (fast MSE-only path) | |
| return a.w_lat * L_lat, float(L_lat), 0.0 | |
| # decoded perceptual on a capped random subset across stems | |
| P = torch.cat(dec_pred, 0); C = torch.cat(dec_clean, 0) | |
| k = min(a.decode_bs, P.shape[0]) | |
| idx = torch.randperm(P.shape[0], device=dev)[:k] | |
| # checkpoint the grad-path decode: recompute decoder in backward -> bounds activation memory | |
| a_pred = checkpoint(sa.decode_audio, P[idx], use_reentrant=False) # grad flows to refiner | |
| with torch.no_grad(): | |
| a_clean = sa.decode_audio(C[idx]) | |
| L_stft = mrstft(a_pred, a_clean) | |
| return a.w_lat * L_lat + a.w_stft * L_stft, float(L_lat), float(L_stft) | |
| def validate(): | |
| if ema: raw = ema.state(); ema.copy_to() | |
| refiner.eval(); tot = 0.0; nb = 0 | |
| for batch in val_loader: | |
| for si, s, sid, base, c_std, mx in gather(batch): | |
| pred = refiner(base, sid, mx) | |
| P = _destd(pred, mu[s], sd[s]); C = _destd(c_std, mu[s], sd[s]) | |
| k = min(a.decode_bs, P.shape[0]) | |
| tot += float(mrstft(sa.decode_audio(P[:k]), sa.decode_audio(C[:k]))); nb += 1 | |
| refiner.train() | |
| if ema: ema.restore(raw) | |
| return tot / max(1, nb) | |
| _last_gn = [0.0] | |
| best = float("inf"); total_steps = a.epochs * max(1, len(tr_idx) // a.batch) | |
| print(f"[refine] total_steps~{total_steps}", flush=True) | |
| step = 0 | |
| for ep in range(a.epochs): | |
| for batch in loader: | |
| r = step_loss(batch, train=True) | |
| if r is None: continue | |
| loss, ll, ls = r | |
| opt.zero_grad(set_to_none=True); loss.backward() | |
| gn = torch.nn.utils.clip_grad_norm_(refiner.parameters(), a.clip) if a.clip > 0 else \ | |
| torch.nn.utils.clip_grad_norm_(refiner.parameters(), 1e9) | |
| opt.step() | |
| if step % 50 == 0: _last_gn[0] = float(gn) | |
| if ema: ema.update() | |
| if step % 50 == 0: | |
| print(f"ep{ep} step{step} loss={float(loss):.4f} lat={ll:.4f} stft={ls:.4f} gnorm={_last_gn[0]:.3f}", flush=True) | |
| step += 1 | |
| if a.smoke and step >= 4: break | |
| vm = validate() | |
| tag = " *BEST" if vm < best else "" | |
| print(f"[refine] ep{ep} val_decoded_mrstft={vm:.4f}{tag}", flush=True) | |
| if vm < best: | |
| best = vm | |
| saved = ema.state() if ema else None | |
| if ema: ema.copy_to() | |
| torch.save({"refiner": refiner.state_dict(), | |
| "refiner_cfg": {"hidden": a.hidden, "depth": a.depth, "stem_emb_dim": rc["stem_emb_dim"], | |
| "heads": a.heads, "use_mix": use_mix, "restore_ids": RESTORE_IDS}, | |
| "source_restorer": Path(a.restorer).name, "epoch": ep + 1, "val_mrstft": vm}, out / "ckpt_best.pt") | |
| torch.save(stats, out / "norm_stats.pt") | |
| if ema: ema.restore(saved) | |
| if a.smoke: break | |
| print(f"done. refiner -> {out/'ckpt_best.pt'} (best val_decoded_mrstft={best:.4f})", flush=True) | |
| print("REFINE_DONE") | |
| if __name__ == "__main__": | |
| main() | |