Luke A Kist commited on
Commit
9d9e608
·
verified ·
1 Parent(s): 160cd60

Upload model.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. model.py +120 -0
model.py ADDED
@@ -0,0 +1,120 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ AAC Micro Brain — 16M parameter conversational flow model.
3
+ Tiny transformer that only knows how humans talk in everyday situations.
4
+ No world knowledge. No encyclopedia. Just conversation patterns.
5
+
6
+ Architecture: ~16M params
7
+ - vocab_size: 8192
8
+ - d_model: 512
9
+ - n_heads: 8
10
+ - n_layers: 6
11
+ - d_ff: 1024
12
+ - max_seq_len: 128
13
+ """
14
+
15
+ import mlx.core as mx
16
+ import mlx.nn as nn
17
+ import math
18
+
19
+
20
+ class MultiHeadAttention(nn.Module):
21
+ def __init__(self, d_model: int, n_heads: int):
22
+ super().__init__()
23
+ self.n_heads = n_heads
24
+ self.d_head = d_model // n_heads
25
+ self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
26
+ self.out = nn.Linear(d_model, d_model, bias=False)
27
+
28
+ def __call__(self, x, mask=None):
29
+ B, T, C = x.shape
30
+ qkv = self.qkv(x)
31
+ q, k, v = mx.split(qkv, 3, axis=-1)
32
+
33
+ q = q.reshape(B, T, self.n_heads, self.d_head).transpose(0, 2, 1, 3)
34
+ k = k.reshape(B, T, self.n_heads, self.d_head).transpose(0, 2, 1, 3)
35
+ v = v.reshape(B, T, self.n_heads, self.d_head).transpose(0, 2, 1, 3)
36
+
37
+ scale = math.sqrt(self.d_head)
38
+ attn = (q @ k.transpose(0, 1, 3, 2)) / scale
39
+
40
+ if mask is not None:
41
+ attn = attn + mask
42
+
43
+ attn = mx.softmax(attn, axis=-1)
44
+ out = (attn @ v).transpose(0, 2, 1, 3).reshape(B, T, C)
45
+ return self.out(out)
46
+
47
+
48
+ class TransformerBlock(nn.Module):
49
+ def __init__(self, d_model: int, n_heads: int, d_ff: int):
50
+ super().__init__()
51
+ self.attn = MultiHeadAttention(d_model, n_heads)
52
+ self.ff = nn.Sequential(
53
+ nn.Linear(d_model, d_ff, bias=False),
54
+ nn.GELU(),
55
+ nn.Linear(d_ff, d_model, bias=False),
56
+ )
57
+ self.ln1 = nn.RMSNorm(d_model)
58
+ self.ln2 = nn.RMSNorm(d_model)
59
+
60
+ def __call__(self, x, mask=None):
61
+ x = x + self.attn(self.ln1(x), mask=mask)
62
+ x = x + self.ff(self.ln2(x))
63
+ return x
64
+
65
+
66
+ class MicroBrain(nn.Module):
67
+ """16M param conversational flow predictor."""
68
+
69
+ def __init__(
70
+ self,
71
+ vocab_size: int = 8192,
72
+ d_model: int = 512,
73
+ n_heads: int = 8,
74
+ n_layers: int = 6,
75
+ d_ff: int = 1024,
76
+ max_seq_len: int = 128,
77
+ ):
78
+ super().__init__()
79
+ self.d_model = d_model
80
+ self.max_seq_len = max_seq_len
81
+
82
+ self.token_emb = nn.Embedding(vocab_size, d_model)
83
+ self.pos_emb = nn.Embedding(max_seq_len, d_model)
84
+
85
+ self.layers = [TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)]
86
+ self.ln_final = nn.RMSNorm(d_model)
87
+ self.output = nn.Linear(d_model, vocab_size, bias=False)
88
+
89
+ def __call__(self, tokens):
90
+ B, T = tokens.shape
91
+ positions = mx.arange(T)
92
+
93
+ x = self.token_emb(tokens) + self.pos_emb(positions)
94
+
95
+ # Causal mask
96
+ mask = nn.MultiHeadAttention.create_additive_causal_mask(T)
97
+
98
+ for layer in self.layers:
99
+ x = layer(x, mask=mask)
100
+
101
+ x = self.ln_final(x)
102
+ logits = self.output(x)
103
+ return logits
104
+
105
+ def count_params(self):
106
+ """Count total parameters."""
107
+ from mlx.utils import tree_flatten
108
+ return sum(v.size for _, v in tree_flatten(self.parameters()))
109
+
110
+
111
+ def create_model(**kwargs):
112
+ model = MicroBrain(**kwargs)
113
+ mx.eval(model.parameters())
114
+ n_params = model.count_params()
115
+ print(f"MicroBrain: {n_params:,} parameters ({n_params / 1e6:.1f}M)")
116
+ return model
117
+
118
+
119
+ if __name__ == "__main__":
120
+ model = create_model()