"""GoLLeM v4 250M — self-contained model + loader + generate (custom arch, NIE HF-transformers). nanoGPT-style: pre-LN, GELU MLP, SDPA causal attention, tied embeddings, learned positions. Repo zawiera SERIE migawek trajektorii treningu: ckpt_.safetensors (2000/6000/8000/...). To sa MIGAWKI TRENINGU, nie warianty do uzycia - jakosc rosnie ze stepem. Uzycie: from modeling_gollem import GollemGPT m = GollemGPT.from_pretrained("./") # najnowsza migawka m = GollemGPT.from_pretrained("./", ckpt="ckpt_6000.safetensors") # konkretny etap print(m.generate("Polska to kraj", max_new_tokens=60)) """ import os, re, json, glob, torch import torch.nn as nn, torch.nn.functional as F class Block(nn.Module): def __init__(s, d, nh): super().__init__(); s.nh = nh s.ln1 = nn.LayerNorm(d); s.ln2 = nn.LayerNorm(d) s.qkv = nn.Linear(d, 3 * d); s.proj = nn.Linear(d, d) s.fc = nn.Linear(d, 4 * d); s.fc2 = nn.Linear(4 * d, d) def forward(s, x): B, T, D = x.shape q, k, v = s.qkv(s.ln1(x)).split(D, 2) q = q.view(B, T, s.nh, D // s.nh).transpose(1, 2) k = k.view(B, T, s.nh, D // s.nh).transpose(1, 2) v = v.view(B, T, s.nh, D // s.nh).transpose(1, 2) y = F.scaled_dot_product_attention(q, k, v, is_causal=True) y = y.transpose(1, 2).contiguous().view(B, T, D) x = x + s.proj(y) x = x + s.fc2(F.gelu(s.fc(s.ln2(x)))) return x class GollemGPT(nn.Module): def __init__(s, vocab, d, nl, nh, block): super().__init__() s.vocab, s.d, s.nl, s.nh, s.block = vocab, d, nl, nh, block s.tok = nn.Embedding(vocab, d); s.pos = nn.Embedding(block, d) s.blocks = nn.ModuleList([Block(d, nh) for _ in range(nl)]) s.lnf = nn.LayerNorm(d); s.head = nn.Linear(d, vocab, bias=False) s.tok.weight = s.head.weight # weight tying s._tok = None def forward(s, idx, tgt=None): B, T = idx.shape x = s.tok(idx) + s.pos(torch.arange(T, device=idx.device))[None] for b in s.blocks: x = b(x) logits = s.head(s.lnf(x)) loss = None if tgt is None else F.cross_entropy(logits.view(-1, logits.size(-1)), tgt.view(-1)) return logits, loss @staticmethod def _latest_ckpt(path): cks = glob.glob(os.path.join(path, "ckpt_*.safetensors")) assert cks, f"brak ckpt_*.safetensors w {path}" return os.path.basename(max(cks, key=lambda p: int(re.search(r"ckpt_(\d+)", p).group(1)))) @classmethod def from_pretrained(cls, path, ckpt=None, device="cpu"): cfg = json.load(open(os.path.join(path, "config.json"))) m = cls(cfg["vocab_size"], cfg["d_model"], cfg["n_layer"], cfg["n_head"], cfg["block_size"]) ckpt = ckpt or cls._latest_ckpt(path) from safetensors.torch import load_file sd = load_file(os.path.join(path, ckpt)) missing, unexpected = m.load_state_dict(sd, strict=False) assert not unexpected and missing in ([], ["head.weight"]), f"load mismatch: {missing}/{unexpected}" m.eval().to(device) tokp = os.path.join(path, "tokenizer.json") if os.path.exists(tokp): from tokenizers import Tokenizer m._tok = Tokenizer.from_file(tokp) m._ckpt = ckpt return m @torch.no_grad() def generate(s, prompt, max_new_tokens=60, temperature=0.8, top_k=40, seed=None): assert s._tok is not None, "brak tokenizer.json przy modelu" if seed is not None: torch.manual_seed(seed) dev = next(s.parameters()).device ids = s._tok.encode(prompt).ids x = torch.tensor([ids], dtype=torch.long, device=dev) for _ in range(max_new_tokens): logits, _ = s(x[:, -s.block:]) logits = logits[0, -1] / max(1e-6, temperature) if top_k: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[-1]] = -float("inf") probs = torch.softmax(logits, -1) nxt = torch.multinomial(probs, 1) x = torch.cat([x, nxt[None]], 1) return s._tok.decode(x[0].tolist())