File size: 5,283 Bytes
3f431df | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | """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
|