tinystories-40m / modeling_tinystories.py
Compactbot's picture
Add model class (TinyStoriesGPT, trust_remote_code)
2c168a1 verified
Raw History Blame Contribute Delete
7.29 kB
"""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 # </s>
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