"""Transformers-compatible loading module for the TinyStoriesGPT model. Architecture (matches the exported model.safetensors exactly): - LLaMA-style decoder: RMSNorm + RoPE(theta=1e4) + SwiGLU + GQA - Tied input embedding / LM head (weight shared, no separate lm_head) - Fused qkv projection (one Linear producing q,k,v) Tensor names in the safetensors (must match these attribute names): tok.weight, layers.N.attn.qkv.weight, layers.N.attn.proj.weight, layers.N.ln1.w, layers.N.ln2.w, layers.N.mlp.{gate,up,down}.weight, ln_f.w """ import math import torch import torch.nn as nn import torch.nn.functional as F from dataclasses import dataclass from typing import Optional, Tuple from transformers import PreTrainedModel class RMSNorm(nn.Module): def __init__(self, dim, eps=1e-6): super().__init__() self.w = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x): return self.w * (x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps)).type_as(x) class CausalSelfAttn(nn.Module): def __init__(self, c): super().__init__() self.nq = c["n_head_q"] self.nkv = c["n_head_kv"] self.hd = c["head_dim"] self.qkv = nn.Linear(c["n_embd"], (self.nq + 2 * self.nkv) * self.hd, bias=False) self.proj = nn.Linear(self.nq * self.hd, c["n_embd"], bias=False) def forward(self, x, cos, sin, kv_cache=None): B, T, _ = x.shape cos = cos[:, :, :T] sin = sin[:, :, :T] qkv = self.qkv(x).view(B, T, self.nq + 2 * self.nkv, self.hd).transpose(1, 2) q, k, v = qkv.split([self.nq, self.nkv, self.nkv], dim=1) def _rope(t): t1 = t[..., 0::2] t2 = t[..., 1::2] return torch.stack([t1 * cos - t2 * sin, t2 * cos + t1 * sin], dim=-1).reshape(t.shape) q = _rope(q) k = _rope(k) if kv_cache is not None: k = torch.cat([kv_cache[0], k], dim=2) v = torch.cat([kv_cache[1], v], dim=2) y = F.scaled_dot_product_attention(q, k, v, is_causal=kv_cache is None, enable_gqa=True) y = y.transpose(1, 2).contiguous().view(B, T, -1) return self.proj(y), (k, v) class MLP(nn.Module): def __init__(self, c): super().__init__() self.gate = nn.Linear(c["n_embd"], c["ffn"], bias=False) self.down = nn.Linear(c["ffn"], c["n_embd"], bias=False) self.up = nn.Linear(c["n_embd"], c["ffn"], bias=False) def forward(self, x): return self.down(F.silu(self.gate(x)) * self.up(x)) class Block(nn.Module): def __init__(self, c): super().__init__() self.ln1 = RMSNorm(c["n_embd"]) self.attn = CausalSelfAttn(c) self.ln2 = RMSNorm(c["n_embd"]) self.mlp = MLP(c) def forward(self, x, cos, sin, kv_cache=None): a, kv = self.attn(self.ln1(x), cos, sin, kv_cache) x = x + a x = x + self.mlp(self.ln2(x)) return x, kv @dataclass class Output: logits: torch.Tensor past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None class TinyStoriesGPT(PreTrainedModel): """Decoder-only GQA LLaMA with tied embedding/head. Reads the standard transformers config attributes from config.json: num_hidden_layers, hidden_size, num_attention_heads, num_key_value_heads, head_dim, intermediate_size, vocab_size, rope_theta, max_position_embeddings """ def __init__(self, config): super().__init__(config) # rope_theta: top-level attr, or nested under rope_scaling (newer transformers) if hasattr(config, "rope_theta") and config.rope_theta is not None: rope_theta = config.rope_theta else: rs = getattr(config, "rope_scaling", None) or {} rope_theta = rs.get("rope_theta", 10000.0) c = dict( n_layer=config.num_hidden_layers, n_embd=config.hidden_size, n_head_q=config.num_attention_heads, n_head_kv=config.num_key_value_heads, head_dim=config.head_dim, ffn=config.intermediate_size, rope_theta=rope_theta, vocab=config.vocab_size, ) self.c = c self.config = config self.tok = nn.Embedding(c["vocab"], c["n_embd"]) self.layers = nn.ModuleList([Block(c) for _ in range(c["n_layer"])]) self.ln_f = RMSNorm(c["n_embd"]) # No separate LM head: logits are computed directly from the input # embedding (weight tying). This keeps the parameter set identical to # the exported model.safetensors (which has no head.weight key). self.max_seq = config.max_position_embeddings self._rope_cache = {} # No tied weights in this model (logits come directly from the input # embedding). Make sure the transformers 5.x tied-weight machinery sees # an empty mapping so from_pretrained finalization does not trip. self._tied_weights_keys = {} self.all_tied_weights_keys = {} def _rope(self, T, device): key = (T, device) if key not in self._rope_cache: c = self.c freqs = 1.0 / (c["rope_theta"] ** (torch.arange(0, c["head_dim"], 2, device=device).float() / c["head_dim"])) t = torch.arange(T, device=device).float() ang = torch.outer(t, freqs) self._rope_cache[key] = (ang.cos()[None, None], ang.sin()[None, None]) return self._rope_cache[key] def forward(self, input_ids, attention_mask=None, past_key_values=None, use_cache=False): B = input_ids.shape[0] T = input_ids.shape[1] start = 0 if past_key_values is None else past_key_values[0][0].shape[2] cos, sin = self._rope(start + T, input_ids.device) x = self.tok(input_ids) kvs = [] for i, blk in enumerate(self.layers): kc = past_key_values[i] if past_key_values else None x, kv = blk(x, cos, sin, kc) kvs.append(kv) x = self.ln_f(x) logits = F.linear(x, self.tok.weight) # tied head: embedding as output projection return Output(logits=logits, past_key_values=tuple(kvs) if use_cache else None) @torch.no_grad() def generate(self, input_ids, max_new_tokens=100, temperature=0.8, top_k=0, eos_token_id=None): if eos_token_id is None: eos_token_id = 2 # device = input_ids.shape[0] and input_ids.device B = input_ids.shape[0] ids = input_ids past = None for _ in range(max_new_tokens): out = self(ids, past_key_values=past, use_cache=True) past = out.past_key_values logits = out.logits[:, -1, :] if temperature and temperature != 1.0: logits = logits / temperature if top_k and top_k > 0: v, _ = torch.topk(logits, top_k) logits[logits < v[:, [-1]]] = float("-inf") probs = F.softmax(logits, dim=-1) nxt = torch.multinomial(probs, 1) ids = torch.cat([ids, nxt], dim=1) if (nxt.item() == eos_token_id) or (B == 1 and nxt.item() == eos_token_id): break return ids