# Adapted from Stability AI's stable-audio-3 (MIT License). See LICENSE. from __future__ import annotations from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from safetensors.torch import load_file WEIGHTS = Path(__file__).parent / "samel" / "model.safetensors" SAMPLE_RATE = 44100 LATENT_DIM = 256 DIM = 1536 DEPTH = 12 HEADS = 24 HEAD_DIM = 64 ROPE_DIM = 32 FF_HIDDEN = 3 * DIM SINUSOIDAL_FROM = 5 PATCH = 256 STRIDE = 16 GROUP = STRIDE + 1 WINDOW = GROUP CHUNK = 128 OVERLAP = 32 SEQ = CHUNK * GROUP SAMPLES_PER_FRAME = PATCH * STRIDE MASK_NOISE = 0.1 LATENT_NOISE = 1e-3 class DynamicTanh(nn.Module): def __init__(self, dim: int): super().__init__() self.alpha = nn.Parameter(torch.ones(1)) self.gamma = nn.Parameter(torch.ones(dim)) self.beta = nn.Parameter(torch.zeros(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.gamma * F.tanh(self.alpha * x) + self.beta def apply_rope(t: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor: rot, keep = t[..., :ROPE_DIM], t[..., ROPE_DIM:] x1, x2 = rot.chunk(2, dim=-1) rotated = torch.cat([-x2, x1], dim=-1) return torch.cat([rot * freqs.cos() + rotated * freqs.sin(), keep], dim=-1) class Attention(nn.Module): """Differential attention: two attention maps per head, subtracted.""" def __init__(self): super().__init__() self.to_qkv = nn.Linear(DIM, 5 * DIM, bias=False) self.to_out = nn.Linear(DIM, DIM, bias=False) self.q_norm = DynamicTanh(HEAD_DIM) self.k_norm = DynamicTanh(HEAD_DIM) def forward(self, x: torch.Tensor, freqs: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: B, N, _ = x.shape q, k, v, q_diff, k_diff = ( self.to_qkv(x).view(B, N, 5, HEADS, HEAD_DIM).permute(2, 0, 3, 1, 4)) q, q_diff = apply_rope(self.q_norm(q).float(), freqs), apply_rope( self.q_norm(q_diff).float(), freqs) k, k_diff = apply_rope(self.k_norm(k).float(), freqs), apply_rope( self.k_norm(k_diff).float(), freqs) out = F.scaled_dot_product_attention(q.to(v.dtype), k.to(v.dtype), v, attn_mask=mask) out = out - F.scaled_dot_product_attention( q_diff.to(v.dtype), k_diff.to(v.dtype), v, attn_mask=mask) return self.to_out(out.transpose(1, 2).reshape(B, N, DIM)) class Sin(nn.Module): def forward(self, x: torch.Tensor) -> torch.Tensor: return torch.sin(torch.pi * x) class GatedProjection(nn.Module): def __init__(self, activation: nn.Module): super().__init__() self.proj = nn.Linear(DIM, 2 * FF_HIDDEN) self.act = activation def forward(self, x: torch.Tensor) -> torch.Tensor: value, gate = self.proj(x).chunk(2, dim=-1) return value * self.act(gate) class FeedForward(nn.Module): def __init__(self, activation: nn.Module): super().__init__() self.ff = nn.Sequential( GatedProjection(activation), nn.Identity(), nn.Linear(FF_HIDDEN, DIM)) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.ff(x) class Block(nn.Module): def __init__(self, sinusoidal: bool): super().__init__() self.pre_norm = DynamicTanh(DIM) self.self_attn = Attention() self.ff_norm = DynamicTanh(DIM) self.ff = FeedForward(Sin() if sinusoidal else nn.SiLU()) def forward(self, x: torch.Tensor, freqs: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: x = x + self.self_attn(self.pre_norm(x), freqs, mask) return x + self.ff(self.ff_norm(x)) class Resampler(nn.Module): def __init__(self): super().__init__() self.new_tokens = nn.Parameter(torch.zeros(1, 1, DIM)) self.blocks = nn.ModuleList(Block(i >= SINUSOIDAL_FROM) for i in range(DEPTH)) self.mapping = nn.Conv1d(DIM, 2 * PATCH, 1) def forward(self, x: torch.Tensor, freqs: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: B, T, _ = x.shape slots = self.new_tokens.expand(B * T, STRIDE, DIM) slots = slots + torch.randn_like(slots) * MASK_NOISE x = torch.cat([x.reshape(B * T, 1, DIM), slots], dim=1).reshape(B, T * GROUP, DIM) for block in self.blocks: x = block(x, freqs, mask) x = x.reshape(B * T, GROUP, DIM)[:, 1:] return self.mapping(x.reshape(B, T * STRIDE, DIM).transpose(1, 2)) class SameLDecoder(nn.Module): def __init__(self): super().__init__() self.proj_in = nn.Linear(LATENT_DIM, DIM) self.resampler = Resampler() self.register_buffer("running_std", torch.ones(1)) position = torch.arange(SEQ) self.register_buffer( "mask", (position[None, :] - position[:, None]).abs() <= WINDOW, persistent=False) self.register_buffer( "inv_freq", 1.0 / (10000.0 ** (torch.arange(0, ROPE_DIM, 2) / ROPE_DIM)), persistent=False) def _rope_freqs(self) -> torch.Tensor: freqs = torch.outer(torch.arange(SEQ, device=self.inv_freq.device).float(), self.inv_freq.float()) return torch.cat([freqs, freqs], dim=-1) def _decode_chunk(self, latents: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor: x = latents * self.running_std x = x + torch.randn_like(x) * self.running_std * LATENT_NOISE x = self.resampler(self.proj_in(x.transpose(1, 2)), freqs, self.mask) B, _, L = x.shape return x.view(B, 2, PATCH, L).permute(0, 1, 3, 2).reshape(B, 2, L * PATCH) @torch.no_grad() def decode(self, latents: torch.Tensor) -> torch.Tensor: """(B, 256, T) latents, T >= CHUNK, to (B, 2, T * 4096) audio.""" frames = latents.shape[-1] starts = list(range(0, frames - CHUNK + 1, CHUNK - OVERLAP)) if starts[-1] != frames - CHUNK: starts.append(frames - CHUNK) freqs = self._rope_freqs() edge = OVERLAP // 2 * SAMPLES_PER_FRAME audio = latents.new_zeros(latents.shape[0], 2, frames * SAMPLES_PER_FRAME) for i, start in enumerate(starts): chunk = self._decode_chunk(latents[..., start:start + CHUNK], freqs) left = 0 if i == 0 else edge right = chunk.shape[-1] if i == len(starts) - 1 else chunk.shape[-1] - edge at = start * SAMPLES_PER_FRAME audio[..., at + left:at + right] = chunk[..., left:right] return audio def load_decoder(device: str = "cuda", dtype: torch.dtype = torch.float16) -> SameLDecoder: """Load the bundled SAME-L weights, keeping only what decoding needs.""" raw = load_file(WEIGHTS) state = {"running_std": raw["bottleneck.running_std"]} g, v = raw["decoder.layers.3.mapping.weight_g"], raw["decoder.layers.3.mapping.weight_v"] state["resampler.mapping.weight"] = g * v / v.norm(dim=(1, 2), keepdim=True) for key, tensor in raw.items(): if key.startswith("decoder.layers.1."): state[key.replace("decoder.layers.1.", "proj_in.")] = tensor elif key.startswith("decoder.layers.3.") and not key.endswith( ("rope.inv_freq", "mapping.weight_g", "mapping.weight_v")): state[key.replace("decoder.layers.3.", "resampler.") .replace("transformers.", "blocks.")] = tensor decoder = SameLDecoder() decoder.load_state_dict(state) return decoder.to(device=device, dtype=dtype).eval().requires_grad_(False)