Download experiments/M31-Python-Agent-220M-v5/runtime/modeling_m31.py from eshanized/M31Tesla: direct link, hf CLI and curl.
- Browser
- Download file 4.59 kB
-
https://huggingface.co/eshanized/M31Tesla/resolve/main/experiments/M31-Python-Agent-220M-v5/runtime/modeling_m31.py
- Command line
-
hf download hf://eshanized/M31Tesla/experiments/M31-Python-Agent-220M-v5/runtime/modeling_m31.py
-
curl -L -o modeling_m31.py https://huggingface.co/eshanized/M31Tesla/resolve/main/experiments/M31-Python-Agent-220M-v5/runtime/modeling_m31.py
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.') | |
| 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 | |