"""Standalone architecture for the Harmonic GPT 128M byte transformer.""" from __future__ import annotations import math import torch import torch.nn as nn import torch.nn.functional as F def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: cos = cos[: x.shape[-2]].to(device=x.device, dtype=x.dtype)[None, None, :, :] sin = sin[: x.shape[-2]].to(device=x.device, dtype=x.dtype)[None, None, :, :] even = x[..., 0::2] odd = x[..., 1::2] return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2) class ModernSelfAttention(nn.Module): def __init__(self, d_model: int, n_heads: int, max_seq_len: int, rope_base: float) -> None: super().__init__() if d_model % n_heads: raise ValueError(f"d_model={d_model} must be divisible by n_heads={n_heads}") self.n_heads = n_heads self.head_dim = d_model // n_heads if self.head_dim % 2: raise ValueError(f"RoPE requires an even head dimension, got {self.head_dim}") self.qkv_proj = nn.Linear(d_model, 3 * d_model, bias=False) self.out_proj = nn.Linear(d_model, d_model, bias=False) inv_freq = 1.0 / ( rope_base ** (torch.arange(0, self.head_dim, 2, dtype=torch.float32) / self.head_dim) ) positions = torch.arange(max_seq_len, dtype=torch.float32) angles = torch.outer(positions, inv_freq) self.register_buffer("rope_cos", angles.cos(), persistent=False) self.register_buffer("rope_sin", angles.sin(), persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: batch, seq_len, d_model = x.shape qkv = self.qkv_proj(x).view(batch, seq_len, 3, self.n_heads, self.head_dim) q, k, v = qkv.unbind(dim=2) q = _apply_rope(q.transpose(1, 2), self.rope_cos, self.rope_sin) k = _apply_rope(k.transpose(1, 2), self.rope_cos, self.rope_sin) v = v.transpose(1, 2) attended = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=True) return self.out_proj(attended.transpose(1, 2).contiguous().view(batch, seq_len, d_model)) class SwiGLU(nn.Module): def __init__(self, d_model: int, hidden_dim: int) -> None: super().__init__() self.gate_proj = nn.Linear(d_model, hidden_dim, bias=False) self.up_proj = nn.Linear(d_model, hidden_dim, bias=False) self.down_proj = nn.Linear(hidden_dim, d_model, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class ModernTransformerBlock(nn.Module): def __init__( self, d_model: int, n_heads: int, hidden_dim: int, max_seq_len: int, rope_base: float, rms_norm_eps: float, ) -> None: super().__init__() self.attn_norm = nn.RMSNorm(d_model, eps=rms_norm_eps) self.attn = ModernSelfAttention(d_model, n_heads, max_seq_len, rope_base) self.ffn_norm = nn.RMSNorm(d_model, eps=rms_norm_eps) self.ffn = SwiGLU(d_model, hidden_dim) def forward(self, x: torch.Tensor) -> torch.Tensor: x = x + self.attn(self.attn_norm(x)) return x + self.ffn(self.ffn_norm(x)) class ModernByteTransformer(nn.Module): """LLaMA/OLMo-style causal decoder over the 256 raw byte values.""" def __init__( self, vocab_size: int = 256, d_model: int = 768, n_layers: int = 18, n_heads: int = 12, hidden_dim: int = 2048, max_seq_len: int = 2048, rope_base: float = 10_000.0, rms_norm_eps: float = 1e-5, ) -> None: super().__init__() self.max_seq_len = max_seq_len self.tok_emb = nn.Embedding(vocab_size, d_model) self.blocks = nn.ModuleList( [ ModernTransformerBlock( d_model, n_heads, hidden_dim, max_seq_len, rope_base, rms_norm_eps, ) for _ in range(n_layers) ] ) self.norm = nn.RMSNorm(d_model, eps=rms_norm_eps) self.lm_head = nn.Linear(d_model, vocab_size, bias=False) self.lm_head.weight = self.tok_emb.weight self.apply(self._init_weights) residual_std = 0.02 / math.sqrt(2 * n_layers) for block in self.blocks: nn.init.normal_(block.attn.out_proj.weight, mean=0.0, std=residual_std) nn.init.normal_(block.ffn.down_proj.weight, mean=0.0, std=residual_std) @staticmethod def _init_weights(module: nn.Module) -> None: if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=0.02) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) elif isinstance(module, nn.RMSNorm) and module.weight is not None: nn.init.ones_(module.weight) def forward(self, idx: torch.Tensor) -> torch.Tensor: _, seq_len = idx.shape if seq_len > self.max_seq_len: raise ValueError(f"seq_len={seq_len} exceeds max_seq_len={self.max_seq_len}") x = self.tok_emb(idx) for block in self.blocks: x = block(x) return self.lm_head(self.norm(x))