stem-restoration / restoflow /stem_gluer.py
soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
10.6 kB
"""STEM GLUER (audio-domain, stem-aware). The pipeline restores/decodes each stem independently and SUMS
them -> the mix sounds un-glued + has per-stem spectral errors (e.g. 'other'/'vocals' carry spurious sub,
overall dull). This learns to GLUE: it takes the decoded restored stems and predicts per-stem TF gain masks
(cross-stem context), applies them, and sums to a cohesive mix supervised against the real clean mix with a
multi-resolution STFT loss (decoded domain => audible, unlike latent edits which decoded ~inert).
Gains-only => cannot hallucinate content (safe); it can only rebalance/shape what the stems already have.
prep+train: python -m restoflow.stem_gluer --n 800 --epochs 60 --device cuda:0
smoke: python -m restoflow.stem_gluer --smoke --device cuda:0
"""
from __future__ import annotations
import argparse, random, collections, csv, time
from pathlib import Path
import torch, torch.nn as nn, torch.nn.functional as F
from .config import Cfg, STEMS, STEM_ID
from . import eval as E, data as Dm
from .model import AttnRestorer
BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra")
ROOT = BASE / "demucs_results_all"; T = 33; SR = 44100
NF, HOP = 2048, 512 # gluer's working STFT
# ---------------- multi-res STFT loss ----------------
def mrstft(a, b, eps=1e-5):
a = a.mean(1); b = b.mean(1) # [B,2,S]->mono
tot = 0.0
for nf in (512, 1024, 2048):
w = torch.hann_window(nf, device=a.device)
A = torch.stft(a, nf, hop_length=nf // 4, window=w, return_complex=True).abs()
B = torch.stft(b, nf, hop_length=nf // 4, window=w, return_complex=True).abs()
tot = tot + torch.linalg.norm(B - A) / (torch.linalg.norm(B) + eps) \
+ (torch.log(A + eps) - torch.log(B + eps)).abs().mean()
return tot / 3.0
# ---------------- gluer net: per-stem TF gain masks from cross-stem context ----------------
class GlueNet(nn.Module):
"""Input per-stem log-mag [B,nstem,F,T] -> per-stem real gain mask [B,nstem,F,T] in [0, gmax].
Cross-stem: stems folded into channels so convs see all stems jointly (learns inter-stem balance)."""
def __init__(self, n_stems=4, F_bins=NF // 2 + 1, hidden=64, gmax=4.0):
super().__init__()
self.n = n_stems; self.gmax = gmax
c = n_stems
self.net = nn.Sequential(
nn.Conv2d(c, hidden, 3, padding=1), nn.GELU(),
nn.Conv2d(hidden, hidden, 3, padding=1), nn.GELU(),
nn.Conv2d(hidden, hidden, 3, padding=1), nn.GELU(),
nn.Conv2d(hidden, c, 3, padding=1),
)
nn.init.zeros_(self.net[-1].weight); nn.init.zeros_(self.net[-1].bias) # start at gain=1 (identity glue)
def forward(self, logmag): # [B,n,F,T]
d = self.net(logmag) # zero-init -> d=0 -> gain=1 (identity glue)
return torch.exp(d.clamp(-2.5, 1.4)) # per-stem gain in ~[0.08, 4.0]
def _stft(x): # x [B,2,S] -> complex [B,2,F,Tt]
w = torch.hann_window(NF, device=x.device)
X = torch.stft(x.reshape(-1, x.shape[-1]), NF, hop_length=HOP, window=w, return_complex=True)
return X.reshape(x.shape[0], 2, X.shape[-2], X.shape[-1])
def _istft(X, length): # complex [B,2,F,Tt] -> [B,2,S]
w = torch.hann_window(NF, device=X.device)
x = torch.istft(X.reshape(-1, X.shape[-2], X.shape[-1]), NF, hop_length=HOP, window=w, length=length)
return x.reshape(X.shape[0], 2, length)
def glue(net, stems): # stems [B,n,2,S] -> mix [B,2,S]
B, n, _, S = stems.shape
dry = stems.sum(1) # exact time-domain dry mix [B,2,S]
Xs = torch.stack([_stft(stems[:, i]) for i in range(n)], 1) # [B,n,2,F,Tt]
mag = Xs.abs().mean(2) # [B,n,F,Tt] (mono mag for context)
g = net(torch.log(mag + 1e-5)) # [B,n,F,Tt] per-stem gain (=1 at init)
corr = (Xs * (g.unsqueeze(2) - 1.0)).sum(1) # gain DEVIATION -> 0 at init
mix = dry + _istft(corr, S) # residual: identity is EXACT dry sum
return mix
# ---------------- decode-once dataset (RAM) ----------------
def build_cache(n, dev, restorer_run):
ck = torch.load(Path(restorer_run) / "ckpt_best.pt", map_location="cpu"); c = ck["cfg"]
rest = 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"]).to(dev).eval()
rest.load_state_dict(ck["model"])
for p in rest.parameters(): p.requires_grad_(False)
stats = torch.load(Path(restorer_run) / "norm_stats.pt", map_location="cpu")
mu = {s: stats[s]["mu"].to(dev) for s in STEMS}; sd = {s: stats[s]["sd"].to(dev) for s in STEMS}
mm = stats["__mix__"]; cfg = Cfg(); cfg.device = dev; sa, _ = E.load_same(cfg)
def lz(p): x = torch.load(ROOT / p, map_location="cpu"); return Dm._fit_T((x if torch.is_tensor(x) else torch.from_numpy(x)).float(), T)
def dec(latent): return sa.decode_audio(latent[None].to(dev))[0].clamp(-1, 1).float().cpu()
rows = list(csv.DictReader(open(ROOT / "metadata.csv"))); bypair = collections.defaultdict(dict)
for r in rows: bypair[r["pair_id"]][r["stem"]] = r
ok = lambda p: p and (ROOT / p).is_file()
pairs = [pid for pid in bypair if (ROOT / pid / "latents" / "degraded_mix.pt").exists()]
random.Random(0).shuffle(pairs)
IN, TGT = [], []; built = 0
with torch.inference_mode():
for pid in pairs:
if built >= n: break
rs = bypair[pid]
if not any(r.get("clean_present") == "1" and ok(r.get("clean_latent_path")) for r in rs.values()): continue
mixp = ROOT / pid / "latents" / "degraded_mix.pt"
mix_lat = Dm._fit_T(torch.load(mixp, map_location="cpu").float(), T)
sd_m = ((mix_lat - mm["mu"][:, None]) / mm["sd"][:, None]).to(dev)[None]
in_stems = []; clean_audios = []; L = None
for s in STEMS:
r = rs.get(s, {}); si = STEM_ID[s]
# input stem = restored audio (if deg present) else silence; clean = decode(clean) if present
in_a = None; cl_a = None
if r.get("clean_present") == "1" and ok(r.get("clean_latent_path")):
cl = lz(r["clean_latent_path"]); cl_a = dec(cl)
if r.get("class") == "restoration" and ok(r.get("deg_latent_path")):
deg = lz(r["deg_latent_path"]).to(dev)
o = (rest(((deg - mu[s][:, None]) / sd[s][:, None])[None], torch.tensor([si], device=dev), sd_m)[0] * sd[s][:, None] + mu[s][:, None]).cpu()
in_a = dec(o)
L = cl_a.shape[-1] if cl_a is not None else (in_a.shape[-1] if in_a is not None else L)
in_stems.append(in_a); clean_audios.append(cl_a)
if L is None: continue
z = torch.zeros(2, L)
IN.append(torch.stack([(a if a is not None else z)[:, :L] for a in in_stems])) # [n,2,L]
TGT.append(sum((a[:, :L] if a is not None else z) for a in clean_audios)) # [2,L]
built += 1
if built % 50 == 0: print(f" cached {built}/{n}", flush=True)
Lm = min(x.shape[-1] for x in TGT)
IN = torch.stack([x[..., :Lm] for x in IN]); TGT = torch.stack([x[..., :Lm] for x in TGT])
print(f"[glue] cache built: IN {tuple(IN.shape)} TGT {tuple(TGT.shape)}", flush=True)
return IN, TGT
def main():
p = argparse.ArgumentParser()
p.add_argument("--restorer", default=str(BASE / "restoflow_runs/restorer_attn_advramp"))
p.add_argument("--out-dir", default=str(BASE / "restoflow_runs/glue_v1"))
p.add_argument("--n", type=int, default=800); p.add_argument("--epochs", type=int, default=60)
p.add_argument("--batch", type=int, default=8); p.add_argument("--lr", type=float, default=1e-4)
p.add_argument("--hidden", type=int, default=32); p.add_argument("--val-frac", type=float, default=0.1)
p.add_argument("--device", default="cuda:0"); p.add_argument("--smoke", action="store_true")
a = p.parse_args(); dev = a.device
if a.smoke: a.n = 24; a.epochs = 2
out = Path(a.out_dir); out.mkdir(parents=True, exist_ok=True)
IN, TGT = build_cache(a.n, dev, a.restorer)
N = IN.shape[0]; nval = max(2, int(N * a.val_frac)); idx = list(range(N)); random.Random(1).shuffle(idx)
val, tr = idx[:nval], idx[nval:]
net = GlueNet(n_stems=len(STEMS), hidden=a.hidden).to(dev)
opt = torch.optim.AdamW(net.parameters(), lr=a.lr)
# baseline: naive sum (gluer identity) distance to clean
def mix_loss(ids, train=False):
tot = 0.0
for i in range(0, len(ids), a.batch):
b = ids[i:i + a.batch]; s = IN[b].to(dev); t = TGT[b].to(dev)
try:
if train:
m = glue(net, s); loss = mrstft(m, t); opt.zero_grad(); loss.backward(); opt.step()
else:
with torch.no_grad(): loss = mrstft(glue(net, s), t)
except torch.cuda.OutOfMemoryError: # survive transient GPU contention
torch.cuda.empty_cache(); opt.zero_grad(set_to_none=True); continue
tot += float(loss) * len(b)
return tot / len(ids)
with torch.no_grad():
base = sum(float(mrstft(IN[val].to(dev)[i:i+1].sum(1), TGT[val].to(dev)[i:i+1])) for i in range(len(val))) / len(val)
print(f"[glue] N={N} train={len(tr)} val={len(val)} baseline(sum) val mrstft={base:.4f}", flush=True)
best = float("inf")
for ep in range(a.epochs):
random.shuffle(tr); tl = mix_loss(tr, train=True); vl = mix_loss(val, train=False)
tag = " *BEST" if vl < best else ""
print(f"ep{ep} train={tl:.4f} val={vl:.4f} (base {base:.4f}){tag}", flush=True)
if vl < best:
best = vl; torch.save({"net": net.state_dict(), "hidden": a.hidden, "n_stems": len(STEMS),
"nf": NF, "hop": HOP, "epoch": ep + 1, "val": vl, "base": base}, out / "ckpt_best.pt")
print(f"done. best val mrstft={best:.4f} vs baseline-sum {base:.4f} -> {out/'ckpt_best.pt'}", flush=True)
print("GLUE_DONE")
if __name__ == "__main__":
main()