Buckets:
| """Tag-conditioned diffusion transformer over SAME-L music latents (0.84B params).""" | |
| from __future__ import annotations | |
| import math | |
| 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 / "audio_dit.safetensors" | |
| CHANNELS = 256 # SAME-L latent channels | |
| FRAMES = 1024 # latent frames per sample, at 10.77 frames/second | |
| DIM = 1536 | |
| HEADS = 24 | |
| HEAD_DIM = 64 | |
| MLP_HIDDEN = 4096 | |
| NUM_TAGS = 1470 # vocabulary rows; index 0 is padding | |
| MAX_TAGS = 8 | |
| ROPE_THETA = 10000.0 | |
| # SPRINT sparse-dense fusion: a dense encoder, a deep middle stack, a dense | |
| # decoder. Encoder and decoder blocks cross-attend to the tags, middle blocks | |
| # every second block. Training subsampled the middle stack; inference runs it on | |
| # every frame. | |
| ENCODER_BLOCKS = 2 | |
| MIDDLE_BLOCKS = 20 | |
| DECODER_BLOCKS = 2 | |
| def rms_norm(x: torch.Tensor) -> torch.Tensor: | |
| return F.rms_norm(x.float(), (x.shape[-1],), eps=1e-6).to(x.dtype) | |
| class RMSNorm(nn.Module): | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(dim)) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return rms_norm(x) * self.weight | |
| def rope_tables() -> tuple[torch.Tensor, torch.Tensor]: | |
| freqs = torch.outer( | |
| torch.arange(FRAMES).float(), | |
| 1.0 / (ROPE_THETA ** (torch.arange(0, HEAD_DIM, 2).float() / HEAD_DIM)), | |
| ) | |
| return freqs.cos()[None, None], freqs.sin()[None, None] | |
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| x1, x2 = x.float().chunk(2, dim=-1) | |
| return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1).to(x.dtype) | |
| class SelfAttention(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.qkv_x = nn.Linear(DIM, 3 * DIM, bias=False) | |
| self.q_norm = RMSNorm(HEAD_DIM) | |
| self.k_norm = RMSNorm(HEAD_DIM) | |
| self.proj = nn.Linear(DIM, DIM, bias=False) | |
| def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| B, S, _ = x.shape | |
| q, k, v = self.qkv_x(x).view(B, S, 3, HEADS, HEAD_DIM).permute(2, 0, 3, 1, 4) | |
| q = apply_rope(self.q_norm(q), cos, sin) | |
| k = apply_rope(self.k_norm(k), cos, sin) | |
| out = F.scaled_dot_product_attention(q, k, v) | |
| return self.proj(out.transpose(1, 2).reshape(B, S, DIM)) | |
| class CrossAttention(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.q_y = nn.Linear(DIM, DIM, bias=False) | |
| self.kv_y = nn.Linear(DIM, 2 * DIM, bias=False) | |
| self.q_norm = RMSNorm(HEAD_DIM) | |
| self.k_norm = RMSNorm(HEAD_DIM) | |
| self.proj_y = nn.Linear(DIM, DIM, bias=False) | |
| def forward(self, x: torch.Tensor, y: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: | |
| B, S, _ = x.shape | |
| q = self.q_norm(self.q_y(x).view(B, S, HEADS, HEAD_DIM).transpose(1, 2)) | |
| k, v = self.kv_y(y).view(B, MAX_TAGS, 2, HEADS, HEAD_DIM).permute(2, 0, 3, 1, 4) | |
| out = F.scaled_dot_product_attention(q, self.k_norm(k), v, | |
| attn_mask=mask[:, None, None, :]) | |
| return self.proj_y(out.transpose(1, 2).reshape(B, S, DIM)) | |
| class SwiGLU(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.w12 = nn.Linear(DIM, 2 * MLP_HIDDEN, bias=False) | |
| self.w3 = nn.Linear(MLP_HIDDEN, DIM, bias=False) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| gate, up = self.w12(x).chunk(2, dim=-1) | |
| return self.w3(F.silu(gate) * up) | |
| class Block(nn.Module): | |
| """adaLN-single: shared timestep modulation plus a learned per-block offset.""" | |
| def __init__(self): | |
| super().__init__() | |
| self.attn = SelfAttention() | |
| self.mlp = SwiGLU() | |
| self.adaLN_offset = nn.Parameter(torch.zeros(6 * DIM)) | |
| def forward(self, x, mod, y, mask, cos, sin): | |
| shift1, scale1, gate1, shift2, scale2, gate2 = (mod + self.adaLN_offset).chunk(6, -1) | |
| x = x + gate1 * self.attn(rms_norm(x) * (1 + scale1) + shift1, cos, sin) | |
| return x + gate2 * self.mlp(rms_norm(x) * (1 + scale2) + shift2) | |
| class CrossBlock(Block): | |
| def __init__(self): | |
| super().__init__() | |
| self.cross = CrossAttention() | |
| self.cross_gate = nn.Parameter(torch.zeros(DIM)) | |
| def forward(self, x, mod, y, mask, cos, sin): | |
| shift1, scale1, gate1, shift2, scale2, gate2 = (mod + self.adaLN_offset).chunk(6, -1) | |
| x = x + gate1 * self.attn(rms_norm(x) * (1 + scale1) + shift1, cos, sin) | |
| x = x + self.cross_gate * self.cross(rms_norm(x), y, mask) | |
| return x + gate2 * self.mlp(rms_norm(x) * (1 + scale2) + shift2) | |
| class Timestep(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.mlp = nn.Sequential(nn.Linear(256, DIM), nn.SiLU(), nn.Linear(DIM, DIM)) | |
| def forward(self, t: float, batch: int, device) -> torch.Tensor: | |
| freqs = torch.exp(-math.log(10000.0) * torch.arange(128, device=device) / 128) | |
| args = t * 1000.0 * freqs[None] | |
| return self.mlp(torch.cat([args.cos(), args.sin()], dim=-1)).expand(batch, DIM) | |
| class FinalLayer(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.adaLN_modulation = nn.Linear(DIM, 2 * DIM) | |
| self.linear = nn.Linear(DIM, CHANNELS) | |
| def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor: | |
| shift, scale = self.adaLN_modulation(F.silu(c)).unsqueeze(1).chunk(2, dim=-1) | |
| return self.linear(rms_norm(x) * (1 + scale) + shift) | |
| class AudioDiT(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.x_embedder = nn.Linear(CHANNELS, DIM) | |
| self.t_embedder = Timestep() | |
| self.tag_embedding = nn.Embedding(NUM_TAGS, DIM, padding_idx=0) | |
| self.tag_norm = RMSNorm(DIM) | |
| self.null_tag = nn.Parameter(torch.zeros(DIM)) | |
| self.mask_token = nn.Parameter(torch.zeros(DIM)) | |
| self.fusion_proj = nn.Linear(2 * DIM, DIM) | |
| self.adaLN_shared = nn.Linear(DIM, 6 * DIM) | |
| self.encoder = nn.ModuleList(CrossBlock() for _ in range(ENCODER_BLOCKS)) | |
| self.middle = nn.ModuleList( | |
| (CrossBlock if i % 2 == 0 else Block)() for i in range(MIDDLE_BLOCKS)) | |
| self.decoder = nn.ModuleList(CrossBlock() for _ in range(DECODER_BLOCKS)) | |
| self.final_layer = FinalLayer() | |
| cos, sin = rope_tables() | |
| self.register_buffer("cos", cos, persistent=False) | |
| self.register_buffer("sin", sin, persistent=False) | |
| def _condition(self, t: float, tags: torch.Tensor): | |
| y = self.tag_norm(self.tag_embedding(tags)) | |
| valid = tags != 0 | |
| # An all-padding row is unconditional: it attends to the learned null tag. | |
| empty = ~valid.any(dim=1) | |
| y = torch.where(empty[:, None, None], self.tag_norm(self.null_tag).expand_as(y), y) | |
| mask = valid.clone() | |
| mask[:, 0] |= empty | |
| c = self.t_embedder(t, tags.shape[0], tags.device) | |
| return c, self.adaLN_shared(F.silu(c)).unsqueeze(1), y, mask | |
| def _run(self, blocks, x, mod, y, mask): | |
| for block in blocks: | |
| x = block(x, mod, y, mask, self.cos, self.sin) | |
| return x | |
| def _fuse(self, f, g, c, mod, y, mask): | |
| h = self.fusion_proj(torch.cat([f, g], dim=-1)) | |
| return self.final_layer(self._run(self.decoder, h, mod, y, mask), c).transpose(1, 2) | |
| def forward(self, x: torch.Tensor, t: float, tags: torch.Tensor) -> torch.Tensor: | |
| """Velocity at (x, t), for latents (B, 256, 1024) and tag indices (B, 8).""" | |
| c, mod, y, mask = self._condition(t, tags) | |
| f = self._run(self.encoder, self.x_embedder(x.transpose(1, 2)), mod, y, mask) | |
| return self._fuse(f, self._run(self.middle, f, mod, y, mask), c, mod, y, mask) | |
| def shallow(self, x: torch.Tensor, t: float, tags: torch.Tensor) -> torch.Tensor: | |
| """Path-drop branch: the middle stack is replaced by [MASK] tokens.""" | |
| c, mod, y, mask = self._condition(t, tags) | |
| f = self._run(self.encoder, self.x_embedder(x.transpose(1, 2)), mod, y, mask) | |
| return self._fuse(f, self.mask_token.expand_as(f), c, mod, y, mask) | |
| def load_model(device: str = "cuda") -> tuple[AudioDiT, torch.Tensor, torch.Tensor]: | |
| """The model plus the per-channel mean/std the training latents were scaled by.""" | |
| state = load_file(WEIGHTS, device=device) | |
| mean = state.pop("latent_mean").view(1, CHANNELS, 1) | |
| std = state.pop("latent_std").view(1, CHANNELS, 1) | |
| model = AudioDiT() | |
| model.load_state_dict(state) | |
| return model.to(device).eval().requires_grad_(False), mean, std | |
Xet Storage Details
- Size:
- 8.66 kB
- Xet hash:
- f160d41c40c9accfcc85c1c5fa1de386ee649c7940bd0c2ec745744dbe95413d
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.