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