#!/usr/bin/env python3 """Load SealGlazer v11 from model.safetensors and generate text. Usage: python load_sealglazer.py --prompt "The harbor seal is" --seed 999 """ import argparse, math, os, torch, torch.nn.functional as F from safetensors.torch import load_file from tokenizers import Tokenizer HERE = os.path.dirname(os.path.abspath(__file__)) class RMSNorm(torch.nn.Module): def __init__(s, c, eps=1e-6): super().__init__(); s.eps = eps; s.weight = torch.nn.Parameter(torch.ones(c)) def forward(s, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + s.eps) * s.weight def rope(hd, ms, base=10000.0): f = 1.0 / (base ** (torch.arange(0, hd, 2).float() / hd)) t = torch.arange(ms); a = torch.outer(t, f) return torch.cos(a), torch.sin(a) class Attn(torch.nn.Module): def __init__(s, C, h): super().__init__(); HD = C // h s.HD = HD s.wq = torch.nn.Linear(C, C, bias=False); s.wk = torch.nn.Linear(C, C, bias=False) s.wv = torch.nn.Linear(C, C, bias=False); s.wo = torch.nn.Linear(C, C, bias=False) def forward(s, x, cos, sin): B, T, C = x.shape; h = C // s.HD q = s.wq(x).view(B, T, h, s.HD).transpose(1, 2) k = s.wk(x).view(B, T, h, s.HD).transpose(1, 2) v = s.wv(x).view(B, T, h, s.HD).transpose(1, 2) def rot(t): t1 = t[..., :s.HD//2]; t2 = t[..., s.HD//2:]; c = cos[:T].unsqueeze(0); si = sin[:T].unsqueeze(0) return torch.cat((c*t1 - si*t2, c*t2 + si*t1), dim=-1) q, k = rot(q), rot(k) att = F.softmax((q @ k.transpose(-2, -1)) / math.sqrt(s.HD), dim=-1) return s.wo((att @ v).transpose(1, 2).contiguous().view(B, T, C)) class MLP(torch.nn.Module): def __init__(s, C, f): super().__init__() s.w1 = torch.nn.Linear(C, f, bias=False); s.w2 = torch.nn.Linear(C, f, bias=False) s.w3 = torch.nn.Linear(f, C, bias=False) def forward(s, x): return s.w3(F.silu(s.w1(x)) * s.w2(x)) class Block(torch.nn.Module): def __init__(s, C, h, f): super().__init__(); s.ln1 = RMSNorm(C); s.attn = Attn(C, h); s.ln2 = RMSNorm(C); s.mlp = MLP(C, f) def forward(s, x, cos, sin): x = x + s.attn(s.ln1(x), cos, sin); x = x + s.mlp(s.ln2(x)); return x class Model(torch.nn.Module): def __init__(s, V, C, L, h, BLOCK, f=384): super().__init__() s.tok = torch.nn.Embedding(V, C) s.blocks = torch.nn.ModuleList([Block(C, h, f) for _ in range(L)]) s.ln_f = RMSNorm(C); s.cos, s.sin = rope(C // h, BLOCK) def forward(s, idx): h = s.tok(idx); cos = s.cos.to(idx.device); sin = s.sin.to(idx.device) for b in s.blocks: h = b(h, cos, sin) return F.linear(s.ln_f(h), s.tok.weight) # tied lm_head def load(dev="cpu"): sd = load_file(os.path.join(HERE, "model.safetensors")) V = sd["tok.weight"].shape[0]; C = sd["tok.weight"].shape[1] L = len([k for k in sd if k.startswith("blocks.")]) // 4 m = Model(V, C, L, h=4, BLOCK=256).to(dev) m.load_state_dict(sd); m.eval() return m def generate(m, tok, prompt, max_new=120, temp=0.7, top_p=0.9, seed=0, dev="cpu"): g = torch.Generator(device=dev).manual_seed(seed) pids = tok.encode(prompt, add_special_tokens=False).ids x = torch.tensor([[pids]], device=dev) with torch.no_grad(): for _ in range(max_new): T = x.shape[1] if T > 256: x = x[:, -256:]; T = 256 logits = m(x)[:, -1, :] / temp v = torch.log_softmax(logits, dim=-1) sv, si = torch.sort(v, descending=True) cp = torch.cumsum(torch.softmax(sv, dim=-1), dim=-1) rm = cp > top_p; rm[0] = False v[si[rm]] = float("-inf") nxt = torch.multinomial(torch.softmax(v, dim=-1), 1, generator=g) x = torch.cat([x, nxt], dim=1) return tok.decode(x[0].tolist()) if __name__ == "__main__": ap = argparse.ArgumentParser() ap.add_argument("--prompt", default="The harbor seal is") ap.add_argument("--seed", type=int, default=999) ap.add_argument("--max-new", type=int, default=120) ap.add_argument("--temp", type=float, default=0.7) args = ap.parse_args() dev = "cuda" if torch.cuda.is_available() else "cpu" m = load(dev) tok = Tokenizer.from_file(os.path.join(HERE, "tokenizer.json")) print(generate(m, tok, args.prompt, args.max_new, args.temp, 0.9, args.seed, dev))