Slayer149 / modeling_gollem.py
kacperwikiel's picture
Release GoLLeM 149M after 20B continuation tokens with full GLINT evaluation
3f431df verified
Raw History Blame Contribute Delete
5.28 kB
"""GoLLeM inference architecture, extracted unchanged from the pinned r6 trainer.
Source: SlayerLab/gollem-v5-ckpts; Apache-2.0. See README and NOTICE.
"""
import torch
from torch import nn
from torch.nn import functional as F
class RMSNorm(nn.Module):
"""Qwen3-style RMSNorm (fp32-compute dla stabilnosci). 1D weight -> AdamW w split-Muon."""
def __init__(self, d, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(d))
self.eps = eps
def forward(self, x):
return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight
def make_norm(d, cfg):
return RMSNorm(d, cfg.norm_eps) if cfg.norm == "rmsnorm" else nn.LayerNorm(d)
def apply_rope(x, base=100000.0):
"""Parameter-free RoPE na [B,H,T,D] (interleaved-conv, port z qwen_model.py). Train==eval
MUSZA uzywac tej samej konwencji (self-contained eval -> spojne)."""
_, _, T, dim = x.shape
pos = torch.arange(T, device=x.device, dtype=torch.float32)
freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=x.device, dtype=torch.float32) / dim))
ang = torch.outer(pos, freq)
cos, sin = ang.cos().to(x.dtype)[None, None], ang.sin().to(x.dtype)[None, None]
even, odd = x[..., ::2], x[..., 1::2]
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)
class SwiGLU(nn.Module):
"""Qwen3 gated-MLP: down(silu(gate(x))*up(x)). 3x 2D bez-bias -> wszystkie do Muon."""
def __init__(self, d, hidden):
super().__init__()
self.gate = nn.Linear(d, hidden, bias=False)
self.up = nn.Linear(d, hidden, bias=False)
self.down = nn.Linear(hidden, d, bias=False)
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))
class Block(nn.Module):
def __init__(self, d, nh, block, cfg, is_first=False):
super().__init__()
self.ln1 = make_norm(d, cfg)
self.ln2 = make_norm(d, cfg)
self.qkv = nn.Linear(d, 3 * d)
self.proj = nn.Linear(d, d)
if cfg.ffn == "swiglu":
self.mlp = SwiGLU(d, int(round(cfg.ffn_mult * d)))
else:
self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
self.nh = nh
self.d = d
self.cfg = cfg
self.is_first = is_first
if cfg.value_residual and not is_first:
self.vr_lambda = nn.Parameter(torch.zeros(1))
if cfg.qk_norm:
hd = d // nh
self.q_norm = RMSNorm(hd, cfg.norm_eps)
self.k_norm = RMSNorm(hd, cfg.norm_eps)
def forward(self, x, v0=None):
B, T, D = x.size()
h = self.ln1(x)
q, k, v = self.qkv(h).split(self.d, dim=2)
hd = D // self.nh
q = q.view(B, T, self.nh, hd).transpose(1, 2)
k = k.view(B, T, self.nh, hd).transpose(1, 2)
v = v.view(B, T, self.nh, hd).transpose(1, 2)
if self.cfg.qk_norm:
q = self.q_norm(q)
k = self.k_norm(k)
if self.cfg.pos == "rope":
q = apply_rope(q, self.cfg.rope_theta)
k = apply_rope(k, self.cfg.rope_theta)
if self.cfg.value_residual:
if self.is_first:
v0 = v
else:
v = v + self.vr_lambda * v0
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
y = y.transpose(1, 2).contiguous().view(B, T, D)
x = x + self.proj(y)
x = x + self.mlp(self.ln2(x))
return x, v0
class GPT(nn.Module):
def __init__(self, vocab, n_layer, n_embd, n_head, block, cfg):
super().__init__()
self.cfg = cfg
self.tok = nn.Embedding(vocab, n_embd)
self.use_rope = cfg.pos == "rope"
if not self.use_rope:
self.pos = nn.Embedding(block, n_embd)
self.blocks = nn.ModuleList([Block(n_embd, n_head, block, cfg, is_first=(i == 0)) for i in range(n_layer)])
self.lnf = make_norm(n_embd, cfg)
self.head = nn.Linear(n_embd, vocab, bias=False)
self.head.weight = self.tok.weight # tie
self.block = block
self.apply(self._init)
def _init(self, m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, 0.0, 0.02)
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, 0.0, 0.02)
def forward(self, idx, targets=None):
B, T = idx.size()
x = self.tok(idx)
if not self.use_rope:
pos = torch.arange(T, device=idx.device)
x = x + self.pos(pos)[None]
v0 = None
for b in self.blocks:
x, v0 = b(x, v0)
logits = self.head(self.lnf(x))
cap = getattr(self.cfg, "logit_cap", 0.0)
if cap and cap > 0:
logits = cap * torch.tanh(logits / cap)
loss = None
if targets is not None:
flat = logits.view(-1, logits.size(-1))
loss = F.cross_entropy(flat, targets.view(-1))
zc = getattr(self.cfg, "z_loss", 0.0)
if zc and zc > 0:
lse = torch.logsumexp(flat, dim=-1)
loss = loss + zc * (lse * lse).mean()
return logits, loss