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