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