solhost commited on
Commit
d78b3c4
·
verified ·
1 Parent(s): 652fa5b

Create model.py

Browse files
Files changed (1) hide show
  1. model.py +192 -0
model.py ADDED
@@ -0,0 +1,192 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ model.py -- standalone architecture definition for GTM-3-base.
3
+
4
+ This is a plain PyTorch nanoGPT-style GPT model with RoPE (rotary position
5
+ embeddings), NOT a HuggingFace `transformers` AutoModel. To load the
6
+ released weights:
7
+
8
+ pip install torch safetensors tiktoken
9
+
10
+ import json, torch
11
+ from safetensors.torch import load_file
12
+ from model import GPT, GPTConfig
13
+
14
+ with open("config.json") as f:
15
+ config = GPTConfig(**json.load(f))
16
+ model = GPT(config)
17
+ state_dict = load_file("model.safetensors")
18
+ model.load_state_dict(state_dict)
19
+ model.eval()
20
+
21
+ import tiktoken
22
+ enc = tiktoken.get_encoding("gpt2")
23
+ ids = enc.encode_ordinary("Once upon a time,")
24
+ x = torch.tensor([ids], dtype=torch.long)
25
+ out = model.generate(x, max_new_tokens=100, temperature=0.8, top_k=50,
26
+ eot_token=enc.eot_token, repetition_penalty=1.3)
27
+ print(enc.decode(out[0].tolist()))
28
+ """
29
+
30
+ import math
31
+ from dataclasses import dataclass
32
+
33
+ import torch
34
+ import torch.nn as nn
35
+ import torch.nn.functional as F
36
+
37
+
38
+ @dataclass
39
+ class GPTConfig:
40
+ vocab_size: int = 50257
41
+ block_size: int = 1024
42
+ n_layer: int = 10
43
+ n_head: int = 8
44
+ n_embd: int = 608
45
+ dropout: float = 0.0
46
+ bias: bool = True
47
+ rope_theta: float = 10000.0 # standard RoPE base frequency
48
+
49
+
50
+ def precompute_rope_freqs(head_dim, max_seq_len, theta=10000.0, device="cpu"):
51
+ """Precompute the complex rotation frequencies used by RoPE, one pair
52
+ per (position, frequency-band). Standard formula: theta_i = theta^(-2i/dim)."""
53
+ assert head_dim % 2 == 0, "RoPE requires an even head_dim"
54
+ freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
55
+ t = torch.arange(max_seq_len, device=device).float()
56
+ freqs = torch.outer(t, freqs) # (max_seq_len, head_dim/2)
57
+ return torch.polar(torch.ones_like(freqs), freqs) # complex64, (max_seq_len, head_dim/2)
58
+
59
+
60
+ def apply_rope(x, freqs_cis):
61
+ """Apply rotary position embeddings to a (B, n_head, T, head_dim) tensor."""
62
+ B, n_head, T, head_dim = x.shape
63
+ x_complex = torch.view_as_complex(x.float().reshape(B, n_head, T, head_dim // 2, 2))
64
+ freqs_cis = freqs_cis[:T].view(1, 1, T, head_dim // 2)
65
+ x_rotated = x_complex * freqs_cis
66
+ x_out = torch.view_as_real(x_rotated).reshape(B, n_head, T, head_dim)
67
+ return x_out.type_as(x)
68
+
69
+
70
+ class CausalSelfAttention(nn.Module):
71
+ def __init__(self, config):
72
+ super().__init__()
73
+ assert config.n_embd % config.n_head == 0
74
+ self.n_head = config.n_head
75
+ self.n_embd = config.n_embd
76
+ self.head_dim = config.n_embd // config.n_head
77
+ self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=config.bias)
78
+ self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias)
79
+ self.attn_dropout = nn.Dropout(config.dropout)
80
+ self.resid_dropout = nn.Dropout(config.dropout)
81
+ self.dropout = config.dropout
82
+
83
+ def forward(self, x, freqs_cis):
84
+ B, T, C = x.shape
85
+ q, k, v = self.c_attn(x).split(self.n_embd, dim=2)
86
+ q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
87
+ k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
88
+ v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
89
+ # RoPE is applied to queries and keys only, not values -- this is what
90
+ # makes attention scores depend on *relative* position between tokens
91
+ q = apply_rope(q, freqs_cis)
92
+ k = apply_rope(k, freqs_cis)
93
+ y = F.scaled_dot_product_attention(
94
+ q, k, v, is_causal=True,
95
+ dropout_p=self.dropout if self.training else 0.0,
96
+ )
97
+ y = y.transpose(1, 2).contiguous().view(B, T, C)
98
+ return self.resid_dropout(self.c_proj(y))
99
+
100
+
101
+ class MLP(nn.Module):
102
+ def __init__(self, config):
103
+ super().__init__()
104
+ self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=config.bias)
105
+ self.gelu = nn.GELU()
106
+ self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=config.bias)
107
+ self.dropout = nn.Dropout(config.dropout)
108
+
109
+ def forward(self, x):
110
+ return self.dropout(self.c_proj(self.gelu(self.c_fc(x))))
111
+
112
+
113
+ class Block(nn.Module):
114
+ def __init__(self, config):
115
+ super().__init__()
116
+ self.ln_1 = nn.LayerNorm(config.n_embd)
117
+ self.attn = CausalSelfAttention(config)
118
+ self.ln_2 = nn.LayerNorm(config.n_embd)
119
+ self.mlp = MLP(config)
120
+
121
+ def forward(self, x, freqs_cis):
122
+ x = x + self.attn(self.ln_1(x), freqs_cis)
123
+ x = x + self.mlp(self.ln_2(x))
124
+ return x
125
+
126
+
127
+ class GPT(nn.Module):
128
+ def __init__(self, config):
129
+ super().__init__()
130
+ self.config = config
131
+ self.transformer = nn.ModuleDict(dict(
132
+ wte=nn.Embedding(config.vocab_size, config.n_embd),
133
+ # NOTE: no wpe (learned position embedding) -- RoPE replaces it entirely,
134
+ # applied inside attention rather than added to the input embeddings
135
+ drop=nn.Dropout(config.dropout),
136
+ h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]),
137
+ ln_f=nn.LayerNorm(config.n_embd),
138
+ ))
139
+ self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
140
+ self.transformer.wte.weight = self.lm_head.weight
141
+ head_dim = config.n_embd // config.n_head
142
+ freqs_cis = precompute_rope_freqs(head_dim, config.block_size, theta=config.rope_theta)
143
+ self.register_buffer("freqs_cis", freqs_cis, persistent=False)
144
+ self.apply(self._init_weights)
145
+ for pn, p in self.named_parameters():
146
+ if pn.endswith("c_proj.weight"):
147
+ nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
148
+
149
+ def _init_weights(self, module):
150
+ if isinstance(module, nn.Linear):
151
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
152
+ if module.bias is not None:
153
+ nn.init.zeros_(module.bias)
154
+ elif isinstance(module, nn.Embedding):
155
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
156
+
157
+ def forward(self, idx, targets=None):
158
+ B, T = idx.shape
159
+ assert T <= self.config.block_size, "sequence longer than block_size"
160
+ x = self.transformer.drop(self.transformer.wte(idx))
161
+ freqs_cis = self.freqs_cis.to(x.device)
162
+ for block in self.transformer.h:
163
+ x = block(x, freqs_cis)
164
+ x = self.transformer.ln_f(x)
165
+ logits = self.lm_head(x)
166
+ loss = None
167
+ if targets is not None:
168
+ loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-1)
169
+ return logits, loss
170
+
171
+ @torch.no_grad()
172
+ def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None, eot_token=None,
173
+ repetition_penalty=1.0):
174
+ for _ in range(max_new_tokens):
175
+ idx_cond = idx if idx.size(1) <= self.config.block_size else idx[:, -self.config.block_size:]
176
+ logits, _ = self(idx_cond)
177
+ logits = logits[:, -1, :] / temperature
178
+ if repetition_penalty != 1.0:
179
+ for seen_id in set(idx[0].tolist()):
180
+ if logits[0, seen_id] > 0:
181
+ logits[0, seen_id] /= repetition_penalty
182
+ else:
183
+ logits[0, seen_id] *= repetition_penalty
184
+ if top_k is not None:
185
+ v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
186
+ logits[logits < v[:, [-1]]] = float("-inf")
187
+ probs = F.softmax(logits, dim=-1)
188
+ idx_next = torch.multinomial(probs, num_samples=1)
189
+ idx = torch.cat((idx, idx_next), dim=1)
190
+ if eot_token is not None and idx_next.item() == eot_token:
191
+ break
192
+ return idx