rune-0.3b-sft / model.py
samueljayasingh's picture
Upload Rune-R1 Base 351M model checkpoint and model code
84c5892 verified
Raw
History Blame Contribute Delete
8.92 kB
# Architecture adapted from rasbt/LLMs-from-scratch pkg/llms_from_scratch/qwen3.py (Apache 2.0):
# RoPE + RMSNorm + SwiGLU + grouped-query attention, trimmed to a dense ~350M config
# with a GPT-2 (tiktoken) vocab instead of Qwen's tokenizer/MoE variants.
import torch
import torch.nn as nn
CONFIG_350M = {
"vocab_size": 50257, # tiktoken gpt2
"context_length": 1024,
"emb_dim": 1024,
"n_heads": 16,
"n_layers": 22,
"hidden_dim": 2816,
"head_dim": None, # defaults to emb_dim // n_heads
"qk_norm": True,
"n_kv_groups": 4,
"rope_base": 10_000.0,
"dtype": torch.float32, # fp32 master weights; train.py autocasts to bf16 for compute
}
class RMSNorm(nn.Module):
def __init__(self, emb_dim, eps=1e-6):
super().__init__()
self.eps = eps
self.scale = nn.Parameter(torch.ones(emb_dim))
def forward(self, x):
input_dtype = x.dtype
x = x.to(torch.float32)
variance = x.pow(2).mean(dim=-1, keepdim=True)
norm_x = x * torch.rsqrt(variance + self.eps) * self.scale
return norm_x.to(input_dtype)
def compute_rope_params(head_dim, theta_base, context_length, dtype=torch.float32):
assert head_dim % 2 == 0, "Head dimension must be even"
inv_freq = 1.0 / (theta_base ** (torch.arange(0, head_dim, 2, dtype=dtype) / head_dim))
positions = torch.arange(context_length, dtype=dtype)
angles = positions.unsqueeze(1) * inv_freq.unsqueeze(0)
angles = torch.cat([angles, angles], dim=1)
return torch.cos(angles), torch.sin(angles)
def apply_rope(x, cos, sin, offset=0):
# x: (batch, heads, seq_len, head_dim). `offset` is the absolute position
# of x[..., 0, :] — nonzero when x is a new chunk appended after cached
# positions, so rotation angles pick up where the cache left off.
head_dim = x.shape[-1]
x1, x2 = x[..., : head_dim // 2], x[..., head_dim // 2:]
seq_len = x.shape[2]
max_pos = cos.shape[0]
if offset + seq_len > max_pos:
offset = max(0, max_pos - seq_len)
cos = cos[offset:offset + seq_len].unsqueeze(0).unsqueeze(0)
sin = sin[offset:offset + seq_len].unsqueeze(0).unsqueeze(0)
rotated = torch.cat((-x2, x1), dim=-1)
return ((x * cos) + (rotated * sin)).to(dtype=x.dtype)
def new_kv_cache(n_layers):
"""One mutable dict per layer; GroupedQueryAttention fills in 'k'/'v' and
grows them in place across calls sharing the same cache list."""
return [dict() for _ in range(n_layers)]
class GroupedQueryAttention(nn.Module):
def __init__(self, d_in, num_heads, num_kv_groups, head_dim=None, qk_norm=False, dtype=None):
super().__init__()
assert num_heads % num_kv_groups == 0, "num_heads must be divisible by num_kv_groups"
if head_dim is None:
assert d_in % num_heads == 0
head_dim = d_in // num_heads
self.num_heads = num_heads
self.num_kv_groups = num_kv_groups
self.group_size = num_heads // num_kv_groups
self.head_dim = head_dim
self.d_out = num_heads * head_dim
self.W_query = nn.Linear(d_in, self.d_out, bias=False, dtype=dtype)
self.W_key = nn.Linear(d_in, num_kv_groups * head_dim, bias=False, dtype=dtype)
self.W_value = nn.Linear(d_in, num_kv_groups * head_dim, bias=False, dtype=dtype)
self.out_proj = nn.Linear(self.d_out, d_in, bias=False, dtype=dtype)
self.q_norm = RMSNorm(head_dim) if qk_norm else None
self.k_norm = RMSNorm(head_dim) if qk_norm else None
def forward(self, x, mask, cos, sin, cache=None):
b, num_tokens, _ = x.shape
queries = self.W_query(x).view(b, num_tokens, self.num_heads, self.head_dim).transpose(1, 2)
keys = self.W_key(x).view(b, num_tokens, self.num_kv_groups, self.head_dim).transpose(1, 2)
values = self.W_value(x).view(b, num_tokens, self.num_kv_groups, self.head_dim).transpose(1, 2)
if self.q_norm:
queries = self.q_norm(queries)
if self.k_norm:
keys = self.k_norm(keys)
past_len = 0 if cache is None or cache.get("k") is None else cache["k"].shape[2]
queries = apply_rope(queries, cos, sin, offset=past_len)
keys = apply_rope(keys, cos, sin, offset=past_len)
if cache is not None:
if cache.get("k") is not None:
keys = torch.cat([cache["k"], keys], dim=2)
values = torch.cat([cache["v"], values], dim=2)
cache["k"], cache["v"] = keys, values
keys = keys.repeat_interleave(self.group_size, dim=1)
values = values.repeat_interleave(self.group_size, dim=1)
if past_len == 0:
# No cache, or first (prefill) call on an empty cache: query and
# key spans are identical, standard causal mask applies.
context = nn.functional.scaled_dot_product_attention(
queries, keys, values, attn_mask=None, is_causal=True
)
elif num_tokens == 1:
# Single-token decode step: this query is always the newest
# position, so it may attend to every cached key — no mask needed.
context = nn.functional.scaled_dot_product_attention(
queries, keys, values, attn_mask=None, is_causal=False
)
else:
raise NotImplementedError("cache only supports prefill-then-single-token decode")
context = context.transpose(1, 2).reshape(b, num_tokens, self.d_out)
return self.out_proj(context)
class FeedForward(nn.Module):
def __init__(self, cfg):
super().__init__()
self.fc1 = nn.Linear(cfg["emb_dim"], cfg["hidden_dim"], dtype=cfg["dtype"], bias=False)
self.fc2 = nn.Linear(cfg["emb_dim"], cfg["hidden_dim"], dtype=cfg["dtype"], bias=False)
self.fc3 = nn.Linear(cfg["hidden_dim"], cfg["emb_dim"], dtype=cfg["dtype"], bias=False)
def forward(self, x):
return self.fc3(nn.functional.silu(self.fc1(x)) * self.fc2(x))
class TransformerBlock(nn.Module):
def __init__(self, cfg):
super().__init__()
self.att = GroupedQueryAttention(
d_in=cfg["emb_dim"], num_heads=cfg["n_heads"], head_dim=cfg["head_dim"],
num_kv_groups=cfg["n_kv_groups"], qk_norm=cfg["qk_norm"], dtype=cfg["dtype"],
)
self.ff = FeedForward(cfg)
self.norm1 = RMSNorm(cfg["emb_dim"])
self.norm2 = RMSNorm(cfg["emb_dim"])
def forward(self, x, mask, cos, sin, cache=None):
x = x + self.att(self.norm1(x), mask, cos, sin, cache)
x = x + self.ff(self.norm2(x))
return x
class RuneModel(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.tok_emb = nn.Embedding(cfg["vocab_size"], cfg["emb_dim"], dtype=cfg["dtype"])
self.trf_blocks = nn.ModuleList(TransformerBlock(cfg) for _ in range(cfg["n_layers"]))
self.final_norm = RMSNorm(cfg["emb_dim"])
self.out_head = nn.Linear(cfg["emb_dim"], cfg["vocab_size"], bias=False, dtype=cfg["dtype"])
head_dim = cfg["head_dim"] or cfg["emb_dim"] // cfg["n_heads"]
cos, sin = compute_rope_params(head_dim, cfg["rope_base"], cfg["context_length"])
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)
def forward(self, in_idx, cache=None):
x = self.tok_emb(in_idx)
for i, block in enumerate(self.trf_blocks):
x = block(x, None, self.cos, self.sin, cache[i] if cache is not None else None)
x = self.final_norm(x)
return self.out_head(x.to(self.cfg["dtype"]))
def _test_kv_cache_matches_full_forward(cfg):
torch.manual_seed(0)
model = RuneModel(cfg).eval()
seq = torch.randint(0, cfg["vocab_size"], (2, 12))
with torch.no_grad():
full_logits = model(seq)
cache = new_kv_cache(cfg["n_layers"])
chunks = [model(seq[:, :5], cache=cache)]
for i in range(5, 12):
chunks.append(model(seq[:, i:i + 1], cache=cache))
cached_logits = torch.cat(chunks, dim=1)
assert cached_logits.shape == full_logits.shape
max_diff = (full_logits - cached_logits).abs().max().item()
assert torch.allclose(full_logits, cached_logits, atol=1e-4), f"max diff {max_diff}"
print(f"kv-cache self-test ok (max diff vs full forward: {max_diff:.2e})")
if __name__ == "__main__":
cfg = CONFIG_350M
model = RuneModel(cfg)
n_params = sum(p.numel() for p in model.parameters())
print(f"params: {n_params:,} ({n_params / 1e6:.1f}M)")
x = torch.randint(0, cfg["vocab_size"], (2, 16))
logits = model(x)
assert logits.shape == (2, 16, cfg["vocab_size"]), logits.shape
assert torch.isfinite(logits).all()
print("forward pass ok:", logits.shape)
_test_kv_cache_matches_full_forward(dict(cfg, n_layers=2, context_length=64))