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