eshanized commited on
Commit
4fd5eb6
·
verified ·
1 Parent(s): 2012230

add v5 experiment runtime

Browse files
experiments/M31-Python-Agent-220M-v5/runtime/modeling_m31.py ADDED
@@ -0,0 +1,126 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from typing import Optional, Tuple
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+
7
+ VOCAB_SIZE = 32_768
8
+ MAX_CONTEXT = 2_048
9
+ TRAIN_SEQ_LEN = 512
10
+ HIDDEN = 896
11
+ LAYERS = 20
12
+ HEADS = 14
13
+ KV_HEADS = 7
14
+ INTERMEDIATE = 2_816
15
+ ROPE_THETA = 10_000.0
16
+ RMS_EPS = 1e-6
17
+
18
+
19
+ def masked_cross_entropy(logits, labels):
20
+ valid = labels.ne(-100)
21
+ if int(valid.sum().item()) <= 0:
22
+ raise RuntimeError('masked_cross_entropy received zero supervised targets')
23
+ return F.cross_entropy(logits.float().reshape(-1, VOCAB_SIZE), labels.reshape(-1), ignore_index=-100)
24
+
25
+
26
+ class RMSNorm(nn.Module):
27
+ def __init__(self, d=HIDDEN, eps=RMS_EPS):
28
+ super().__init__()
29
+ self.weight = nn.Parameter(torch.ones(d))
30
+ self.eps = eps
31
+
32
+ def forward(self, x):
33
+ return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight
34
+
35
+
36
+ def rotate_half(x):
37
+ half = x.shape[-1] // 2
38
+ return torch.cat((-x[..., half:], x[..., :half]), dim=-1)
39
+
40
+
41
+ class Rotary(nn.Module):
42
+ def __init__(self, head_dim, max_seq=MAX_CONTEXT, theta=ROPE_THETA):
43
+ super().__init__()
44
+ inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
45
+ t = torch.arange(max_seq, dtype=torch.float32)
46
+ freqs = torch.outer(t, inv_freq)
47
+ emb = torch.cat((freqs, freqs), dim=-1)
48
+ self.register_buffer('cos', emb.cos()[None, :, :], persistent=False)
49
+ self.register_buffer('sin', emb.sin()[None, :, :], persistent=False)
50
+
51
+ def forward(self, q, k):
52
+ n = q.shape[-2]
53
+ cos = self.cos[:, :n].to(q.device, q.dtype)
54
+ sin = self.sin[:, :n].to(q.device, q.dtype)
55
+ return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin
56
+
57
+
58
+ class M31Attention(nn.Module):
59
+ def __init__(self):
60
+ super().__init__()
61
+ assert HIDDEN % HEADS == 0
62
+ assert HEADS % KV_HEADS == 0
63
+ self.head_dim = HIDDEN // HEADS
64
+ self.q = nn.Linear(HIDDEN, HIDDEN, bias=False)
65
+ self.k = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False)
66
+ self.v = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False)
67
+ self.o = nn.Linear(HIDDEN, HIDDEN, bias=False)
68
+ self.rope = Rotary(self.head_dim, MAX_CONTEXT, ROPE_THETA)
69
+
70
+ def forward(self, x):
71
+ b, t, _ = x.shape
72
+ q = self.q(x).view(b, t, HEADS, self.head_dim).transpose(1, 2)
73
+ k = self.k(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2)
74
+ v = self.v(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2)
75
+ repeat = HEADS // KV_HEADS
76
+ k = k.repeat_interleave(repeat, dim=1)
77
+ v = v.repeat_interleave(repeat, dim=1)
78
+ q, k = self.rope(q, k)
79
+ y = F.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
80
+ return self.o(y.transpose(1, 2).contiguous().view(b, t, HIDDEN))
81
+
82
+
83
+ class M31Block(nn.Module):
84
+ def __init__(self):
85
+ super().__init__()
86
+ self.n1 = RMSNorm(HIDDEN, RMS_EPS)
87
+ self.attn = M31Attention()
88
+ self.n2 = RMSNorm(HIDDEN, RMS_EPS)
89
+ self.gate = nn.Linear(HIDDEN, INTERMEDIATE, bias=False)
90
+ self.up = nn.Linear(HIDDEN, INTERMEDIATE, bias=False)
91
+ self.down = nn.Linear(INTERMEDIATE, HIDDEN, bias=False)
92
+
93
+ def forward(self, x):
94
+ x = x + self.attn(self.n1(x))
95
+ h = self.n2(x)
96
+ h = F.silu(self.gate(h)) * self.up(h)
97
+ return x + self.down(h)
98
+
99
+
100
+ class M31Model(nn.Module):
101
+ def __init__(self):
102
+ super().__init__()
103
+ self.embed = nn.Embedding(VOCAB_SIZE, HIDDEN)
104
+ self.blocks = nn.ModuleList([M31Block() for _ in range(LAYERS)])
105
+ self.norm = RMSNorm(HIDDEN, RMS_EPS)
106
+ self.apply(self._init_weights)
107
+ nn.init.normal_(self.embed.weight, mean=0.0, std=0.02)
108
+ self.num_parameters = sum(p.numel() for p in self.parameters())
109
+ if self.num_parameters >= 250_000_000:
110
+ raise RuntimeError('Hard parameter ceiling violated.')
111
+
112
+ @staticmethod
113
+ def _init_weights(m):
114
+ if isinstance(m, nn.Linear):
115
+ nn.init.normal_(m.weight, mean=0.0, std=0.02 / math.sqrt(2 * LAYERS))
116
+ elif isinstance(m, nn.Embedding):
117
+ pass
118
+
119
+ def forward(self, input_ids, labels: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
120
+ x = self.embed(input_ids)
121
+ for block in self.blocks:
122
+ x = block(x)
123
+ x = self.norm(x)
124
+ logits = F.linear(x, self.embed.weight)
125
+ loss = masked_cross_entropy(logits, labels) if labels is not None else None
126
+ return logits, loss