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