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