soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
18.5 kB
"""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())