AriaLM / src /s04_model.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
11.6 kB
"""
Music Transformer Model — LLaMA-style architecture for symbolic music generation.
Key innovations combined:
- Rotary Position Embeddings (RoPE) — better long-range modeling than sinusoidal
- RMSNorm — faster than LayerNorm, used in LLaMA/Mistral
- SwiGLU activation — better than GELU/ReLU, used in LLaMA
- Grouped Query Attention (GQA) — reduces KV-cache memory by sharing KV heads
- Gradient checkpointing — cuts memory usage ~50% with ~20% speed cost
- KV-cache — O(1) per-token inference instead of O(n)
"""
import math
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class RMSNorm(nn.Module):
"""Root Mean Square Layer Normalization (faster than LayerNorm)."""
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return (x.float() * norm).type_as(x) * self.weight
def precompute_rope_freqs(dim: int, max_seq_len: int, theta: float = 10000.0) -> torch.Tensor:
"""Precompute RoPE frequency tensor for complex exponentials."""
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
t = torch.arange(max_seq_len, dtype=torch.float32)
freqs = torch.outer(t, freqs)
return torch.polar(torch.ones_like(freqs), freqs) # complex64
def apply_rope(xq: torch.Tensor, xk: torch.Tensor, freqs: torch.Tensor):
"""Apply rotary embeddings to query and key tensors."""
# Reshape to complex
xq_c = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
xk_c = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
# Reshape freqs for broadcasting: (seq_len,) -> (1, seq_len, 1, head_dim//2)
freqs = freqs.unsqueeze(0).unsqueeze(2)
xq_out = torch.view_as_real(xq_c * freqs).flatten(-2)
xk_out = torch.view_as_real(xk_c * freqs).flatten(-2)
return xq_out.type_as(xq), xk_out.type_as(xk)
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
"""Repeat KV heads to match query head count for GQA."""
if n_rep == 1:
return x
bs, seq_len, n_kv_heads, head_dim = x.shape
return (
x[:, :, :, None, :]
.expand(bs, seq_len, n_kv_heads, n_rep, head_dim)
.reshape(bs, seq_len, n_kv_heads * n_rep, head_dim)
)
class GroupedQueryAttention(nn.Module):
"""
Multi-head attention with Grouped Query Attention (GQA).
Uses fewer KV heads than Q heads to reduce memory.
"""
def __init__(self, dim: int, n_heads: int, n_kv_heads: int, dropout: float = 0.1):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.n_rep = n_heads // n_kv_heads
self.head_dim = dim // n_heads
self.wq = nn.Linear(dim, n_heads * self.head_dim, bias=False)
self.wk = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)
self.wv = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)
self.wo = nn.Linear(n_heads * self.head_dim, dim, bias=False)
self.attn_dropout = nn.Dropout(dropout)
self.resid_dropout = nn.Dropout(dropout)
# KV-cache for inference
self.cache_k: Optional[torch.Tensor] = None
self.cache_v: Optional[torch.Tensor] = None
def forward(
self,
x: torch.Tensor,
freqs: torch.Tensor,
mask: Optional[torch.Tensor] = None,
use_cache: bool = False,
) -> torch.Tensor:
bs, seq_len, _ = x.shape
q = self.wq(x).view(bs, seq_len, self.n_heads, self.head_dim)
k = self.wk(x).view(bs, seq_len, self.n_kv_heads, self.head_dim)
v = self.wv(x).view(bs, seq_len, self.n_kv_heads, self.head_dim)
# Apply RoPE to Q and K
q_rope = q.view(bs, seq_len, self.n_heads, self.head_dim)
k_rope = k.view(bs, seq_len, self.n_kv_heads, self.head_dim)
# RoPE needs (bs, seq_len, heads, head_dim) but freqs is (seq_len, head_dim//2)
# Apply per-head
q_for_rope = q_rope.reshape(bs * self.n_heads, seq_len, self.head_dim)
k_for_rope = k_rope.reshape(bs * self.n_kv_heads, seq_len, self.head_dim)
# Simpler RoPE application
q = q.transpose(1, 2) # (bs, n_heads, seq_len, head_dim)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
# Apply RoPE via cos/sin (more compatible than complex)
q, k = self._apply_rope_real(q, k, freqs)
# KV-cache for generation
if use_cache:
if self.cache_k is not None:
k = torch.cat([self.cache_k, k], dim=2)
v = torch.cat([self.cache_v, v], dim=2)
self.cache_k = k.detach()
self.cache_v = v.detach()
# GQA: repeat KV heads
k = repeat_kv(k.transpose(1, 2), self.n_rep).transpose(1, 2)
v = repeat_kv(v.transpose(1, 2), self.n_rep).transpose(1, 2)
# Scaled dot-product attention (uses Flash Attention when available)
scale = 1.0 / math.sqrt(self.head_dim)
try:
# PyTorch 2.0+ SDPA with memory-efficient backend
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=mask,
dropout_p=self.attn_dropout.p if self.training else 0.0,
is_causal=(mask is None and not use_cache),
)
except RuntimeError:
# Fallback for older PyTorch
scores = torch.matmul(q, k.transpose(-2, -1)) * scale
if mask is not None:
scores = scores + mask
elif not use_cache:
causal = torch.triu(
torch.full((seq_len, seq_len), float("-inf"), device=x.device), diagonal=1
)
scores = scores + causal
scores = F.softmax(scores, dim=-1)
scores = self.attn_dropout(scores)
out = torch.matmul(scores, v)
out = out.transpose(1, 2).contiguous().view(bs, seq_len, -1)
return self.resid_dropout(self.wo(out))
def _apply_rope_real(self, q, k, freqs):
"""Apply RoPE using real-valued sin/cos (more device-compatible)."""
# freqs shape: (seq_len, head_dim//2)
seq_len = q.shape[2]
freqs = freqs[:seq_len]
cos_f = freqs.cos().unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, head_dim//2)
sin_f = freqs.sin().unsqueeze(0).unsqueeze(0)
def rotate_half(x):
x1, x2 = x.chunk(2, dim=-1)
return torch.cat((-x2, x1), dim=-1)
q = q * cos_f.repeat(1, 1, 1, 2) + rotate_half(q) * sin_f.repeat(1, 1, 1, 2)
k = k * cos_f.repeat(1, 1, 1, 2) + rotate_half(k) * sin_f.repeat(1, 1, 1, 2)
return q, k
def reset_cache(self):
self.cache_k = None
self.cache_v = None
class SwiGLU(nn.Module):
"""SwiGLU activation — superior to GELU/ReLU, used in LLaMA."""
def __init__(self, dim: int, hidden_dim: int, dropout: float = 0.1):
super().__init__()
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))
class TransformerBlock(nn.Module):
"""Single transformer block with pre-norm architecture."""
def __init__(self, dim: int, n_heads: int, n_kv_heads: int, hidden_dim: int, dropout: float):
super().__init__()
self.attention = GroupedQueryAttention(dim, n_heads, n_kv_heads, dropout)
self.feed_forward = SwiGLU(dim, hidden_dim, dropout)
self.norm1 = RMSNorm(dim)
self.norm2 = RMSNorm(dim)
def forward(
self,
x: torch.Tensor,
freqs: torch.Tensor,
mask: Optional[torch.Tensor] = None,
use_cache: bool = False,
) -> torch.Tensor:
# Pre-norm residual connections
x = x + self.attention(self.norm1(x), freqs, mask, use_cache)
x = x + self.feed_forward(self.norm2(x))
return x
class MusicTransformer(nn.Module):
"""
LLaMA-style Transformer for music generation.
Combines: RoPE + GQA + SwiGLU + RMSNorm + gradient checkpointing.
~5M parameters with default config — suitable for training on consumer GPUs.
"""
def __init__(self, config):
super().__init__()
self.config = config
self.token_emb = nn.Embedding(config.vocab_size, config.dim)
self.dropout = nn.Dropout(config.dropout)
self.layers = nn.ModuleList([
TransformerBlock(
config.dim, config.n_heads, config.n_kv_heads,
config.hidden_dim, config.dropout,
)
for _ in range(config.n_layers)
])
self.norm = RMSNorm(config.dim)
self.output = nn.Linear(config.dim, config.vocab_size, bias=False)
# Weight tying — reduces params and improves generalization
self.token_emb.weight = self.output.weight
# Precompute RoPE frequencies
head_dim = config.dim // config.n_heads
freqs = self._precompute_freqs(head_dim, config.max_seq_len, config.rope_theta)
self.register_buffer("freqs", freqs, persistent=False)
self.grad_checkpoint = False
self._init_weights()
def _precompute_freqs(self, dim: int, max_seq_len: int, theta: float) -> torch.Tensor:
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
t = torch.arange(max_seq_len, dtype=torch.float32)
return torch.outer(t, freqs)
def _init_weights(self):
"""Xavier-style initialization for stable training."""
for module in self.modules():
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def forward(
self,
input_ids: torch.Tensor,
targets: Optional[torch.Tensor] = None,
use_cache: bool = False,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
bs, seq_len = input_ids.shape
h = self.dropout(self.token_emb(input_ids))
freqs = self.freqs[:seq_len].to(h.device)
for layer in self.layers:
if self.grad_checkpoint and self.training:
h = torch.utils.checkpoint.checkpoint(
layer, h, freqs, None, use_cache, use_reentrant=False
)
else:
h = layer(h, freqs, use_cache=use_cache)
h = self.norm(h)
logits = self.output(h)
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
targets.view(-1),
ignore_index=0, # Ignore padding
)
return logits, loss
def reset_caches(self):
for layer in self.layers:
layer.attention.reset_cache()
def count_parameters(self) -> int:
return sum(p.numel() for p in self.parameters() if p.requires_grad)
@classmethod
def from_config(cls, model_config) -> "MusicTransformer":
return cls(model_config)