"""FINAL stage: remix-coupled, recon+distribution-ANCHORED, LIGHT-adversarial fine-tune. Restorer fixes restoration-class stems, generator fills generation-class stems, sum -> reconstructed mix latent. A latent discriminator judges real clean-mix vs reconstructed-mix. BOTH models are the "generators" co-trained to make a coherent realistic full mix. De-risked by construction: - recon (restorer MSE) + var-match distribution anchor DOMINATE; adversarial is a small ramped aux with a kill-switch -> worst case == the non-adversarial model, never "diverged garbage". - feature-matching (stable) + hinge + R1 grad penalty + warm-started D + grad-accum for OOM. Latent-domain (no decode) = cheap, CPU-smokeable. NOT launched automatically. Smoke: python -m restoflow.remix_gan --smoke --device cpu """ from __future__ import annotations import argparse, csv, math, random, time from collections import defaultdict from pathlib import Path import torch, torch.nn as nn, torch.nn.functional as F from torch.utils.data import Dataset, DataLoader from .config import Cfg, STEMS, STEM_ID from . import data as D, router, distloss as DL from .model import (AttnRestorer, DetRestorer, CondFlow, AttnCondFlow, LatentDiscriminator, MultiScaleLatentDiscriminator) BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") # ---------------- pair-grouped data ---------------- class PairData(Dataset): """Per pair: clean[4], deg[4] (raw latents, 0 if absent), mix, present[4], role[4] (0 none / 1 restoration / 2 generation).""" def __init__(self, roots, cfg, max_pairs=None): self.cfg = cfg; self.T = cfg.T rows = defaultdict(dict) for root in roots: meta = Path(root) / "metadata.csv" if not meta.exists(): continue for r in csv.DictReader(open(meta)): if r.get("stem") not in STEMS: continue rows[(root, r["pair_id"])][r["stem"]] = r self.items = [] for (root, pid), stemrows in rows.items(): mix = Path(root) / pid / "latents" / "degraded_mix.pt" if not mix.exists(): continue self.items.append((root, pid, stemrows, str(mix))) random.Random(0).shuffle(self.items) if max_pairs: self.items = self.items[:max_pairs] def __len__(self): return len(self.items) def __getitem__(self, i): root, pid, stemrows, mixp = self.items[i] T = self.T clean = torch.zeros(len(STEMS), 256, T); deg = torch.zeros(len(STEMS), 256, T) present = torch.zeros(len(STEMS), dtype=torch.bool); role = torch.zeros(len(STEMS), dtype=torch.long) for s, r in stemrows.items(): si = STEM_ID[s] cl = r.get("clean_latent_path") or "" if cl: p = Path(root) / cl if p.exists(): clean[si] = D._fit_T(torch.load(p, map_location="cpu").float(), T); present[si] = True if not present[si]: continue if router.is_restoration(r, self.cfg.rms_thr, self.cfg.peak_thr, self.cfg.max_drop_db): dp = r.get("deg_latent_path") or "" if dp and (Path(root) / dp).exists(): deg[si] = D._fit_T(torch.load(Path(root) / dp, map_location="cpu").float(), T) role[si] = 1 # restoration else: role[si] = 2 # generation (clean present, deg gutted) mix = D._fit_T(torch.load(mixp, map_location="cpu").float(), T) return {"clean": clean, "deg": deg, "mix": mix, "present": present, "role": role} # ---------------- helpers ---------------- def _std(x, mu, sd): return (x - mu[:, None]) / sd[:, None] def _destd(x, mu, sd): return x * sd[:, None] + mu[:, None] def load_restorer(run, dev): ck = torch.load(run / "ckpt_best.pt", map_location="cpu"); c = ck["cfg"] cls = AttnRestorer if c.get("model_kind") == "attn" else DetRestorer kw = dict(stem_emb=c["stem_emb_dim"], use_mix=c["use_mix"]) if c.get("model_kind") == "attn": kw["heads"] = c.get("heads", 4) m = cls(256, c["hidden"], c["depth"], n_stems=len(STEMS), **kw).to(dev) m.load_state_dict(ck["model"]) return m, torch.load(run / "norm_stats.pt", map_location="cpu"), c def load_generator(run, dev): ck = torch.load(run / "ckpt_best.pt", map_location="cpu"); a = ck["args"] cls = AttnCondFlow if a.get("arch") == "attn" else CondFlow kw = dict(stem_emb=Cfg().stem_emb_dim) if a.get("arch") == "attn": kw["heads"] = a.get("heads", 8) m = cls(256, a["hidden"], a["depth"], n_stems=len(STEMS), **kw).to(dev) m.load_state_dict(ck["model"]) return m, ck["stats"], a def sample_grad(gen, ctx, sid, steps): """Few-step Euler flow sampler WITH gradients (training-time; CFG off).""" z = torch.randn_like(ctx); dt = 1.0 / steps for i in range(steps): t = torch.full((ctx.shape[0],), i * dt, device=ctx.device) z = z + dt * gen(z, t, ctx, sid) return z def r1_penalty(D, real): """Gradient penalty, MEAN over latent dims (not sum). Sum-over-8192-dims made this ~O(80) and, with the old r1=1.0 every step, it dwarfed the hinge (~2) -> D collapsed to a constant (Dacc 0.50). Mean-normalization makes it scale-invariant (~O(grad_per_elem^2)); spectral-norm already bounds D's Lipschitz so this is a light touch, applied lazily.""" real = real.detach().requires_grad_(True) out = D(real).sum() g, = torch.autograd.grad(out, real, create_graph=True) return g.pow(2).flatten(1).mean(1).mean() class EMA: """Exponential moving average of params (shadow weights). Standard in AudioSR/diffusion/GAN — the EMA weights are smoother and generalize better than the raw last-step weights, esp. on small data.""" def __init__(self, params, decay): self.decay = decay self.shadow = [p.detach().clone() for p in params] self.params = list(params) @torch.no_grad() def update(self): for s, p in zip(self.shadow, self.params): s.mul_(self.decay).add_(p.detach(), alpha=1 - self.decay) def copy_to(self, params=None): # write EMA weights into the live params (for eval/save) for s, p in zip(self.shadow, params or self.params): p.data.copy_(s) def state(self): # raw weights, to restore after an EMA save return [p.detach().clone() for p in self.params] @torch.no_grad() def restore(self, saved): for p, q in zip(self.params, saved): p.data.copy_(q) def warmup_cosine(step, warmup, total): if step < warmup: return step / max(1, warmup) prog = (step - warmup) / max(1, total - warmup) return 0.5 * (1 + math.cos(math.pi * min(1.0, prog))) def aug_latent(deg_std, noise, tmask): """SpecAugment-style latent regularizer on the standardized degraded input: gaussian noise + a random contiguous time-mask (zeroed frames). Cheap effective-data multiplier for the small (~20h) cache.""" x = deg_std + noise * torch.randn_like(deg_std) if noise > 0 else deg_std if tmask > 0 and x.shape[-1] > tmask + 1: w = int(torch.randint(0, tmask + 1, (1,)).item()) if w > 0: s = int(torch.randint(0, x.shape[-1] - w, (1,)).item()) x = x.clone(); x[..., s:s + w] = 0.0 return x # ---------------- train ---------------- def main(): p = argparse.ArgumentParser() p.add_argument("--restorer", default=str(BASE / "restoflow_runs/restorer_attn_v1")) p.add_argument("--generator", default=str(BASE / "restoflow_runs_archive/gen_v2")) p.add_argument("--out-dir", default=str(BASE / "restoflow_runs/remix_gan_v1")) p.add_argument("--cache-roots", default=str(BASE / "demucs_results_all")) p.add_argument("--device", default="cuda") p.add_argument("--epochs", type=int, default=20) p.add_argument("--batch", type=int, default=16) p.add_argument("--accum", type=int, default=4, help="grad-accum micro-batches (OOM -> small batch, big effective)") p.add_argument("--lr-g", type=float, default=2e-5); p.add_argument("--lr-d", type=float, default=1e-4) # TTUR p.add_argument("--gen-steps", type=int, default=2) p.add_argument("--w-recon", type=float, default=1.0) p.add_argument("--w-dist", type=float, default=0.02) p.add_argument("--w-fm", type=float, default=2.0) # feature-matching = the stable workhorse p.add_argument("--w-adv", type=float, default=0.05) # SMALL adversarial (ramped) p.add_argument("--adv-warmup", type=int, default=200) # D-only warmup steps before adv on G p.add_argument("--train-gen", action="store_true", help="co-train the generator too (risky: adv-through-sampler). Default: generator " "FROZEN (sampled no-grad), only restorer+D adapt first — the safe stage.") p.add_argument("--freeze-restorer", action="store_true", help="FREEZE the restorer backbone (e.g. advramp) as the stable clean base; only the " "generator + residual-flow detail head finetune & glue. Avoids re-introducing the " "standalone-flow harshness while the glue critic adds coherent detail.") p.add_argument("--gen-restore", action="store_true", help="GENERATIVE-RESIDUAL restorer (Exp2): det backbone + a flow-matching residual that " "SAMPLES the missing detail on content-poor stems (AudioSR/BABE-2/UniverSR: generate, " "don't average -> kills the mean-collapse hiss). Bass stays deterministic.") p.add_argument("--restore-stems", default="drums,other,vocals", help="stems that get the generative residual (content-poor); bass excluded (content present).") p.add_argument("--multiscale-d", action="store_true", help="GATED anti-hiss ablation (Exp3): multi-scale mix-D (EnCodec/DAC/BigVGAN). Off by default.") p.add_argument("--w-rflow", type=float, default=1.0, help="flow-matching weight for the residual flow") p.add_argument("--rflow-warmup-steps", type=int, default=0, help="train ONLY backbone-recon + residual-flow (no D, no adv) for N steps so the residual " "isn't trivially-detectable junk; THEN start the mix-D coupling. Without this the fresh " "residual pins Dacc~0.97 and the kill-switch holds adv off forever (gen-restore never " "gets adversarial shaping).") p.add_argument("--rflow-hidden", type=int, default=512, help="residual-flow capacity (bigger invention head)") p.add_argument("--rflow-depth", type=int, default=8) # --- adapted regime (small-data -> regularize + train long, vs naive bigger+more-epochs overfit) --- p.add_argument("--ema", type=float, default=0.999, help="EMA decay for G params (0 disables). Eval/save EMA. AudioSR/diffusion-std.") p.add_argument("--warmup-steps", type=int, default=500, help="linear LR warmup, then cosine decay") p.add_argument("--aug", action="store_true", help="latent augmentation: SpecAugment time/feat mask + input noise on deg (small-data regularizer)") p.add_argument("--aug-noise", type=float, default=0.03, help="gaussian noise std on standardized deg input") p.add_argument("--aug-tmask", type=int, default=3, help="max contiguous frames masked (SpecAugment time)") p.add_argument("--r1", type=float, default=0.1) # mean-normalized R1 (spectral-norm already bounds D) p.add_argument("--d-reg-every", type=int, default=16) # lazy R1: every N steps (cost; magnitude is fixed by mean-norm) p.add_argument("--num-workers", type=int, default=4) p.add_argument("--smoke", action="store_true") a = p.parse_args() if a.smoke: a.epochs, a.batch, a.accum, a.num_workers, a.adv_warmup = 1, 3, 1, 0, 2 dev = a.device out = Path(a.out_dir); out.mkdir(parents=True, exist_ok=True) cfg = Cfg(device=dev, T=32) rest, rstats, rcfg = load_restorer(Path(a.restorer), dev) gen, gstats, gargs = load_generator(Path(a.generator), dev) disc = (MultiScaleLatentDiscriminator() if a.multiscale_d else LatentDiscriminator()).to(dev) # residual flow (Exp2): conditions on the det restoration, samples the missing detail per content-poor stem. restore_ids = {STEMS.index(s) for s in a.restore_stems.split(",") if s in STEMS} if a.gen_restore else set() rflow = None if a.gen_restore: rflow = CondFlow(256, a.rflow_hidden, a.rflow_depth, n_stems=len(STEMS), stem_emb=rcfg["stem_emb_dim"]).to(dev) print(f"[gan] restorer {sum(p.numel() for p in rest.parameters())/1e6:.1f}M " f"generator {sum(p.numel() for p in gen.parameters())/1e6:.1f}M disc {disc.num_params()/1e6:.1f}M" + (f" rflow {rflow.num_params()/1e6:.1f}M on {sorted(restore_ids)}" if rflow is not None else "")) rmu = {s: rstats[s]["mu"].to(dev) for s in STEMS}; rsd = {s: rstats[s]["sd"].to(dev) for s in STEMS} mmu, msd = rstats["__mix__"]["mu"].to(dev), rstats["__mix__"]["sd"].to(dev) gtm = {s: gstats["tgt"][s][0].to(dev) for s in STEMS}; gts = {s: gstats["tgt"][s][1].to(dev) for s in STEMS} gcm = {s: gstats["ctx"][s][0].to(dev) for s in STEMS}; gcs = {s: gstats["ctx"][s][1].to(dev) for s in STEMS} ds = PairData([x for x in a.cache_roots.split(",") if x], cfg, max_pairs=60 if a.smoke else None) loader = DataLoader(ds, batch_size=a.batch, shuffle=True, drop_last=True, num_workers=a.num_workers) print(f"[gan] {len(ds)} pairs") gparams = ((list(rest.parameters()) if not a.freeze_restorer else []) + (list(gen.parameters()) if a.train_gen else []) + (list(rflow.parameters()) if rflow is not None else [])) if a.freeze_restorer: rest.eval() for p_ in rest.parameters(): p_.requires_grad_(False) print("[gan] restorer FROZEN (stable base); training generator+rflow+D only") if not a.train_gen: gen.eval() for p_ in gen.parameters(): p_.requires_grad_(False) assert gparams, "no trainable G params — need --train-gen and/or --gen-restore when --freeze-restorer" optG = torch.optim.AdamW(gparams, a.lr_g, betas=(0.5, 0.9)) optD = torch.optim.AdamW(disc.parameters(), a.lr_d, betas=(0.5, 0.9)) total_steps = a.epochs * max(1, len(ds) // a.batch) schG = torch.optim.lr_scheduler.LambdaLR(optG, lambda s: warmup_cosine(s, a.warmup_steps, total_steps)) schD = torch.optim.lr_scheduler.LambdaLR(optD, lambda s: warmup_cosine(s, a.warmup_steps, total_steps)) ema = EMA(gparams, a.ema) if a.ema > 0 else None print(f"[gan] total_steps~{total_steps} warmup={a.warmup_steps} ema={a.ema} aug={a.aug}") def build_remix(b): """-> (recon_mix[B,256,T], real_mix[B,256,T], recon_loss). Restorer for role1, generator role2.""" clean = b["clean"].to(dev); deg = b["deg"].to(dev); mix = b["mix"].to(dev) present = b["present"].to(dev); role = b["role"].to(dev) B = clean.shape[0] std_mix = _std(mix, mmu, msd) restored = torch.zeros_like(clean); recon_l = clean.new_zeros(()); rflow_fm = clean.new_zeros(()) nrec = 0; nrf = 0 for si, s in enumerate(STEMS): m = role[:, si] == 1 if m.any(): sid = torch.full((int(m.sum()),), si, device=dev) sd_in = _std(deg[m, si], rmu[s], rsd[s]) if a.aug: sd_in = aug_latent(sd_in, a.aug_noise, a.aug_tmask) out = rest(sd_in, sid, std_mix[m]) # deterministic backbone (std) clean_std = _std(clean[m, si], rmu[s], rsd[s]) recon_l = recon_l + F.mse_loss(out, clean_std); nrec += 1 if rflow is not None and si in restore_ids: # generative residual: flow SAMPLES the missing detail r=clean-backbone, conditioned # on the backbone (AudioSR/BABE-2: generate the gap, don't average it away). ctx = out.detach() r = (clean_std - ctx).detach() x0 = torch.randn_like(r); t = torch.rand(r.shape[0], device=dev) z = (1 - t[:, None, None]) * x0 + t[:, None, None] * r rflow_fm = rflow_fm + F.mse_loss(rflow(z, t, ctx, sid), r - x0); nrf += 1 final = out + sample_grad(rflow, ctx, sid, a.gen_steps) # backbone + sampled residual else: final = out # bass / non-restore stems: det only restored[m, si] = _destd(final, rmu[s], rsd[s]) recon_l = recon_l / max(nrec, 1); rflow_fm = rflow_fm / max(nrf, 1) # generation stems: context = sum of restored present-others, normalize by gen ctx stats. # ALSO compute the generator's own FM loss toward the clean target = ANCHOR (keeps the # generator a valid flow under adversarial pressure; prevents the drift/explosion). gen_out = torch.zeros_like(clean); gen_fm = clean.new_zeros(()); ngen = 0 for si, s in enumerate(STEMS): m = role[:, si] == 2 if m.any(): others = present.clone(); others[:, si] = False ctx_raw = (restored * others[:, :, None, None].float())[m].sum(1) # [n,256,T] sum others sid = torch.full((int(m.sum()),), si, device=dev) ctx = _std(ctx_raw, gcm[s], gcs[s]) if a.train_gen: # co-train: FM anchor + grad sampler tgt = _std(clean[m, si], gtm[s], gts[s]) x0 = torch.randn_like(tgt); t = torch.rand(tgt.shape[0], device=dev) z = (1 - t[:, None, None]) * x0 + t[:, None, None] * tgt gen_fm = gen_fm + F.mse_loss(gen(z, t, ctx, sid), tgt - x0); ngen += 1 g = sample_grad(gen, ctx, sid, a.gen_steps) else: # frozen generator (safe stage) with torch.no_grad(): g = sample_grad(gen, ctx, sid, a.gen_steps) gen_out[m, si] = _destd(g, gtm[s], gts[s]) gen_fm = gen_fm / max(ngen, 1) use = torch.where((role == 2)[:, :, None, None], gen_out, torch.where((role == 1)[:, :, None, None], restored, clean)) use = use * present[:, :, None, None].float() recon_mix = use.sum(1) real_mix = (clean * present[:, :, None, None].float()).sum(1) return recon_mix, real_mix, recon_l, gen_fm, rflow_fm step = 0 for ep in range(a.epochs): for b in loader: recon_mix, real_mix, recon_l, gen_fm, rflow_fm = build_remix(b) real_in = _std(real_mix, mmu, msd); recon_in = _std(recon_mix, mmu, msd) # standardize D inputs warming = step < a.rflow_warmup_steps # rflow-only phase: no D, no adv (residual matures first) # ---- D step (skipped during rflow warmup so D doesn't get a head start on junk residuals) ---- if not warming: optD.zero_grad() d_real = disc(real_in); d_fake = disc(recon_in.detach()) lossD = F.relu(1 - d_real).mean() + F.relu(1 + d_fake).mean() if a.r1 > 0 and step % a.d_reg_every == 0: lossD = lossD + a.r1 * r1_penalty(disc, real_in) if torch.isfinite(lossD): lossD.backward() torch.nn.utils.clip_grad_norm_(disc.parameters(), 1.0) optD.step() else: d_real = recon_in.new_zeros(1); d_fake = recon_in.new_zeros(1); lossD = recon_in.new_zeros(1) # ---- G step ---- optG.zero_grad() d_fake_logit, f_fake = disc(_std(recon_mix, mmu, msd), return_feats=True) with torch.no_grad(): _, f_real = disc(real_in, return_feats=True) fm = sum(F.l1_loss(ff, fr) for ff, fr in zip(f_fake, f_real)) / len(f_fake) dist = DL.mmd_loss(recon_mix, real_mix) # bounded -> stable mix anchor adv = -d_fake_logit.mean() dacc = ((d_real > 0).float().mean() + (d_fake < 0).float().mean()).item() / 2 # KILL-SWITCH: adversarial only after warmup AND while D isn't dominating (Dacc<0.8). # When D wins, adv/fm gradients blow up -> drop them; recon+gen_fm+mmd anchors carry on. # GLUE TERM (HiFi-GAN/MelGAN): feature-matching is the STABLE workhorse — bounded L1 on D feats, # does NOT blow up when D wins -> keep it ON after warmup regardless of Dacc. Only the raw # adversarial logit (which explodes when D dominates) is gated by Dacc<0.8. This is the fix for # remix_gan_v1, where the kill-switch killed BOTH -> glue never engaged -> result == advramp. past_warm = (not warming) and step >= (a.rflow_warmup_steps + a.adv_warmup) adv_on = past_warm and dacc < 0.8 w_adv = a.w_adv if adv_on else 0.0 w_fm = a.w_fm if past_warm else 0.0 lossG = (a.w_recon * (recon_l + gen_fm) + a.w_rflow * rflow_fm + a.w_dist * dist + w_fm * fm + w_adv * adv) if torch.isfinite(lossG): try: lossG.backward() torch.nn.utils.clip_grad_norm_(gparams, 1.0) optG.step() if ema is not None: ema.update() except RuntimeError as e: # transient autograd glitch ("tensor does not have optG.zero_grad(set_to_none=True) # a device") -> skip this step, don't kill the run print(f" [skip] G backward RuntimeError @ step {step}: {str(e)[:80]}") else: print(f" [skip] non-finite G loss @ step {step}") schG.step(); schD.step() step += 1 if step % (1 if a.smoke else 50) == 0: print(f"ep{ep} step{step} recon={recon_l.item():.4f} genfm={gen_fm.item():.4f} " f"rflow={rflow_fm.item():.4f} " f"mmd={dist.item():.4f} fm={fm.item():.4f} adv={adv.item():.3f} | " f"D={lossD.item():.3f} Dacc={dacc:.2f} dR={d_real.mean().item():+.2f} " f"dF={d_fake.mean().item():+.2f} w_adv={w_adv}") # save after each epoch in the STANDARD restorer run-dir schema (model+cfg+norm_stats), so the # adversarially-refined restorer is a drop-in for quality_probe / select_best / app.py — plus # the discriminator for resuming. Recoverable + directly FAD-comparable vs the source restorer. saved = ema.state() if ema is not None else None # swap in EMA weights for the saved ckpt if ema is not None: ema.copy_to() ckpt = {"model": rest.state_dict(), "cfg": rcfg, "epoch": ep + 1, "disc": disc.state_dict(), "gan_args": vars(a), "source_restorer": Path(a.restorer).name} if rflow is not None: # generative-residual head (Exp2): det + sampled residual ckpt["rflow"] = rflow.state_dict() ckpt["rflow_cfg"] = {"hidden": a.rflow_hidden, "depth": a.rflow_depth, "stem_emb": rcfg["stem_emb_dim"], "restore_ids": sorted(restore_ids), "gen_steps": a.gen_steps} torch.save(ckpt, out / "ckpt_best.pt") torch.save(rstats, out / "norm_stats.pt") if a.train_gen: # capture the CO-TRAINED generator as a standalone gout = out.parent / f"{out.name}_gen" # generator run-dir (loadable by gen eval / app) gout.mkdir(parents=True, exist_ok=True) torch.save({"model": gen.state_dict(), "args": gargs, "stats": gstats, # EMA weights (still copied in) "epoch": ep + 1, "co_trained_with": Path(a.restorer).name}, gout / "ckpt_best.pt") if ema is not None: ema.restore(saved) # restore raw weights to continue training print(f"done. refined restorer -> {out/'ckpt_best.pt'} (drop-in run-dir; FAD-compare vs " f"{Path(a.restorer).name} via quality_probe --runs). " f"kill-switch = drop w_adv if Dacc->1 (D wins) or quality/FAD regress.") if __name__ == "__main__": main()