"""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) @torch.no_grad() 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()