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