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