sealglazer-1.9m / load_sealglazer.py
Compactbot's picture
Add SealGlazer v11 card, config, and loader (public release per requester) (#1)
71cacc6
Raw
History Blame Contribute Delete
4.46 kB
#!/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))