Spaces:
Sleeping
Sleeping
File size: 10,635 Bytes
7153194 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 | """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()
|