Compactbot commited on
Commit
d73d9e7
·
verified ·
1 Parent(s): 2cc3bfa

Add model code

Browse files
Files changed (1) hide show
  1. model.py +137 -0
model.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """TinyChat-5M: 5.1M param transformer. d=192, 6L, 3H, head_dim=64, SwiGLU 4x, RoPE, RMSNorm, tied embeddings."""
2
+ import torch
3
+ import torch.nn as nn
4
+ import torch.nn.functional as F
5
+
6
+
7
+ class RMSNorm(nn.Module):
8
+ def __init__(self, d, eps=1e-6):
9
+ super().__init__()
10
+ self.weight = nn.Parameter(torch.ones(d))
11
+ self.eps = eps
12
+
13
+ def forward(self, x):
14
+ rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
15
+ return x * rms * self.weight
16
+
17
+
18
+ def rotate_half(x):
19
+ x1, x2 = x.chunk(2, dim=-1)
20
+ return torch.cat([-x2, x1], dim=-1)
21
+
22
+
23
+ class Attention(nn.Module):
24
+ def __init__(self, d, n_heads, head_dim):
25
+ super().__init__()
26
+ self.n_heads = n_heads
27
+ self.head_dim = head_dim
28
+ self.wq = nn.Linear(d, n_heads * head_dim, bias=False)
29
+ self.wk = nn.Linear(d, n_heads * head_dim, bias=False)
30
+ self.wv = nn.Linear(d, n_heads * head_dim, bias=False)
31
+ self.wo = nn.Linear(n_heads * head_dim, d, bias=False)
32
+
33
+ def forward(self, x, cos, sin):
34
+ B, S, _ = x.shape
35
+ q = self.wq(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
36
+ k = self.wk(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
37
+ v = self.wv(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)
38
+ cos_f = cos[:S].unsqueeze(0).unsqueeze(0)
39
+ sin_f = sin[:S].unsqueeze(0).unsqueeze(0)
40
+ q = (q * cos_f) + (rotate_half(q) * sin_f)
41
+ k = (k * cos_f) + (rotate_half(k) * sin_f)
42
+ scale = self.head_dim ** -0.5
43
+ scores = (q @ k.transpose(-2, -1)) * scale
44
+ mask = torch.triu(torch.ones(S, S, device=x.device, dtype=torch.bool), diagonal=1)
45
+ scores = scores.masked_fill(mask, float('-inf'))
46
+ attn = F.softmax(scores, dim=-1)
47
+ out = (attn @ v).transpose(1, 2).reshape(B, S, -1)
48
+ return self.wo(out)
49
+
50
+
51
+ class FFN(nn.Module):
52
+ def __init__(self, d, mult):
53
+ super().__init__()
54
+ hidden = d * mult
55
+ self.w1 = nn.Linear(d, hidden, bias=False)
56
+ self.w2 = nn.Linear(hidden, d, bias=False)
57
+ self.w3 = nn.Linear(d, hidden, bias=False)
58
+
59
+ def forward(self, x):
60
+ return self.w2(F.silu(self.w1(x)) * self.w3(x))
61
+
62
+
63
+ class Block(nn.Module):
64
+ def __init__(self, d, n_heads, head_dim, ffn_mult):
65
+ super().__init__()
66
+ self.norm1 = RMSNorm(d)
67
+ self.attn = Attention(d, n_heads, head_dim)
68
+ self.norm2 = RMSNorm(d)
69
+ self.ffn = FFN(d, ffn_mult)
70
+
71
+ def forward(self, x, cos, sin):
72
+ x = x + self.attn(self.norm1(x), cos, sin)
73
+ x = x + self.ffn(self.norm2(x))
74
+ return x
75
+
76
+
77
+ class TinyLM(nn.Module):
78
+ def __init__(self, vocab=4096, d=192, n_layers=6, n_heads=3, head_dim=64, ffn_mult=4):
79
+ super().__init__()
80
+ self.tok_emb = nn.Embedding(vocab, d)
81
+ self.layers = nn.ModuleList([Block(d, n_heads, head_dim, ffn_mult) for _ in range(n_layers)])
82
+ self.norm_f = RMSNorm(d)
83
+ self.head = nn.Linear(d, vocab, bias=False)
84
+ # Tied embeddings
85
+ self.head.weight = self.tok_emb.weight
86
+ # RoPE
87
+ freqs = 1.0 / (10000 ** (torch.arange(0, head_dim, 2).float() / head_dim))
88
+ t = torch.arange(2048).float()
89
+ outer = torch.outer(t, freqs)
90
+ angles = torch.cat([outer, outer], dim=-1)
91
+ self.register_buffer("cos", angles.cos())
92
+ self.register_buffer("sin", angles.sin())
93
+
94
+ def forward(self, x, y=None):
95
+ h = self.tok_emb(x)
96
+ for layer in self.layers:
97
+ h = layer(h, self.cos, self.sin)
98
+ h = self.norm_f(h)
99
+ logits = self.head(h)
100
+ if y is not None:
101
+ return F.cross_entropy(logits.reshape(-1, logits.size(-1)), y.reshape(-1))
102
+ return logits
103
+
104
+
105
+ def load_model(path="model.pt", device="cpu"):
106
+ model = TinyLM()
107
+ ckpt = torch.load(path, map_location=device, weights_only=True)
108
+ model.load_state_dict(ckpt["model"])
109
+ model.eval()
110
+ return model
111
+
112
+
113
+ def generate(model, tokenizer, prompt, max_new=100, temperature=0.8, rep_penalty=1.5):
114
+ ids = tokenizer.encode(prompt, add_special_tokens=False).ids
115
+ ids = torch.tensor([ids])
116
+ with torch.no_grad():
117
+ for _ in range(max_new):
118
+ inp = ids[:, -512:]
119
+ logits = model(inp)[0, -1]
120
+ for tid in ids[0].unique():
121
+ if logits[tid] > 0:
122
+ logits[tid] /= rep_penalty
123
+ else:
124
+ logits[tid] *= rep_penalty
125
+ logits = logits / temperature
126
+ probs = F.softmax(logits, dim=-1)
127
+ next_id = torch.multinomial(probs, 1).item()
128
+ ids = torch.cat([ids, torch.tensor([[next_id]])], dim=1)
129
+ return tokenizer.decode(ids[0].tolist(), skip_special_tokens=True)
130
+
131
+
132
+ if __name__ == "__main__":
133
+ from tokenizers import Tokenizer
134
+ tok = Tokenizer.from_file("tokenizer.json")
135
+ model = load_model()
136
+ print(f"Params: {sum(p.numel() for p in model.parameters()):,}")
137
+ print(generate(model, tok, "[INST] Hello! [/INST]", max_new=100))