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