Download model.py from AdvaithMagic/tinyindianlm-ml: direct link, hf CLI and curl.
- Browser
- Download file 13.6 kB
-
https://huggingface.co/AdvaithMagic/tinyindianlm-ml/resolve/main/model.py
- Command line
-
hf download hf://AdvaithMagic/tinyindianlm-ml/model.py
-
curl -L -o model.py https://huggingface.co/AdvaithMagic/tinyindianlm-ml/resolve/main/model.py
13.6 kB
| """ | |
| 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 ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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 | |
| 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}") | |