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