import math import torch import torch.nn as nn import torch.nn.functional as F from dataclasses import dataclass from typing import Optional @dataclass class ModelConfig: vocab_size: int = 32000 dim: int = 768 n_layers: int = 12 n_heads: int = 12 n_kv_heads: Optional[int] = 12 # For Multi-Query / Grouped Query Attention multiple_of: int = 256 # For SwiGLU hidden dim alignment ffn_dim_multiplier: Optional[float] = None norm_eps: float = 1e-5 max_seq_len: int = 2048 dropout: float = 0.0 @classmethod def get_125m(cls, vocab_size: int = 32000): # ~125 Million Parameters return cls( vocab_size=vocab_size, dim=768, n_layers=12, n_heads=12, n_kv_heads=12, max_seq_len=2048 ) @classmethod def get_350m(cls, vocab_size: int = 32000): # ~350 Million Parameters return cls( vocab_size=vocab_size, dim=1024, n_layers=24, n_heads=16, n_kv_heads=16, max_seq_len=2048 ) class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: variance = x.pow(2).mean(-1, keepdim=True) return x * torch.rsqrt(variance + self.eps) * self.weight def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0) -> torch.Tensor: freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(end, device=freqs.device) freqs = torch.outer(t, freqs).float() freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 return freqs_cis def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor) -> torch.Tensor: ndim = x.ndim assert 0 <= 1 < ndim assert freqs_cis.shape == (x.shape[1], x.shape[-1]) shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)] return freqs_cis.view(*shape) def apply_rotary_emb( xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2)) freqs_cis = reshape_for_broadcast(freqs_cis, xq_) xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3) xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3) return xq_out.type_as(xq), xk_out.type_as(xk) class FeedForward(nn.Module): def __init__(self, config: ModelConfig): super().__init__() hidden_dim = int(2 * (4 * config.dim) / 3) if config.ffn_dim_multiplier is not None: hidden_dim = int(config.ffn_dim_multiplier * hidden_dim) hidden_dim = config.multiple_of * ((hidden_dim + config.multiple_of - 1) // config.multiple_of) self.w1 = nn.Linear(config.dim, hidden_dim, bias=False) self.w2 = nn.Linear(hidden_dim, config.dim, bias=False) self.w3 = nn.Linear(config.dim, hidden_dim, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.w2(F.silu(self.w1(x)) * self.w3(x)) class Attention(nn.Module): def __init__(self, config: ModelConfig): super().__init__() self.n_kv_heads = config.n_heads if config.n_kv_heads is None else config.n_kv_heads self.n_heads = config.n_heads self.head_dim = config.dim // config.n_heads self.n_rep = self.n_heads // self.n_kv_heads self.wq = nn.Linear(config.dim, config.n_heads * self.head_dim, bias=False) self.wk = nn.Linear(config.dim, self.n_kv_heads * self.head_dim, bias=False) self.wv = nn.Linear(config.dim, self.n_kv_heads * self.head_dim, bias=False) self.wo = nn.Linear(config.n_heads * self.head_dim, config.dim, bias=False) self.dropout = config.dropout def forward(self, x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: bsz, seqlen, _ = x.shape xq, xk, xv = self.wq(x), self.wk(x), self.wv(x) xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim) xk = xk.view(bsz, seqlen, self.n_kv_heads, self.head_dim) xv = xv.view(bsz, seqlen, self.n_kv_heads, self.head_dim) xq, xk = apply_rotary_emb(xq, xk, freqs_cis=freqs_cis) if self.n_rep > 1: xk = xk.repeat_interleave(self.n_rep, dim=2) xv = xv.repeat_interleave(self.n_rep, dim=2) # Transpose for PyTorch Scaled Dot Product Attention [bsz, n_heads, seqlen, head_dim] xq = xq.transpose(1, 2) xk = xk.transpose(1, 2) xv = xv.transpose(1, 2) # Fast FlashAttention / SDPA kernel output = F.scaled_dot_product_attention( xq, xk, xv, is_causal=True, dropout_p=self.dropout if self.training else 0.0 ) output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1) return self.wo(output) class TransformerBlock(nn.Module): def __init__(self, layer_id: int, config: ModelConfig): super().__init__() self.layer_id = layer_id self.attention = Attention(config) self.feed_forward = FeedForward(config) self.attention_norm = RMSNorm(config.dim, eps=config.norm_eps) self.ffn_norm = RMSNorm(config.dim, eps=config.norm_eps) def forward(self, x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: h = x + self.attention(self.attention_norm(x), freqs_cis) out = h + self.feed_forward(self.ffn_norm(h)) return out class Transformer(nn.Module): def __init__(self, config: ModelConfig): super().__init__() self.config = config self.tok_embeddings = nn.Embedding(config.vocab_size, config.dim) self.layers = nn.ModuleList([TransformerBlock(i, config) for i in range(config.n_layers)]) self.norm = RMSNorm(config.dim, eps=config.norm_eps) self.output = nn.Linear(config.dim, config.vocab_size, bias=False) # Weight tying (optional, but standard for small models) self.tok_embeddings.weight = self.output.weight # Precompute rotary frequencies freqs_cis = precompute_freqs_cis(config.dim // config.n_heads, config.max_seq_len * 2) self.register_buffer("freqs_cis", freqs_cis, persistent=False) def forward(self, tokens: torch.Tensor) -> torch.Tensor: _bsz, seqlen = tokens.shape h = self.tok_embeddings(tokens) freqs_cis = self.freqs_cis[:seqlen] for layer in self.layers: h = layer(h, freqs_cis) h = self.norm(h) logits = self.output(h) return logits def count_parameters(self) -> int: return sum(p.numel() for p in self.parameters() if p.requires_grad) if __name__ == "__main__": cfg = ModelConfig.get_125m() model = Transformer(cfg) print(f"Initialized 125M Model. Parameter Count: {model.count_parameters():,}") cfg_350m = ModelConfig.get_350m() model_350m = Transformer(cfg_350m) print(f"Initialized 350M Model. Parameter Count: {model_350m.count_parameters():,}")