catgirl-1m / model.py
Compactbot's picture
Add self-contained model.py (class, loader, generate CLI)
93699e9 verified
Raw History Blame Contribute Delete
6.36 kB
"""catgirl-1m — a tiny tsundere-catgirl chat model (Glint-2 architecture).
Self-contained: this file holds the model class, a loader, and a generate CLI.
Requires only: torch, tokenizers.
python model.py "User: hi there
Assistant:"
python model.py --max-new-tokens 80 --temperature 0.7 "User: do you like cats?
Assistant:"
The model trains at 8 loops of its shared block. Running more loops produces
gibberish, so loops defaults to 8; change it only if you know why.
"""
import argparse
import json
import math
import os
import torch
from tokenizers import Tokenizer
from torch import nn
from torch.nn import functional
LOOPS = 8
ATTENTION_WINDOW = 256
class SwiGlu(nn.Module):
def __init__(self, dim, hidden):
super().__init__()
self.gate_up = nn.Linear(dim, 2 * hidden, bias=False)
self.down = nn.Linear(hidden, dim, bias=False)
def forward(self, x):
gate, up = self.gate_up(x).chunk(2, dim=-1)
return self.down(functional.silu(gate) * up)
def apply_rope(x, cos, sin):
x_even, x_odd = x[..., 0::2], x[..., 1::2]
return torch.stack(
(x_even * cos - x_odd * sin, x_even * sin + x_odd * cos), dim=-1
).flatten(-2)
class Attention(nn.Module):
def __init__(self, dim, n_heads):
super().__init__()
self.n_heads = n_heads
self.head_dim = dim // n_heads
self.qkv = nn.Linear(dim, 3 * dim, bias=False)
self.out = nn.Linear(dim, dim, bias=False)
def forward(self, x, cos, sin, qkv_delta):
batch, seq_len, dim = x.shape
qkv = self.qkv(x) + qkv_delta
q, k, v = qkv.split(dim, dim=-1)
q = q.view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
k = k.view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
v = v.view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
pos = torch.arange(seq_len, device=x.device)
mask = (pos[None, :] <= pos[:, None]) & (pos[:, None] - pos[None, :] < ATTENTION_WINDOW)
attended = functional.scaled_dot_product_attention(
q, k, v, attn_mask=mask, scale=1.0 / math.sqrt(self.head_dim)
)
return self.out(attended.transpose(1, 2).reshape(batch, seq_len, dim))
class Block(nn.Module):
def __init__(self, dim, n_heads, ffn_hidden):
super().__init__()
self.attn_norm = nn.RMSNorm(dim)
self.attn = Attention(dim, n_heads)
self.ffn_norm = nn.RMSNorm(dim)
self.ffn = SwiGlu(dim, ffn_hidden)
def forward(self, x, cos, sin, qkv_delta):
x = x + self.attn(self.attn_norm(x), cos, sin, qkv_delta)
return x + self.ffn(self.ffn_norm(x))
class LoopLora(nn.Module):
def __init__(self, dim, rank, max_loops):
super().__init__()
self.down = nn.ModuleList(nn.Linear(dim, rank, bias=False) for _ in range(max_loops))
self.up = nn.ModuleList(nn.Linear(rank, 3 * dim, bias=False) for _ in range(max_loops))
def forward(self, x, loop_index):
i = min(loop_index, len(self.down) - 1)
return self.up[i](self.down[i](x))
class Indexer(nn.Module):
def __init__(self):
super().__init__()
self.gate = nn.Parameter(torch.tensor([0.1]))
class Glint2(nn.Module):
def __init__(self, cfg, max_loops):
super().__init__()
self.cfg = cfg
self.max_loops = max_loops
self.embed = nn.Embedding(cfg["vocab_size"], cfg["dim"])
self.indexer = Indexer()
self.shared = Block(cfg["dim"], cfg["n_heads"], cfg["ffn_hidden"])
self.loop_lora = LoopLora(cfg["dim"], cfg["lora_rank"], max_loops)
self.loop_embed = nn.Embedding(max_loops, cfg["dim"])
self.final_norm = nn.RMSNorm(cfg["dim"])
def rope(self, seq_len):
head_dim = self.cfg["dim"] // self.cfg["n_heads"]
positions = torch.arange(seq_len, dtype=torch.float32)
inv_freq = 1.0 / (
self.cfg["rope_base"]
** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
)
angles = torch.outer(positions, inv_freq)
return torch.cos(angles), torch.sin(angles)
def forward(self, tokens, loops=LOOPS):
cos, sin = self.rope(tokens.shape[1])
x = self.embed(tokens)
for loop_index in range(loops):
clamped = min(loop_index, self.max_loops - 1)
gated = x + self.loop_embed.weight[clamped]
delta = self.loop_lora(gated, loop_index)
x = self.shared(gated, cos, sin, delta)
x = self.final_norm(x)
return functional.linear(x, self.embed.weight)
def load_model(model_dir):
with open(os.path.join(model_dir, "config.json")) as f:
cfg = json.load(f)
from safetensors.torch import load_file
state = load_file(os.path.join(model_dir, "catgirl-1m.safetensors"))
model = Glint2(cfg, max_loops=cfg["max_loops"])
model.load_state_dict(state, strict=False)
model.eval()
tok = Tokenizer.from_file(os.path.join(model_dir, "tokenizer.json"))
return model, tok
def sample(logits, ids, temperature, top_k, rep, win):
if rep != 1.0:
for t in set(ids[-win:]):
logits[t] = logits[t] / rep if logits[t] > 0 else logits[t] * rep
if temperature <= 0:
return int(logits.argmax())
values, idx = logits.topk(top_k)
probs = functional.softmax(values / temperature, dim=-1)
return int(idx[torch.multinomial(probs, 1)])
if __name__ == "__main__":
ap = argparse.ArgumentParser()
ap.add_argument("prompt")
ap.add_argument("--model-dir", default=os.path.dirname(os.path.abspath(__file__)))
ap.add_argument("--max-new-tokens", type=int, default=64)
ap.add_argument("--temperature", type=float, default=0.7)
ap.add_argument("--top-k", type=int, default=5)
ap.add_argument("--repetition-penalty", type=float, default=1.05)
args = ap.parse_args()
model, tok = load_model(args.model_dir)
ids = tok.encode(args.prompt).ids
with torch.no_grad():
for _ in range(args.max_new_tokens):
logits = model(torch.tensor([ids]), loops=LOOPS)[0, -1]
ids.append(sample(logits, ids, args.temperature, args.top_k,
args.repetition_penalty, 128))
print(tok.decode(ids))