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