""" Decoder-only Transformer with RoPE positional encoding. Target: ~30M parameters. Architecture choices (informed by MobileLLM paper + Vizuara results): - d_model = 512, n_layers = 4 → ~30M params - n_heads = 8, head_dim = 64 - FFN: SwiGLU activation (d_ffn = 4 * d_model, but gated so 2/3 effective) Actually: d_ffn = int(2/3 * 4 * d_model) rounded to nearest 64 → 1408 - RoPE positional encoding (no learned position embeddings) - RMSNorm (no bias, more stable than LayerNorm for small models) - No dropout during training (small model on small data, dropout hurts) - Causal (autoregressive) mask Parameter count breakdown (vocab=16000, d=512, layers=4): Embedding: 16000 × 512 = 8.19M Each layer: Attention: 4 × 512 × 512 = 1.05M FFN: 512×1408 + 1408×512 + 512×1408 = ~2.16M Norms: 2 × 512 negligible 4 layers: 4 × 3.21M = 12.84M Output head: tied to embedding = 0M (weight tying) TOTAL: ~21M (with weight tying) → add unembedding = ~29M without With weight tying (output = embedding.T): ~21M ← we use this This is standard practice for small models (GPT-2 style). """ import math import torch import torch.nn as nn import torch.nn.functional as F from dataclasses import dataclass # ── Config ──────────────────────────────────────────────────────────────────── @dataclass class ModelConfig: vocab_size: int = 16000 d_model: int = 512 n_layers: int = 4 n_heads: int = 8 n_kv_heads: int = 4 # GQA: 4 KV heads, 8 Q heads (reduces params) max_seq_len: int = 512 # FFN hidden dim: SwiGLU convention = 2/3 * 4 * d_model, rounded to 64 d_ffn: int = 1408 # = round(2/3 * 4 * 512 / 64) * 64 # Regularization dropout: float = 0.0 # set to 0.1 for fine-tuning if needed # RoPE rope_theta: float = 10000.0 def __post_init__(self): assert self.d_model % self.n_heads == 0 assert self.n_heads % self.n_kv_heads == 0 self.head_dim = self.d_model // self.n_heads self.n_rep = self.n_heads // self.n_kv_heads # for GQA repeat # ── RoPE ───────────────────────────────────────────────────────────────────── def precompute_rope_freqs(head_dim: int, max_seq_len: int, theta: float = 10000.0): """ Precompute RoPE frequency tensor. Returns: (max_seq_len, head_dim//2) complex tensor. """ freqs = 1.0 / ( theta ** (torch.arange(0, head_dim, 2).float() / head_dim) ) t = torch.arange(max_seq_len) freqs = torch.outer(t, freqs) return torch.polar(torch.ones_like(freqs), freqs) # complex def apply_rope(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor: """ Apply RoPE to query or key tensor. x: (batch, seq_len, n_heads, head_dim) freqs: (seq_len, head_dim//2) complex """ # Reshape to pairs for complex multiplication x_r = x.float().reshape(*x.shape[:-1], -1, 2) x_c = torch.view_as_complex(x_r) freqs = freqs[:x.shape[1]].unsqueeze(0).unsqueeze(2) # (1, seq, 1, dim//2) x_out = torch.view_as_real(x_c * freqs).flatten(-2) return x_out.type_as(x) # ── RMSNorm ─────────────────────────────────────────────────────────────────── class RMSNorm(nn.Module): def __init__(self, d_model: int, eps: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(d_model)) self.eps = eps def forward(self, x: torch.Tensor) -> torch.Tensor: norm = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return norm * self.weight # ── Attention ───────────────────────────────────────────────────────────────── class GroupedQueryAttention(nn.Module): def __init__(self, cfg: ModelConfig): super().__init__() self.n_heads = cfg.n_heads self.n_kv_heads = cfg.n_kv_heads self.n_rep = cfg.n_rep self.head_dim = cfg.head_dim self.d_model = cfg.d_model self.Wq = nn.Linear(cfg.d_model, cfg.n_heads * cfg.head_dim, bias=False) self.Wk = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False) self.Wv = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False) self.Wo = nn.Linear(cfg.n_heads * cfg.head_dim, cfg.d_model, bias=False) self.dropout = nn.Dropout(cfg.dropout) def forward( self, x: torch.Tensor, # (B, T, d_model) freqs: torch.Tensor, # (T, head_dim//2) complex mask: torch.Tensor | None = None, # (T, T) causal mask ) -> torch.Tensor: B, T, _ = x.shape q = self.Wq(x).view(B, T, self.n_heads, self.head_dim) k = self.Wk(x).view(B, T, self.n_kv_heads, self.head_dim) v = self.Wv(x).view(B, T, self.n_kv_heads, self.head_dim) # RoPE q = apply_rope(q, freqs) k = apply_rope(k, freqs) # GQA: repeat K/V to match Q heads if self.n_rep > 1: k = k.repeat_interleave(self.n_rep, dim=2) v = v.repeat_interleave(self.n_rep, dim=2) # Attention: (B, n_heads, T, head_dim) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) # Use PyTorch's flash attention when available (much faster on GPU) if hasattr(F, "scaled_dot_product_attention"): # is_causal=True handles the mask automatically and uses FlashAttention out = F.scaled_dot_product_attention( q, k, v, attn_mask=None, dropout_p=self.dropout.p if self.training else 0.0, is_causal=True, ) else: scale = self.head_dim ** -0.5 scores = torch.matmul(q, k.transpose(-2, -1)) * scale if mask is not None: scores = scores + mask scores = F.softmax(scores.float(), dim=-1).type_as(q) scores = self.dropout(scores) out = torch.matmul(scores, v) # Merge heads out = out.transpose(1, 2).contiguous().view(B, T, -1) return self.Wo(out) # ── SwiGLU FFN ──────────────────────────────────────────────────────────────── class SwiGLUFFN(nn.Module): def __init__(self, cfg: ModelConfig): super().__init__() self.gate = nn.Linear(cfg.d_model, cfg.d_ffn, bias=False) self.up = nn.Linear(cfg.d_model, cfg.d_ffn, bias=False) self.down = nn.Linear(cfg.d_ffn, cfg.d_model, bias=False) self.dropout = nn.Dropout(cfg.dropout) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.dropout(self.down(F.silu(self.gate(x)) * self.up(x))) # ── Transformer Block ───────────────────────────────────────────────────────── class TransformerBlock(nn.Module): def __init__(self, cfg: ModelConfig): super().__init__() self.attn_norm = RMSNorm(cfg.d_model) self.attn = GroupedQueryAttention(cfg) self.ffn_norm = RMSNorm(cfg.d_model) self.ffn = SwiGLUFFN(cfg) def forward( self, x: torch.Tensor, freqs: torch.Tensor, mask: torch.Tensor | None = None, ) -> torch.Tensor: # Pre-norm (LLaMA style) x = x + self.attn(self.attn_norm(x), freqs, mask) x = x + self.ffn(self.ffn_norm(x)) return x # ── Full Model ──────────────────────────────────────────────────────────────── class TinyIndianLM(nn.Module): def __init__(self, cfg: ModelConfig): super().__init__() self.cfg = cfg self.embedding = nn.Embedding(cfg.vocab_size, cfg.d_model, padding_idx=0) self.layers = nn.ModuleList([TransformerBlock(cfg) for _ in range(cfg.n_layers)]) self.norm = RMSNorm(cfg.d_model) self.output = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) # Weight tying: output projection shares weights with embedding self.output.weight = self.embedding.weight # Precompute RoPE frequencies (register as buffer → moves to device) freqs = precompute_rope_freqs(cfg.head_dim, cfg.max_seq_len, cfg.rope_theta) self.register_buffer("rope_freqs", freqs, persistent=False) # Causal mask (optional fallback when not using F.scaled_dot_product_attention) mask = torch.full((cfg.max_seq_len, cfg.max_seq_len), float("-inf")) mask = torch.triu(mask, diagonal=1) self.register_buffer("causal_mask", mask, persistent=False) # Init weights self.apply(self._init_weights) # Scale residual projections (GPT-2 style) for pn, p in self.named_parameters(): if pn.endswith(("Wo.weight", "down.weight")): nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * cfg.n_layers)) def _init_weights(self, module): 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, # (B, T) targets: torch.Tensor | None = None, # (B, T) for training pad_id: int = 0, ) -> tuple[torch.Tensor, torch.Tensor | None]: B, T = input_ids.shape assert T <= self.cfg.max_seq_len, f"Sequence length {T} > max {self.cfg.max_seq_len}" x = self.embedding(input_ids) # (B, T, d_model) freqs = self.rope_freqs[:T] for layer in self.layers: x = layer(x, freqs) x = self.norm(x) logits = self.output(x) # (B, T, vocab_size) loss = None if targets is not None: # Shift: predict token[i+1] from token[i] # input: [BOS, t1, t2, ..., tN, EOS] # targets: [t1, t2, ..., tN, EOS, PAD] # But we already have aligned input/target from dataloader # Mask out PAD tokens in the loss loss = F.cross_entropy( logits.view(-1, self.cfg.vocab_size), targets.view(-1), ignore_index=pad_id, ) return logits, loss @torch.no_grad() def generate( self, input_ids: torch.Tensor, # (1, T) prompt max_new_tokens: int = 200, temperature: float = 1.0, top_k: int = 50, eos_id: int = 3, pad_id: int = 0, ) -> list[int]: self.eval() generated = input_ids.tolist()[0] for _ in range(max_new_tokens): ids_tensor = torch.tensor([generated], device=input_ids.device) # Truncate to max_seq_len if ids_tensor.shape[1] > self.cfg.max_seq_len: ids_tensor = ids_tensor[:, -self.cfg.max_seq_len:] logits, _ = self.forward(ids_tensor) logits = logits[0, -1, :] / temperature # (vocab_size,) # Remove PAD from generation logits[pad_id] = float("-inf") # Top-k sampling if top_k > 0: top_vals, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < top_vals[-1]] = float("-inf") probs = F.softmax(logits, dim=-1) next_id = torch.multinomial(probs, num_samples=1).item() generated.append(next_id) if next_id == eos_id: break return generated def num_parameters(self, exclude_embeddings: bool = False) -> int: if exclude_embeddings: return sum(p.numel() for n, p in self.named_parameters() if "embedding" not in n and p.requires_grad) return sum(p.numel() for p in self.parameters() if p.requires_grad) # ── Quick test ──────────────────────────────────────────────────────────────── if __name__ == "__main__": cfg = ModelConfig() model = TinyIndianLM(cfg) total = model.num_parameters() print(f"Model config: d_model={cfg.d_model}, n_layers={cfg.n_layers}, " f"n_heads={cfg.n_heads}, d_ffn={cfg.d_ffn}") print(f"Total parameters: {total:,} ({total/1e6:.1f}M)") # Forward pass test B, T = 2, 64 x = torch.randint(0, cfg.vocab_size, (B, T)) logits, loss = model(x, targets=x) print(f"Logits shape: {logits.shape}") print(f"Initial loss (should be ~ln({cfg.vocab_size})={math.log(cfg.vocab_size):.2f}): {loss.item():.4f}")