tinyindianlm-ml / model.py
AdvaithMagic's picture
Export TinyIndianLM checkpoint in HF format
a15230d verified
Raw History Blame Contribute Delete
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 ────────────────────────────────────────────────────────────────────
@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}")