M31Tesla / modeling_m31.py
eshanized's picture
add M31 v5 standalone inference runtime
422afe4 verified
Raw History Blame Contribute Delete
4.59 kB
import math
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
VOCAB_SIZE = 32_768
MAX_CONTEXT = 2_048
TRAIN_SEQ_LEN = 512
HIDDEN = 896
LAYERS = 20
HEADS = 14
KV_HEADS = 7
INTERMEDIATE = 2_816
ROPE_THETA = 10_000.0
RMS_EPS = 1e-6
def masked_cross_entropy(logits, labels):
valid = labels.ne(-100)
if int(valid.sum().item()) <= 0:
raise RuntimeError('masked_cross_entropy received zero supervised targets')
return F.cross_entropy(logits.float().reshape(-1, VOCAB_SIZE), labels.reshape(-1), ignore_index=-100)
class RMSNorm(nn.Module):
def __init__(self, d=HIDDEN, eps=RMS_EPS):
super().__init__()
self.weight = nn.Parameter(torch.ones(d))
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):
half = x.shape[-1] // 2
return torch.cat((-x[..., half:], x[..., :half]), dim=-1)
class Rotary(nn.Module):
def __init__(self, head_dim, max_seq=MAX_CONTEXT, theta=ROPE_THETA):
super().__init__()
inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
t = torch.arange(max_seq, dtype=torch.float32)
freqs = torch.outer(t, inv_freq)
emb = torch.cat((freqs, freqs), dim=-1)
self.register_buffer('cos', emb.cos()[None, :, :], persistent=False)
self.register_buffer('sin', emb.sin()[None, :, :], persistent=False)
def forward(self, q, k):
n = q.shape[-2]
cos = self.cos[:, :n].to(q.device, q.dtype)
sin = self.sin[:, :n].to(q.device, q.dtype)
return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin
class M31Attention(nn.Module):
def __init__(self):
super().__init__()
assert HIDDEN % HEADS == 0
assert HEADS % KV_HEADS == 0
self.head_dim = HIDDEN // HEADS
self.q = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.k = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False)
self.v = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False)
self.o = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.rope = Rotary(self.head_dim, MAX_CONTEXT, ROPE_THETA)
def forward(self, x):
b, t, _ = x.shape
q = self.q(x).view(b, t, HEADS, self.head_dim).transpose(1, 2)
k = self.k(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2)
v = self.v(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2)
repeat = HEADS // KV_HEADS
k = k.repeat_interleave(repeat, dim=1)
v = v.repeat_interleave(repeat, dim=1)
q, k = self.rope(q, k)
y = F.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
return self.o(y.transpose(1, 2).contiguous().view(b, t, HIDDEN))
class M31Block(nn.Module):
def __init__(self):
super().__init__()
self.n1 = RMSNorm(HIDDEN, RMS_EPS)
self.attn = M31Attention()
self.n2 = RMSNorm(HIDDEN, RMS_EPS)
self.gate = nn.Linear(HIDDEN, INTERMEDIATE, bias=False)
self.up = nn.Linear(HIDDEN, INTERMEDIATE, bias=False)
self.down = nn.Linear(INTERMEDIATE, HIDDEN, bias=False)
def forward(self, x):
x = x + self.attn(self.n1(x))
h = self.n2(x)
h = F.silu(self.gate(h)) * self.up(h)
return x + self.down(h)
class M31Model(nn.Module):
def __init__(self):
super().__init__()
self.embed = nn.Embedding(VOCAB_SIZE, HIDDEN)
self.blocks = nn.ModuleList([M31Block() for _ in range(LAYERS)])
self.norm = RMSNorm(HIDDEN, RMS_EPS)
self.apply(self._init_weights)
nn.init.normal_(self.embed.weight, mean=0.0, std=0.02)
self.num_parameters = sum(p.numel() for p in self.parameters())
if self.num_parameters >= 250_000_000:
raise RuntimeError('Hard parameter ceiling violated.')
@staticmethod
def _init_weights(m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, mean=0.0, std=0.02 / math.sqrt(2 * LAYERS))
elif isinstance(m, nn.Embedding):
pass
def forward(self, input_ids, labels: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
x = self.embed(input_ids)
for block in self.blocks:
x = block(x)
x = self.norm(x)
logits = F.linear(x, self.embed.weight)
loss = masked_cross_entropy(logits, labels) if labels is not None else None
return logits, loss