Download same_l_decoder.py from AIWizard76/Localsong: direct link, hf CLI and curl.
- Browser
- Download file 7.56 kB
-
https://huggingface.co/AIWizard76/Localsong/resolve/main/same_l_decoder.py
- Command line
-
hf download hf://AIWizard76/Localsong/same_l_decoder.py
-
curl -L -o same_l_decoder.py https://huggingface.co/AIWizard76/Localsong/resolve/main/same_l_decoder.py
7.56 kB
| # 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) | |
| 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) | |