Views
No views yet
28m_model.pt and model_config.pt in the same directory as the script below.1import torch
2import torch.nn as nn
3from torch.nn import functional as F
4import math
5import tiktoken
6import time
7import os
8
9device = 'cuda' if torch.cuda.is_available() else 'cpu'
10print(f"Using device: {device}")
11
12PROMPT = "INSERT YOUR PROMPT HERE"
13
14n_embd = 384
15n_head = 6
16n_layer = 8
17block_size = 128
18dropout = 0.20
19
20enc = tiktoken.get_encoding("gpt2")
21vocab_size = enc.n_vocab
22
23class MultiHeadAttention(nn.Module):
24 def __init__(self, n_embd, n_head, block_size, dropout=0.1):
25 super().__init__()
26 assert n_embd % n_head == 0
27 self.n_head = n_head
28 self.head_size = n_embd // n_head
29
30 self.qkv = nn.Linear(n_embd, 3 * n_embd, bias=False)
31 self.proj = nn.Linear(n_embd, n_embd, bias=False)
32 self.dropout = nn.Dropout(dropout)
33
34 self.register_buffer("bias", torch.tril(torch.ones(block_size, block_size))
35 .view(1, 1, block_size, block_size))
36
37 def forward(self, x):
38 B, T, C = x.size()
39 qkv = self.qkv(x).view(B, T, 3, self.n_head, self.head_size).permute(2, 0, 3, 1, 4)
40 q, k, v = qkv[0], qkv[1], qkv[2]
41
42 att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.head_size))
43 att = att.masked_fill(self.bias[:,:,:T,:T] == 0, float('-inf'))
44 att = F.softmax(att, dim=-1)
45 att = self.dropout(att)
46
47 y = att @ v
48 y = y.transpose(1, 2).contiguous().view(B, T, C)
49 y = self.proj(y)
50 return y
51
52class FeedForward(nn.Module):
53 def __init__(self, n_embd, dropout=0.1):
54 super().__init__()
55 self.net = nn.Sequential(
56 nn.Linear(n_embd, 2 * n_embd, bias=False),
57 nn.GELU(),
58 nn.Linear(2 * n_embd, n_embd, bias=False),
59 nn.Dropout(dropout),
60 )
61
62 def forward(self, x):
63 return self.net(x)
64
65class Block(nn.Module):
66 def __init__(self, n_embd, n_head, block_size, dropout=0.1):
67 super().__init__()
68 self.ln1 = nn.LayerNorm(n_embd, bias=False)
69 self.attn = MultiHeadAttention(n_embd, n_head, block_size, dropout)
70 self.ln2 = nn.LayerNorm(n_embd, bias=False)
71 self.mlp = FeedForward(n_embd, dropout)
72
73 def forward(self, x):
74 x = x + self.attn(self.ln1(x))
75 x = x + self.mlp(self.ln2(x))
76 return x
77
78class BigramLanguageModel(nn.Module):
79 def __init__(self, vocab_size, n_embd, n_head, n_layer, block_size, dropout=0.1):
80 super().__init__()
81 self.block_size = block_size
82
83 self.token_embedding_table = nn.Embedding(vocab_size, n_embd)
84 self.position_embedding_table = nn.Embedding(block_size, n_embd)
85 self.blocks = nn.Sequential(*[Block(n_embd, n_head, block_size, dropout) for _ in range(n_layer)])
86 self.ln_f = nn.LayerNorm(n_embd, bias=False)
87 self.lm_head = nn.Linear(n_embd, vocab_size, bias=False)
88
89 self.lm_head.weight = self.token_embedding_table.weight
90
91 self.apply(self._init_weights)
92
93 def _init_weights(self, module):
94 if isinstance(module, nn.Linear):
95 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
96 elif isinstance(module, nn.Embedding):
97 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
98
99 def forward(self, idx, targets=None):
100 B, T = idx.shape
101
102 tok_emb = self.token_embedding_table(idx)
103 pos_emb = self.position_embedding_table(torch.arange(T, device=idx.device))
104 x = tok_emb + pos_emb
105 x = self.blocks(x)
106 x = self.ln_f(x)
107 logits = self.lm_head(x)
108
109 if targets is None:
110 loss = None
111 else:
112 B, T, C = logits.shape
113 logits = logits.view(B*T, C)
114 targets = targets.view(B*T)
115 loss = F.cross_entropy(logits, targets)
116
117 return logits, loss
118
119def encode(text):
120 return enc.encode(text)
121
122def decode(tokens):
123 return enc.decode(tokens)
124
125def generate_text(model, prompt="", max_tokens=500, temperature=0.8, top_k=50):
126 model.eval()
127
128 if prompt:
129 context = torch.tensor([encode(prompt)], dtype=torch.long, device=device)
130 else:
131 context = torch.zeros((1, 1), dtype=torch.long, device=device)
132
133 generated_tokens = []
134
135 with torch.no_grad():
136 for i in range(max_tokens):
137 idx_cond = context if context.size(1) <= block_size else context[:, -block_size:]
138 logits, _ = model(idx_cond)
139 logits = logits[:, -1, :] / temperature
140
141 if top_k is not None:
142 v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
143 logits[logits < v[:, [-1]]] = -float('Inf')
144
145 probs = F.softmax(logits, dim=-1)
146 idx_next = torch.multinomial(probs, num_samples=1)
147 context = torch.cat((context, idx_next), dim=1)
148 generated_tokens.append(idx_next.item())
149
150 return decode(generated_tokens)
151
152def load_model():
153 if not os.path.exists('model_config.pt'):
154 print("Model config file not found: model_config.pt")
155 return None
156
157 if not os.path.exists('28m_model.pt'):
158 print("Model weights file not found: 28m_model.pt")
159 return None
160
161 config = torch.load('model_config.pt', map_location=device)
162 print("Configuration loaded")
163
164 model = BigramLanguageModel(
165 vocab_size=vocab_size,
166 n_embd=n_embd,
167 n_head=n_head,
168 n_layer=n_layer,
169 block_size=block_size,
170 dropout=dropout
171 )
172
173 checkpoint = torch.load('28m_model.pt', map_location=device)
174 model.load_state_dict(checkpoint['model_state_dict'])
175 model = model.to(device)
176 model.eval()
177
178 print("Model loaded successfully!")
179
180 return model
181
182if __name__ == "__main__":
183 print("Loading model...")
184
185 try:
186 model = load_model()
187 if model is None:
188 exit()
189
190 start_time = time.time()
191 generated_text = generate_text(
192 model,
193 prompt=PROMPT,
194 max_tokens=150,
195 temperature=0.8,
196 top_k=50
197 )
198 generation_time = time.time() - start_time
199
200 print(f"\nGENERATED TEXT ({generation_time:.1f}s):")
201 print("-" * 50)
202 print(f"{PROMPT}{generated_text}")
203 print("-" * 50)
204
205 except Exception as e:
206 print(f"Error: {e}")
207 print("Please make sure:")
208 print("1. The model files (28m_model.pt and model_config.pt) are in the same directory")
209 print("2. The model architecture matches the code")