File size: 4,714 Bytes
d6a86af
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Native Qwen3-style GQA model with explicit value residuals and tied embeddings.

State names match the reviewed JugnuVR checkpoint. Explicit v0 avoids a shared
mutable context dictionary in the training forward. No KV cache is implemented.
"""
import torch
from torch import nn
from torch.nn import functional as F

class RMSNorm(nn.Module):
    def __init__(self, width, eps=1e-6):
        super().__init__(); self.weight=nn.Parameter(torch.ones(width)); self.eps=eps
    def forward(self,x):
        dtype=x.dtype; x=x.float(); x=x*torch.rsqrt(x.square().mean(-1,keepdim=True)+self.eps)
        return self.weight*x.to(dtype)

class ValueProjection(nn.Linear):
    def __init__(self,width,out,first):
        super().__init__(width,out,bias=False)
        if not first:self.vr_lambda=nn.Parameter(torch.zeros(1))

class Attention(nn.Module):
    def __init__(self,c,index):
        super().__init__();h=c['hidden_size'];self.heads=c['num_attention_heads'];self.kv=c['num_key_value_heads'];self.d=c.get('head_dim',h//self.heads);self.first=index==0
        self.q_proj=nn.Linear(h,self.heads*self.d,bias=False);self.k_proj=nn.Linear(h,self.kv*self.d,bias=False)
        self.v_proj=ValueProjection(h,self.kv*self.d,self.first);self.o_proj=nn.Linear(self.heads*self.d,h,bias=False)
        self.q_norm=RMSNorm(self.d,c.get('rms_norm_eps',1e-6));self.k_norm=RMSNorm(self.d,c.get('rms_norm_eps',1e-6))
    def forward(self,x,cos,sin,v0):
        b,t,_=x.shape
        q=self.q_norm(self.q_proj(x).view(b,t,self.heads,self.d)).transpose(1,2)
        k=self.k_norm(self.k_proj(x).view(b,t,self.kv,self.d)).transpose(1,2)
        v=self.v_proj(x)
        if self.first:v0=v
        else:v=v+self.v_proj.vr_lambda*v0
        v=v.view(b,t,self.kv,self.d).transpose(1,2)
        def rotate(z):return torch.cat((-z[...,self.d//2:],z[...,:self.d//2]),dim=-1)
        q=q*cos+rotate(q)*sin;k=k*cos+rotate(k)*sin
        k=k.repeat_interleave(self.heads//self.kv,dim=1);v=v.repeat_interleave(self.heads//self.kv,dim=1)
        out=F.scaled_dot_product_attention(q,k,v,is_causal=True)
        return self.o_proj(out.transpose(1,2).contiguous().view(b,t,-1)),v0

class MLP(nn.Module):
    def __init__(self,c):
        super().__init__();h=c['hidden_size'];f=c['intermediate_size']
        self.gate_proj=nn.Linear(h,f,bias=False);self.up_proj=nn.Linear(h,f,bias=False);self.down_proj=nn.Linear(f,h,bias=False)
    def forward(self,x):return self.down_proj(F.silu(self.gate_proj(x))*self.up_proj(x))

class Block(nn.Module):
    def __init__(self,c,i):
        super().__init__();self.self_attn=Attention(c,i);self.mlp=MLP(c)
        self.input_layernorm=RMSNorm(c['hidden_size'],c.get('rms_norm_eps',1e-6));self.post_attention_layernorm=RMSNorm(c['hidden_size'],c.get('rms_norm_eps',1e-6))
    def forward(self,x,cos,sin,v0):
        a,v0=self.self_attn(self.input_layernorm(x),cos,sin,v0);x=x+a
        return x+self.mlp(self.post_attention_layernorm(x)),v0

class Backbone(nn.Module):
    def __init__(self,c):
        super().__init__();self.embed_tokens=nn.Embedding(c['vocab_size'],c['hidden_size']);self.layers=nn.ModuleList([Block(c,i) for i in range(c['num_hidden_layers'])]);self.norm=RMSNorm(c['hidden_size'],c.get('rms_norm_eps',1e-6))

class BalancedLM(nn.Module):
    def __init__(self,c):
        super().__init__();self.config=dict(c)
        assert c['num_attention_heads']%c['num_key_value_heads']==0
        assert c['hidden_size']==c['num_attention_heads']*c.get('head_dim',64)
        self.model=Backbone(c);self.lm_head=nn.Linear(c['hidden_size'],c['vocab_size'],bias=False);self.lm_head.weight=self.model.embed_tokens.weight
        d=c.get('head_dim',64);theta=c.get('rope_parameters',{}).get('rope_theta',10000.)
        inv=1./(theta**(torch.arange(0,d,2,dtype=torch.float32)/d))
        phases=torch.outer(torch.arange(c['max_position_embeddings'],dtype=torch.float32),inv);phases=torch.cat((phases,phases),dim=-1)
        self.register_buffer('cos',phases.cos()[None,None],persistent=False);self.register_buffer('sin',phases.sin()[None,None],persistent=False)
    def forward(self,ids,targets=None):
        x=self.model.embed_tokens(ids);v0=None;t=ids.shape[1]
        for layer in self.model.layers:x,v0=layer(x,self.cos[:,:,:t],self.sin[:,:,:t],v0)
        logits=self.lm_head(self.model.norm(x))
        loss=None if targets is None else F.cross_entropy(logits.float().flatten(0,1),targets.flatten())
        return logits,loss
    @torch.no_grad()
    def initialize(self,seed):
        torch.manual_seed(seed)
        for name,p in self.named_parameters():
            if p.ndim==2:p.normal_(mean=0,std=0.02)
            elif name.endswith('vr_lambda'):p.zero_()
            else:p.fill_(1)