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()