File size: 18,522 Bytes
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7153194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af4583e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
"""Conditional flow velocity network over the SAME latent sequence [B,256,T].

Predicts velocity v(z_t, t | degraded, stem). Small (latents are tiny). One SHARED
net conditioned on stem id (FiLM) — validated >= separate per-stem networks.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn


def sinusoidal(t: torch.Tensor, dim: int) -> torch.Tensor:  # t [B] in [0,1]
    half = dim // 2
    freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / max(half - 1, 1))
    a = t[:, None] * freqs[None] * 2 * math.pi
    return torch.cat([a.sin(), a.cos()], dim=-1)


class FiLMBlock(nn.Module):
    def __init__(self, h, cond_dim, k=5):
        super().__init__()
        self.conv1 = nn.Conv1d(h, h, k, padding=k // 2)
        self.conv2 = nn.Conv1d(h, h, 1)
        self.film = nn.Linear(cond_dim, 2 * h)
        self.act = nn.GELU()

    def forward(self, x, cond):
        g, b = self.film(cond).chunk(2, dim=-1)             # [B,h] each
        h = self.conv1(x)
        h = h * (1 + g[:, :, None]) + b[:, :, None]
        h = self.act(h)
        h = self.conv2(h)
        return x + h


class CondFlow(nn.Module):
    def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, t_dim=64):
        super().__init__()
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        cond_dim = t_dim + stem_emb
        self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim))
        self.t_dim = t_dim
        self.inp = nn.Conv1d(2 * d, hidden, 1)              # concat(z_t, degraded)
        self.blocks = nn.ModuleList([FiLMBlock(hidden, cond_dim) for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)  # start near identity velocity

    def forward(self, z_t, t, deg, stem_id):
        # z_t,deg [B,256,T]; t [B]; stem_id [B]
        cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1)
        h = self.inp(torch.cat([z_t, deg], dim=1))
        for blk in self.blocks:
            h = blk(h, cond)
        return self.out(h)

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class AttnCondFlow(nn.Module):
    """Attention velocity net for the generator: same Conformer-lite block that won for the
    restorer, conditioned on (t, stem) via FiLM. Global self-attention lets the invented stem
    see the whole context window (conv-only CondFlow could not), which lands flow endpoints
    closer to the clean manifold (the off-manifold endpoints are what decode too loud)."""
    def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64,
                 heads=8, t_max=64):
        super().__init__()
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim))
        self.t_dim = t_dim
        cond_dim = t_dim + stem_emb
        self.inp = nn.Conv1d(2 * d, hidden, 1)                  # concat(z_t, ctx)
        self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False)
        self.blocks = nn.ModuleList([ConformerBlock(hidden, cond_dim, heads) for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)

    def forward(self, z_t, t, deg, stem_id):
        cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1)
        h = self.inp(torch.cat([z_t, deg], dim=1)) + self.pos[:, :, :z_t.shape[2]]
        for blk in self.blocks:
            h = blk(h, cond)
        return self.out(h)

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class ConformerBlock(nn.Module):
    """Pre-norm Conformer-lite: half-FFN -> MHSA -> FiLM depthwise-conv -> half-FFN.
    Global self-attention (T~=32 is tiny) fixes the conv-only receptive-field gap;
    per-block FiLM injects stem id at every depth instead of once at the input."""
    def __init__(self, h, cond_dim, heads=4, k=7, ff_mult=2):
        super().__init__()
        self.ln_ff1 = nn.LayerNorm(h)
        self.ff1 = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h))
        self.ln_attn = nn.LayerNorm(h)
        self.attn = nn.MultiheadAttention(h, heads, batch_first=True)
        self.ln_conv = nn.LayerNorm(h)
        self.pw1 = nn.Conv1d(h, 2 * h, 1)
        self.dw = nn.Conv1d(h, h, k, padding=k // 2, groups=h)
        self.gn = nn.GroupNorm(1, h)
        self.film = nn.Linear(cond_dim, 2 * h)
        self.pw2 = nn.Conv1d(h, h, 1)
        self.act = nn.GELU()
        self.ln_ff2 = nn.LayerNorm(h)
        self.ff2 = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h))

    def forward(self, x, cond):                                  # x [B,h,T], cond [B,cond_dim]
        xt = x.transpose(1, 2)                                   # [B,T,h]
        xt = xt + 0.5 * self.ff1(self.ln_ff1(xt))
        a = self.ln_attn(xt)
        a, _ = self.attn(a, a, a, need_weights=False)
        xt = xt + a
        # conv module (channels-first)
        c = self.ln_conv(xt).transpose(1, 2)                    # [B,h,T]
        c = nn.functional.glu(self.pw1(c), dim=1)
        c = self.gn(self.dw(c))
        g, b = self.film(cond).chunk(2, dim=-1)
        c = self.act(c * (1 + g[:, :, None]) + b[:, :, None])
        c = self.pw2(c)
        xt = xt + c.transpose(1, 2)
        xt = xt + 0.5 * self.ff2(self.ln_ff2(xt))
        return xt.transpose(1, 2)


class CrossAttnBlock(nn.Module):
    """DiT-style block for the generator: self-attn over the noised target frames ->
    CROSS-attn into per-stem context tokens (keeps which instrument is which, vs the
    old summed context) -> FiLM(t,stem) depthwise-conv -> FFN. A key-padding mask hides
    absent context stems (and, on CFG drop, ALL real context -> only a learned null token)."""
    def __init__(self, h, cond_dim, heads=8, k=7, ff_mult=2):
        super().__init__()
        self.ln_sa = nn.LayerNorm(h)
        self.self_attn = nn.MultiheadAttention(h, heads, batch_first=True)
        self.ln_ca = nn.LayerNorm(h)
        self.cross_attn = nn.MultiheadAttention(h, heads, batch_first=True)
        self.ln_conv = nn.LayerNorm(h)
        self.pw1 = nn.Conv1d(h, 2 * h, 1)
        self.dw = nn.Conv1d(h, h, k, padding=k // 2, groups=h)
        self.gn = nn.GroupNorm(1, h)
        self.film = nn.Linear(cond_dim, 2 * h)
        self.pw2 = nn.Conv1d(h, h, 1)
        self.act = nn.GELU()
        self.ln_ff = nn.LayerNorm(h)
        self.ff = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h))

    def forward(self, x, ctx_tok, cond, ctx_pad):       # x [B,T,h], ctx_tok [B,M,h], cond [B,cond_dim], ctx_pad [B,M] (True=ignore)
        a = self.ln_sa(x)
        a, _ = self.self_attn(a, a, a, need_weights=False)
        x = x + a
        q = self.ln_ca(x)
        c, _ = self.cross_attn(q, ctx_tok, ctx_tok, key_padding_mask=ctx_pad, need_weights=False)
        x = x + c
        cc = self.ln_conv(x).transpose(1, 2)            # [B,h,T]
        cc = nn.functional.glu(self.pw1(cc), dim=1)
        cc = self.gn(self.dw(cc))
        g, b = self.film(cond).chunk(2, dim=-1)
        cc = self.act(cc * (1 + g[:, :, None]) + b[:, :, None])
        cc = self.pw2(cc)
        x = x + cc.transpose(1, 2)
        x = x + self.ff(self.ln_ff(x))
        return x


class XAttnCondFlow(nn.Module):
    """Cross-attention (DiT-style) conditional flow generator. Each CONTEXT stem latent is
    encoded as its own token sequence (projected + per-instrument embedding + positional);
    the noised target cross-attends into the union of those tokens. This replaces summing the
    context (which discarded instrument identity) and is the principled 'transformer encoder'
    conditioning from the bass-accompaniment literature (Sony arXiv:2402.01412). A learned
    null token gives a well-defined unconditional pass for classifier-free guidance."""
    def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64,
                 heads=8, t_max=64):
        super().__init__()
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim))
        self.t_dim = t_dim
        cond_dim = t_dim + stem_emb
        self.inp = nn.Conv1d(d, hidden, 1)                       # noised target only (context via cross-attn)
        self.ctx_proj = nn.Conv1d(d, hidden, 1)                  # shared per-stem context projection
        self.ctx_stem_emb = nn.Embedding(n_stems, hidden)        # which instrument each context token is
        self.null_ctx = nn.Parameter(torch.randn(1, 1, hidden) * 0.02)
        self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False)
        self.blocks = nn.ModuleList([CrossAttnBlock(hidden, cond_dim, heads) for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)

    def _encode_ctx(self, ctx_stems, ctx_ids, ctx_mask):
        # ctx_stems [B,K,d,T]; ctx_ids [B,K]; ctx_mask [B,K] (True=present)
        B, K, d, T = ctx_stems.shape
        x = self.ctx_proj(ctx_stems.reshape(B * K, d, T))        # [B*K,h,T]
        x = x + self.pos[:, :, :T]
        x = x.transpose(1, 2)                                    # [B*K,T,h]
        x = x + self.ctx_stem_emb(ctx_ids.reshape(B * K))[:, None, :]
        tok = x.reshape(B, K * T, -1)
        pad = (~ctx_mask)[:, :, None].expand(B, K, T).reshape(B, K * T)   # True = ignore
        null = self.null_ctx.expand(B, -1, -1)                   # [B,1,h] always attended
        tok = torch.cat([null, tok], dim=1)
        nullpad = torch.zeros(B, 1, dtype=torch.bool, device=tok.device)
        return tok, torch.cat([nullpad, pad], dim=1)

    def forward(self, z_t, t, ctx_stems, ctx_ids, ctx_mask, stem_id):
        cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1)
        tok, pad = self._encode_ctx(ctx_stems, ctx_ids, ctx_mask)
        x = self.inp(z_t) + self.pos[:, :, :z_t.shape[2]]        # [B,h,T]
        x = x.transpose(1, 2)                                    # [B,T,h]
        for blk in self.blocks:
            x = blk(x, tok, cond, pad)
        return self.out(x.transpose(1, 2))

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class MixAttnCondFlow(nn.Module):
    """STOCHASTIC restorer: attention velocity net for the deg-anchored flow bridge, WITH mix
    conditioning. Combines the three validated levers — attention (won deterministically),
    mix-conditioning (config notes: helps drums most), and stochasticity (the literature fix
    for the smeared transients that L2 regression averages away on percussive stems). The
    deterministic restorers can't add transient detail (conditional-mean); this can sample it.
    Signature model(z_t, t, deg, stem_id, mix) matches flow.fm_loss_mix / flow.sample_mix."""
    def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64,
                 heads=8, t_max=64, use_mix=True):
        super().__init__()
        self.use_mix = use_mix
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim))
        self.t_dim = t_dim
        cond_dim = t_dim + stem_emb
        self.inp = nn.Conv1d(d * (3 if use_mix else 2), hidden, 1)   # concat(z_t, deg, [mix])
        self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False)
        self.blocks = nn.ModuleList([ConformerBlock(hidden, cond_dim, heads) for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)

    def forward(self, z_t, t, deg, stem_id, mix=None):
        cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1)
        parts = [z_t, deg, mix] if (self.use_mix and mix is not None) else [z_t, deg]
        h = self.inp(torch.cat(parts, dim=1)) + self.pos[:, :, :z_t.shape[2]]
        for blk in self.blocks:
            h = blk(h, cond)
        return self.out(h)

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class AttnRestorer(nn.Module):
    """Deterministic restorer with global attention (Conformer-lite). Residual on the
    degraded latent, like DetRestorer, so it inherits the identity-passthrough init."""
    def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, use_mix=True,
                 heads=4, t_max=64):
        super().__init__()
        self.use_mix = use_mix
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        in_ch = d * (2 if use_mix else 1)
        self.inp = nn.Conv1d(in_ch, hidden, 1)
        self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False)
        self.blocks = nn.ModuleList([ConformerBlock(hidden, stem_emb, heads) for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)  # start at identity (deg passthrough)

    def forward(self, deg, stem_id, mix=None):
        cond = self.stem_emb(stem_id)                            # [B,stem_emb]
        parts = [deg, mix] if (self.use_mix and mix is not None) else [deg]
        h = self.inp(torch.cat(parts, dim=1))                   # [B,hidden,T]
        h = h + self.pos[:, :, :h.shape[2]]
        for blk in self.blocks:
            h = blk(h, cond)
        return deg + self.out(h)

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class LatentDiscriminator(nn.Module):
    """Spectral-norm conv critic over SAME latents [B,256,T] -> per-frame logits (PatchGAN-style)
    + feature taps (for the stable feature-matching loss). Latent-domain = cheap, no decode
    ('the latent is the mirror'). Judges full-mix realism in the remix-GAN fine-tune stage."""
    def __init__(self, d=256, hidden=256, depth=4, k=5):
        super().__init__()
        from torch.nn.utils import spectral_norm as SN
        self.blocks = nn.ModuleList()
        c = d
        for _ in range(depth):
            self.blocks.append(nn.Sequential(SN(nn.Conv1d(c, hidden, k, padding=k // 2)),
                                             nn.LeakyReLU(0.2)))
            c = hidden
        self.head = SN(nn.Conv1d(c, 1, 1))

    def forward(self, x, return_feats=False):
        feats = []; h = x
        for b in self.blocks:
            h = b(h); feats.append(h)
        logit = self.head(h)                                     # [B,1,T] per-frame
        return (logit, feats) if return_feats else logit

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


class MultiScaleLatentDiscriminator(nn.Module):
    """Multi-SCALE critic over SAME latents [B,256,T]: K independent LatentDiscriminators, each on a
    temporally avg-pooled view (stride 1,2,4). Latent-domain analogue of the multi-scale/multi-resolution
    STFT discriminators that are the STANDARD anti-artifact recipe in neural audio synthesis (EnCodec,
    Defossez 2022; DAC, Kumar 2023; BigVGAN multi-resolution D, Lee 2023). The fine scale (stride 1)
    catches per-frame hiss/buzz; the coarse scales (stride 2,4) judge structure/'glue' across the clip
    -> directly targets both the hiss and the cross-chunk coherence the single-scale critic misses.
    Returns mean-over-scales logit (+ pooled feature taps for HiFi-GAN feature matching, Kong 2020)."""
    def __init__(self, d=256, hidden=256, depth=4, k=5, scales=(1, 2, 4)):
        super().__init__()
        self.scales = scales
        self.discs = nn.ModuleList([LatentDiscriminator(d, hidden, depth, k) for _ in scales])

    def _view(self, x, s):
        return x if s == 1 else torch.nn.functional.avg_pool1d(x, kernel_size=s, stride=s)

    def forward(self, x, return_feats=False):
        logits, feats = [], []
        for s, disc in zip(self.scales, self.discs):
            xv = self._view(x, s)
            if return_feats:
                lg, ft = disc(xv, return_feats=True); feats.extend(ft)
            else:
                lg = disc(xv)
            logits.append(lg.mean(dim=2))                         # [B,1] per-scale clip logit
        logit = torch.cat(logits, dim=1).mean(dim=1, keepdim=True)  # [B,1] mean over scales
        return (logit, feats) if return_feats else logit

    def num_params(self):
        return sum(p.numel() for p in self.parameters())


def _sinusoidal_pos(t_max: int, dim: int) -> torch.Tensor:
    pos = torch.arange(t_max).float()[:, None]
    half = dim // 2
    freqs = torch.exp(-math.log(10000) * torch.arange(half).float() / max(half - 1, 1))
    pe = torch.zeros(t_max, dim)
    pe[:, 0::2] = torch.sin(pos * freqs)[:, :pe[:, 0::2].shape[1]]
    pe[:, 1::2] = torch.cos(pos * freqs)[:, :pe[:, 1::2].shape[1]]
    return pe.t()[None]                                          # [1,dim,t_max]


class DetRestorer(nn.Module):
    """Deterministic restorer: predict clean stem latent from degraded stem (+ mix) + stem id.
    Validated > generative for reference-matching; mix-conditioning helps drums most.
    Predicts a residual on the degraded latent (small, aligned move)."""
    def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, use_mix=True):
        super().__init__()
        self.use_mix = use_mix
        self.stem_emb = nn.Embedding(n_stems, stem_emb)
        in_ch = d * (2 if use_mix else 1) + stem_emb
        self.inp = nn.Conv1d(in_ch, hidden, 1)
        self.blocks = nn.ModuleList([
            nn.Sequential(nn.Conv1d(hidden, hidden, 5, padding=2), nn.GELU(), nn.Conv1d(hidden, hidden, 1))
            for _ in range(depth)])
        self.out = nn.Conv1d(hidden, d, 1)
        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)  # start at identity (deg passthrough)

    def forward(self, deg, stem_id, mix=None):
        em = self.stem_emb(stem_id)[:, :, None].expand(-1, -1, deg.shape[2])
        parts = [deg, mix, em] if (self.use_mix and mix is not None) else [deg, em]
        h = self.inp(torch.cat(parts, dim=1))
        for blk in self.blocks:
            h = h + blk(h)
        return deg + self.out(h)

    def num_params(self):
        return sum(p.numel() for p in self.parameters())