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