"""Slayer139 1.01 model code (inference only), extracted from the training script used for this run. The module definitions are those of the training script; only the weight initialisation and the training loss were removed and comments translated. Weights load with strict=True. Requirements: torch, safetensors, tokenizers. from modeling_slayer139 import load_slayer139, generate model, tok = load_slayer139(".") # directory with config.json, model.safetensors, tokenizer.json print(generate(model, tok, "The capital of France is", max_new_tokens=40)) Notes: RoPE uses the interleaved convention (pairs of adjacent channels), not rotate_half; the output head is tied to the token embedding; context length is 1,024 tokens. """ from __future__ import annotations import json import os import types import torch import torch.nn as nn import torch.nn.functional as F class RMSNorm(nn.Module): """RMSNorm with fp32 compute for stability.""" def __init__(self, d, eps=1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(d)) self.eps = eps def forward(self, x): return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight def make_norm(d, cfg): return RMSNorm(d, cfg.norm_eps) if cfg.norm == "rmsnorm" else nn.LayerNorm(d) def apply_rope(x, base=100000.0): """Parameter-free RoPE on [B, H, T, D], interleaved convention (even/odd channel pairs).""" _, _, T, dim = x.shape pos = torch.arange(T, device=x.device, dtype=torch.float32) freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=x.device, dtype=torch.float32) / dim)) ang = torch.outer(pos, freq) cos, sin = ang.cos().to(x.dtype)[None, None], ang.sin().to(x.dtype)[None, None] even, odd = x[..., ::2], x[..., 1::2] return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2) class SwiGLU(nn.Module): """Gated MLP: down(silu(gate(x)) * up(x)).""" def __init__(self, d, hidden): super().__init__() self.gate = nn.Linear(d, hidden, bias=False) self.up = nn.Linear(d, hidden, bias=False) self.down = nn.Linear(hidden, d, bias=False) def forward(self, x): return self.down(F.silu(self.gate(x)) * self.up(x)) class Block(nn.Module): def __init__(self, d, nh, block, cfg, is_first=False): super().__init__() self.ln1 = make_norm(d, cfg) self.ln2 = make_norm(d, cfg) self.qkv = nn.Linear(d, 3 * d) self.proj = nn.Linear(d, d) if cfg.ffn == "swiglu": self.mlp = SwiGLU(d, int(round(cfg.ffn_mult * d))) else: self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d)) self.nh = nh self.d = d self.cfg = cfg self.is_first = is_first if cfg.value_residual and not is_first: self.vr_lambda = nn.Parameter(torch.zeros(1)) if cfg.qk_norm: hd = d // nh self.q_norm = RMSNorm(hd, cfg.norm_eps) self.k_norm = RMSNorm(hd, cfg.norm_eps) def forward(self, x, v0=None): B, T, D = x.size() h = self.ln1(x) q, k, v = self.qkv(h).split(self.d, dim=2) hd = D // self.nh q = q.view(B, T, self.nh, hd).transpose(1, 2) k = k.view(B, T, self.nh, hd).transpose(1, 2) v = v.view(B, T, self.nh, hd).transpose(1, 2) if self.cfg.qk_norm: q = self.q_norm(q) k = self.k_norm(k) if self.cfg.pos == "rope": q = apply_rope(q, self.cfg.rope_theta) k = apply_rope(k, self.cfg.rope_theta) if self.cfg.value_residual: if self.is_first: v0 = v else: v = v + self.vr_lambda * v0 y = F.scaled_dot_product_attention(q, k, v, is_causal=True) y = y.transpose(1, 2).contiguous().view(B, T, D) x = x + self.proj(y) x = x + self.mlp(self.ln2(x)) return x, v0 class GPT(nn.Module): def __init__(self, vocab, n_layer, n_embd, n_head, block, cfg): super().__init__() self.cfg = cfg self.tok = nn.Embedding(vocab, n_embd) self.use_rope = cfg.pos == "rope" if not self.use_rope: self.pos = nn.Embedding(block, n_embd) self.blocks = nn.ModuleList([Block(n_embd, n_head, block, cfg, is_first=(i == 0)) for i in range(n_layer)]) self.lnf = make_norm(n_embd, cfg) self.head = nn.Linear(n_embd, vocab, bias=False) self.head.weight = self.tok.weight # tied self.block = block def forward(self, idx): B, T = idx.size() x = self.tok(idx) if not self.use_rope: pos = torch.arange(T, device=idx.device) x = x + self.pos(pos)[None] v0 = None for b in self.blocks: x, v0 = b(x, v0) logits = self.head(self.lnf(x)) cap = getattr(self.cfg, "logit_cap", 0.0) if cap and cap > 0: logits = cap * torch.tanh(logits / cap) return logits def load_slayer139(path, device="cpu"): """Model (eval mode) and tokenizer from a directory with config.json, model.safetensors and tokenizer.json.""" from safetensors.torch import load_file from tokenizers import Tokenizer with open(os.path.join(path, "config.json"), encoding="utf-8") as f: cfg = types.SimpleNamespace(**json.load(f)) model = GPT(cfg.vocab, cfg.n_layer, cfg.n_embd, cfg.n_head, cfg.block, cfg) state = load_file(os.path.join(path, "model.safetensors")) # the output head is tied to the token embedding and stored once (tok.weight); restore the tied key for strict loading if "head.weight" in state: raise ValueError("model.safetensors must not contain head.weight (tied to tok.weight, stored once)") state["head.weight"] = state["tok.weight"] model.load_state_dict(state, strict=True) return model.to(device).eval(), Tokenizer.from_file(os.path.join(path, "tokenizer.json")) @torch.no_grad() def generate(model, tok, prompt, max_new_tokens=64, temperature=0.0, top_k=None, seed=None): """Continue `prompt`. temperature 0 = greedy; otherwise sampling, optionally top-k.""" g = torch.Generator().manual_seed(seed) if seed is not None else None dev = next(model.parameters()).device x = torch.tensor([tok.encode(prompt).ids], device=dev) for _ in range(max_new_tokens): logits = model(x[:, -model.block:])[:, -1].float().cpu() if temperature <= 0: nxt = logits.argmax(-1, keepdim=True) else: logits = logits / temperature if top_k: kth = torch.topk(logits, top_k).values[:, -1:] logits = logits.masked_fill(logits < kth, float("-inf")) nxt = torch.multinomial(torch.softmax(logits, -1), 1, generator=g) x = torch.cat([x, nxt.to(dev)], dim=1) return tok.decode(x[0].tolist())