Slayer149-balanced / balanced_model.py
kacperwikiel's picture
Update tokenizer, configuration and training reports
d6a86af verified
Raw History Blame Contribute Delete
4.71 kB
"""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)