M31Genesis / modeling_m31.py
eshanized's picture
update M31Genesis root runtime
b1fd620 verified
Raw History Blame Contribute Delete
3.62 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from tokenizers import Tokenizer
from safetensors.torch import load_model
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
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 rotate_half(x):
a, b = x.chunk(2, dim=-1)
return torch.cat((-b, a), dim=-1)
class M31Attention(nn.Module):
def __init__(self, c):
super().__init__()
h = c['hidden_size']; nh = c['num_attention_heads']; nk = c['num_key_value_heads']; d = h // nh
self.h, self.nh, self.nk, self.d = h, nh, nk, d
self.q = nn.Linear(h, nh*d, bias=False)
self.k = nn.Linear(h, nk*d, bias=False)
self.v = nn.Linear(h, nk*d, bias=False)
self.o = nn.Linear(h, h, bias=False)
inv = 1.0 / (c['rope_theta'] ** (torch.arange(0, d, 2).float() / d))
pos = torch.arange(c['max_position_embeddings'], dtype=torch.float)
f = torch.outer(pos, inv); e = torch.cat([f, f], dim=-1)
self.register_buffer('cos', e.cos()[None, None], persistent=False)
self.register_buffer('sin', e.sin()[None, None], persistent=False)
def forward(self, x):
b, s, _ = x.shape
q = self.q(x).view(b, s, self.nh, self.d).transpose(1, 2)
k = self.k(x).view(b, s, self.nk, self.d).transpose(1, 2)
v = self.v(x).view(b, s, self.nk, self.d).transpose(1, 2)
co = self.cos[:, :, :s, :].to(x.device, x.dtype); si = self.sin[:, :, :s, :].to(x.device, x.dtype)
q = q * co + rotate_half(q) * si; k = k * co + rotate_half(k) * si
try:
y = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)
except Exception:
rep = self.nh // self.nk
y = F.scaled_dot_product_attention(q, k.repeat_interleave(rep, 1), v.repeat_interleave(rep, 1), is_causal=True)
return self.o(y.transpose(1, 2).contiguous().view(b, s, self.h))
class M31MLP(nn.Module):
def __init__(self, c):
super().__init__(); h=c['hidden_size']; i=c['intermediate_size']
self.gate=nn.Linear(h,i,bias=False); self.up=nn.Linear(h,i,bias=False); self.down=nn.Linear(i,h,bias=False)
def forward(self,x): return self.down(F.silu(self.gate(x))*self.up(x))
class M31Block(nn.Module):
def __init__(self,c):
super().__init__(); h=c['hidden_size']
self.n1=RMSNorm(h,c['rms_norm_eps']); self.attn=M31Attention(c); self.n2=RMSNorm(h,c['rms_norm_eps']); self.mlp=M31MLP(c)
def forward(self,x):
x = x + self.attn(self.n1(x))
x = x + self.mlp(self.n2(x))
return x
class M31ForCausalLM(nn.Module):
def __init__(self,c):
super().__init__(); h=c['hidden_size']
self.embed=nn.Embedding(c['vocab_size'],h)
self.blocks=nn.ModuleList([M31Block(c) for _ in range(c['num_hidden_layers'])])
self.norm=RMSNorm(h,c['rms_norm_eps'])
self.lm_head=nn.Linear(h,c['vocab_size'],bias=False)
self.lm_head.weight=self.embed.weight
def forward(self,input_ids):
x=self.embed(input_ids)
for b in self.blocks: x=b(x)
return self.lm_head(self.norm(x))
def load_model_bundle(path='.'):
import json
from pathlib import Path
p=Path(path); c=json.loads((p/'config.json').read_text()); tok=Tokenizer.from_file(str(p/'tokenizer.json'))
m=M31ForCausalLM(c); load_model(m,str(p/'model.safetensors'),strict=True); return m,tok