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