| """ |
| RuqLM architecture — a small transformer trained from scratch. |
| |
| This is the model itself: randomly initialised weights, not derived from any |
| pretrained checkpoint. The recipe is modern and standard: pre-norm, RMSNorm, |
| RoPE, SwiGLU, tied embeddings. |
| |
| Why these choices at 30M parameters specifically: |
| - RMSNorm over LayerNorm: fewer operations, no measurable quality cost. |
| - RoPE over learned positional embeddings: no extra parameters, and better |
| generalisation to lengths not seen during training. |
| - SwiGLU: better than ReLU/GELU at a fixed parameter count. |
| - Tying input and output embeddings: saves 4.2M parameters, 14% of the model |
| at a vocabulary of 8192. At this size that is structural, not a marginal |
| optimisation. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import math |
| from dataclasses import dataclass, asdict |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| @dataclass |
| class ModelArgs: |
| vocab_size: int = 8192 |
| d_model: int = 512 |
| n_layers: int = 8 |
| n_heads: int = 8 |
| n_kv_heads: int | None = None |
| ffn_hidden: int | None = None |
| max_seq_len: int = 512 |
| rope_theta: float = 10000.0 |
| norm_eps: float = 1e-5 |
| dropout: float = 0.0 |
| tie_embeddings: bool = True |
|
|
| def __post_init__(self) -> None: |
| if self.d_model % self.n_heads: |
| raise ValueError("d_model must be divisible by n_heads") |
| if self.n_kv_heads is None: |
| self.n_kv_heads = self.n_heads |
| if self.n_heads % self.n_kv_heads: |
| raise ValueError("n_heads must be divisible by n_kv_heads") |
| if self.ffn_hidden is None: |
| |
| |
| self.ffn_hidden = 64 * math.ceil((8 * self.d_model / 3) / 64) |
|
|
| @property |
| def head_dim(self) -> int: |
| return self.d_model // self.n_heads |
|
|
| def to_dict(self) -> dict: |
| return asdict(self) |
|
|
|
|
| |
| 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: |
| |
| |
| dtype = x.dtype |
| x = x.float() |
| x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) |
| return (x * self.weight.float()).to(dtype) |
|
|
|
|
| def build_rope_cache(seq_len: int, head_dim: int, theta: float, device, dtype): |
| """Returns (cos, sin), each of shape (seq_len, head_dim/2).""" |
| inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim)) |
| pos = torch.arange(seq_len, device=device).float() |
| freqs = torch.outer(pos, inv_freq) |
| return freqs.cos().to(dtype), freqs.sin().to(dtype) |
|
|
|
|
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: |
| """x is (B, H, S, D) — rotates each coordinate pair by an angle set by position.""" |
| x1, x2 = x.chunk(2, dim=-1) |
| cos = cos[None, None, : x.size(-2), :] |
| sin = sin[None, None, : x.size(-2), :] |
| return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1) |
|
|
|
|
| class Attention(nn.Module): |
| def __init__(self, args: ModelArgs): |
| super().__init__() |
| self.n_heads, self.n_kv_heads = args.n_heads, args.n_kv_heads |
| self.head_dim = args.head_dim |
| self.repeat = self.n_heads // self.n_kv_heads |
| self.dropout = args.dropout |
|
|
| self.wq = nn.Linear(args.d_model, self.n_heads * self.head_dim, bias=False) |
| self.wk = nn.Linear(args.d_model, self.n_kv_heads * self.head_dim, bias=False) |
| self.wv = nn.Linear(args.d_model, self.n_kv_heads * self.head_dim, bias=False) |
| self.wo = nn.Linear(self.n_heads * self.head_dim, args.d_model, bias=False) |
|
|
| def forward(self, x, cos, sin): |
| B, S, _ = x.shape |
| q = self.wq(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2) |
| k = self.wk(x).view(B, S, self.n_kv_heads, self.head_dim).transpose(1, 2) |
| v = self.wv(x).view(B, S, self.n_kv_heads, self.head_dim).transpose(1, 2) |
|
|
| q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin) |
|
|
| if self.repeat > 1: |
| k = k.repeat_interleave(self.repeat, dim=1) |
| v = v.repeat_interleave(self.repeat, dim=1) |
|
|
| out = F.scaled_dot_product_attention( |
| q, k, v, is_causal=True, |
| dropout_p=self.dropout if self.training else 0.0, |
| ) |
| return self.wo(out.transpose(1, 2).contiguous().view(B, S, -1)) |
|
|
|
|
| class SwiGLU(nn.Module): |
| def __init__(self, args: ModelArgs): |
| super().__init__() |
| h = args.ffn_hidden |
| self.w_gate = nn.Linear(args.d_model, h, bias=False) |
| self.w_up = nn.Linear(args.d_model, h, bias=False) |
| self.w_down = nn.Linear(h, args.d_model, bias=False) |
|
|
| def forward(self, x): |
| return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x)) |
|
|
|
|
| class Block(nn.Module): |
| def __init__(self, args: ModelArgs): |
| super().__init__() |
| self.attn_norm = RMSNorm(args.d_model, args.norm_eps) |
| self.attn = Attention(args) |
| self.ffn_norm = RMSNorm(args.d_model, args.norm_eps) |
| self.ffn = SwiGLU(args) |
| self.drop = nn.Dropout(args.dropout) |
|
|
| def forward(self, x, cos, sin): |
| x = x + self.drop(self.attn(self.attn_norm(x), cos, sin)) |
| return x + self.drop(self.ffn(self.ffn_norm(x))) |
|
|
|
|
| |
| class RuqLM(nn.Module): |
| def __init__(self, args: ModelArgs): |
| super().__init__() |
| self.args = args |
| self.tok_emb = nn.Embedding(args.vocab_size, args.d_model) |
| self.drop = nn.Dropout(args.dropout) |
| self.blocks = nn.ModuleList(Block(args) for _ in range(args.n_layers)) |
| self.norm = RMSNorm(args.d_model, args.norm_eps) |
| self.lm_head = nn.Linear(args.d_model, args.vocab_size, bias=False) |
|
|
| if args.tie_embeddings: |
| self.lm_head.weight = self.tok_emb.weight |
|
|
| self.apply(self._init) |
| |
| |
| |
| std = 0.02 / math.sqrt(2 * args.n_layers) |
| for block in self.blocks: |
| nn.init.normal_(block.attn.wo.weight, mean=0.0, std=std) |
| nn.init.normal_(block.ffn.w_down.weight, mean=0.0, std=std) |
|
|
| self._cache_key = None |
|
|
| @staticmethod |
| def _init(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 _rope(self, seq_len: int, device, dtype): |
| key = (seq_len, device, dtype) |
| if self._cache_key != key: |
| self._cos, self._sin = build_rope_cache( |
| max(seq_len, self.args.max_seq_len), self.args.head_dim, |
| self.args.rope_theta, device, dtype, |
| ) |
| self._cache_key = key |
| return self._cos[:seq_len], self._sin[:seq_len] |
|
|
| def forward(self, input_ids: torch.Tensor, labels: torch.Tensor | None = None): |
| x = self.drop(self.tok_emb(input_ids)) |
| cos, sin = self._rope(input_ids.size(1), x.device, x.dtype) |
| for block in self.blocks: |
| x = block(x, cos, sin) |
| logits = self.lm_head(self.norm(x)) |
|
|
| loss = None |
| if labels is not None: |
| |
| loss = F.cross_entropy( |
| logits[:, :-1].reshape(-1, logits.size(-1)).float(), |
| labels[:, 1:].reshape(-1), |
| ignore_index=-100, |
| ) |
| return logits, loss |
|
|
| |
| def num_params(self, embeddings: bool = True) -> int: |
| """Tied embeddings are counted once (lm_head.weight is tok_emb.weight).""" |
| seen, total = set(), 0 |
| for name, p in self.named_parameters(): |
| if id(p) in seen: |
| continue |
| seen.add(id(p)) |
| if not embeddings and "tok_emb" in name: |
| continue |
| total += p.numel() |
| return total |
|
|
| @torch.no_grad() |
| def generate(self, input_ids, max_new_tokens=128, temperature=0.8, |
| top_k=50, eos_id=None): |
| """Plain sampling without a KV cache — adequate for short sequences.""" |
| self.eval() |
| for _ in range(max_new_tokens): |
| window = input_ids[:, -self.args.max_seq_len:] |
| logits, _ = self(window) |
| logits = logits[:, -1, :].float() / max(temperature, 1e-6) |
| if top_k: |
| kth = logits.topk(min(top_k, logits.size(-1)), dim=-1).values[:, -1:] |
| logits = logits.masked_fill(logits < kth, float("-inf")) |
| nxt = torch.multinomial(logits.softmax(-1), num_samples=1) |
| input_ids = torch.cat([input_ids, nxt], dim=1) |
| if eos_id is not None and (nxt == eos_id).all(): |
| break |
| return input_ids |
|
|