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