# 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))