Download modeling_gollem.py from SlayerLab/Slayer149: direct link, hf CLI and curl.
- Browser
- Download file 5.28 kB
-
https://huggingface.co/SlayerLab/Slayer149/resolve/main/modeling_gollem.py
- Command line
-
hf download hf://SlayerLab/Slayer149/modeling_gollem.py
-
curl -L -o modeling_gollem.py https://huggingface.co/SlayerLab/Slayer149/resolve/main/modeling_gollem.py
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 | |