DemonKing1234 commited on
Commit
2a78536
·
verified ·
1 Parent(s): 08e18ab

Upload train.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train.py +216 -0
train.py ADDED
@@ -0,0 +1,216 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import math
3
+ import os
4
+ from pathlib import Path
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ from torch.nn import functional as F
9
+
10
+
11
+ DEFAULT_CHARS = (
12
+ "\n"
13
+ " "
14
+ "abcdefghijklmnopqrstuvwxyz"
15
+ "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
16
+ "0123456789"
17
+ ".,!?;:'\"-_/\\()[]{}<>@#$%^&*+=|`~"
18
+ )
19
+
20
+
21
+ class TinyTransformerLM(nn.Module):
22
+ def __init__(self, vocab_size, block_size, n_embd=128, n_head=2, n_layer=2, dropout=0.1):
23
+ super().__init__()
24
+ self.block_size = block_size
25
+ self.token_embedding = nn.Embedding(vocab_size, n_embd)
26
+ self.position_embedding = nn.Embedding(block_size, n_embd)
27
+ encoder_layer = nn.TransformerEncoderLayer(
28
+ d_model=n_embd,
29
+ nhead=n_head,
30
+ dim_feedforward=4 * n_embd,
31
+ dropout=dropout,
32
+ activation="gelu",
33
+ batch_first=True,
34
+ )
35
+ self.blocks = nn.TransformerEncoder(encoder_layer, num_layers=n_layer)
36
+ self.ln_f = nn.LayerNorm(n_embd)
37
+ self.head = nn.Linear(n_embd, vocab_size)
38
+
39
+ def forward(self, idx, targets=None):
40
+ batch, time = idx.shape
41
+ if time > self.block_size:
42
+ raise ValueError("sequence is longer than block_size")
43
+
44
+ token_emb = self.token_embedding(idx)
45
+ pos = torch.arange(time, device=idx.device)
46
+ pos_emb = self.position_embedding(pos)[None, :, :]
47
+ x = token_emb + pos_emb
48
+
49
+ mask = torch.triu(torch.ones(time, time, device=idx.device), diagonal=1).bool()
50
+ x = self.blocks(x, mask=mask)
51
+ x = self.ln_f(x)
52
+ logits = self.head(x)
53
+
54
+ loss = None
55
+ if targets is not None:
56
+ loss = F.cross_entropy(logits.reshape(batch * time, -1), targets.reshape(batch * time))
57
+ return logits, loss
58
+
59
+
60
+ PRESETS = {
61
+ "tiny": {"block_size": 64, "n_embd": 64, "n_head": 2, "n_layer": 1, "batch_size": 4, "steps": 1200, "lr": 3e-4},
62
+ "turbo": {"block_size": 32, "n_embd": 64, "n_head": 4, "n_layer": 2, "batch_size": 16, "steps": 600, "lr": 1e-3},
63
+ "fast": {"block_size": 64, "n_embd": 96, "n_head": 3, "n_layer": 2, "batch_size": 8, "steps": 800, "lr": 5e-4},
64
+ "smart": {"block_size": 128, "n_embd": 160, "n_head": 4, "n_layer": 3, "batch_size": 12, "steps": 1500, "lr": 3e-4},
65
+ "power": {"block_size": 128, "n_embd": 256, "n_head": 8, "n_layer": 4, "batch_size": 16, "steps": 1000, "lr": 4e-4},
66
+ "small": {"block_size": 128, "n_embd": 128, "n_head": 2, "n_layer": 2, "batch_size": 8, "steps": 1200, "lr": 3e-4},
67
+ "big": {"block_size": 128, "n_embd": 192, "n_head": 4, "n_layer": 4, "batch_size": 4, "steps": 1200, "lr": 2e-4},
68
+ "large": {"block_size": 128, "n_embd": 256, "n_head": 8, "n_layer": 6, "batch_size": 2, "steps": 1200, "lr": 1.5e-4},
69
+ }
70
+
71
+
72
+ def build_vocab(text):
73
+ chars = sorted(set(DEFAULT_CHARS + text))
74
+ stoi = {ch: i for i, ch in enumerate(chars)}
75
+ itos = {i: ch for ch, i in stoi.items()}
76
+ return stoi, itos
77
+
78
+
79
+ def encode_text(text, stoi):
80
+ fallback = stoi.get(" ", 0)
81
+ return torch.tensor([stoi.get(ch, fallback) for ch in text], dtype=torch.long)
82
+
83
+
84
+ def make_batch(data, batch_size, block_size, device):
85
+ max_start = len(data) - block_size - 1
86
+ starts = torch.randint(max_start, (batch_size,))
87
+ x = torch.stack([data[i : i + block_size] for i in starts])
88
+ y = torch.stack([data[i + 1 : i + block_size + 1] for i in starts])
89
+ return x.to(device), y.to(device)
90
+
91
+
92
+ @torch.no_grad()
93
+ def estimate_loss(model, train_data, val_data, batch_size, block_size, device, eval_iters=20):
94
+ model.eval()
95
+ out = {}
96
+ for split, data in (("train", train_data), ("val", val_data)):
97
+ losses = []
98
+ for _ in range(eval_iters):
99
+ x, y = make_batch(data, batch_size, block_size, device)
100
+ _, loss = model(x, y)
101
+ losses.append(loss.item())
102
+ out[split] = sum(losses) / len(losses)
103
+ model.train()
104
+ return out
105
+
106
+
107
+ def main():
108
+ parser = argparse.ArgumentParser()
109
+ parser.add_argument("--data", default="data/input.txt")
110
+ parser.add_argument("--out", default="runs/tiny-char-model.pt")
111
+ parser.add_argument("--preset", choices=sorted(PRESETS), default="tiny")
112
+ parser.add_argument("--steps", type=int, default=1200)
113
+ parser.add_argument("--batch-size", type=int, default=16)
114
+ parser.add_argument("--block-size", type=int, default=128)
115
+ parser.add_argument("--n-embd", type=int, default=128)
116
+ parser.add_argument("--n-head", type=int, default=2)
117
+ parser.add_argument("--n-layer", type=int, default=2)
118
+ parser.add_argument("--lr", type=float, default=3e-4)
119
+ args = parser.parse_args()
120
+
121
+ preset = PRESETS[args.preset]
122
+ if args.steps == 1200:
123
+ args.steps = preset["steps"]
124
+ if args.batch_size == 16:
125
+ args.batch_size = preset["batch_size"]
126
+ if args.block_size == 128:
127
+ args.block_size = preset["block_size"]
128
+ if args.n_embd == 128:
129
+ args.n_embd = preset["n_embd"]
130
+ if args.n_head == 2:
131
+ args.n_head = preset["n_head"]
132
+ if args.n_layer == 2:
133
+ args.n_layer = preset["n_layer"]
134
+ if args.lr == 3e-4:
135
+ args.lr = preset["lr"]
136
+
137
+ device = "cuda" if torch.cuda.is_available() else "cpu"
138
+ if device == "cpu":
139
+ threads = os.cpu_count() or 4 # Use the actual number of logical CPU cores
140
+ torch.set_num_threads(threads)
141
+ torch.set_num_interop_threads(1)
142
+ torch.set_float32_matmul_precision("high")
143
+ print(f"CPU optimization: using {threads} threads")
144
+ else:
145
+ torch.backends.cuda.matmul.allow_tf32 = True
146
+ torch.backends.cudnn.benchmark = True
147
+ torch.set_float32_matmul_precision("high")
148
+ print("CUDA optimization: TF32 enabled, cudnn benchmark on")
149
+
150
+ text = Path(args.data).read_text(encoding="utf-8")
151
+ stoi, itos = build_vocab(text)
152
+ encoded = encode_text(text, stoi)
153
+
154
+ if len(encoded) < args.block_size + 2:
155
+ raise SystemExit("Dataset is too small. Add more text or lower --block-size.")
156
+
157
+ split = max(1, int(0.9 * len(encoded)))
158
+ train_data = encoded[:split]
159
+ val_data = encoded[split - args.block_size - 1 :]
160
+ chars = [ch for ch, _ in sorted(stoi.items(), key=lambda item: item[1])]
161
+
162
+ model = TinyTransformerLM(
163
+ vocab_size=len(chars),
164
+ block_size=args.block_size,
165
+ n_embd=args.n_embd,
166
+ n_head=args.n_head,
167
+ n_layer=args.n_layer,
168
+ ).to(device)
169
+ optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr)
170
+ scaler = torch.cuda.amp.GradScaler(enabled=device == "cuda")
171
+
172
+ params = sum(p.numel() for p in model.parameters())
173
+ print(f"device={device} params={params:,} vocab={len(chars)}")
174
+
175
+ for step in range(args.steps + 1):
176
+ if step % 100 == 0:
177
+ losses = estimate_loss(model, train_data, val_data, args.batch_size, args.block_size, device)
178
+ ppl = math.exp(min(losses["val"], 20))
179
+ print(f"step {step:5d} train {losses['train']:.4f} val {losses['val']:.4f} ppl {ppl:.2f}")
180
+
181
+ xb, yb = make_batch(train_data, args.batch_size, args.block_size, device)
182
+ if device == "cuda":
183
+ with torch.cuda.amp.autocast():
184
+ _, loss = model(xb, yb)
185
+ else:
186
+ _, loss = model(xb, yb)
187
+ optimizer.zero_grad(set_to_none=True)
188
+ scaler.scale(loss).backward() if device == "cuda" else loss.backward()
189
+ if device == "cuda":
190
+ scaler.step(optimizer)
191
+ scaler.update()
192
+ else:
193
+ optimizer.step()
194
+
195
+ out_path = Path(args.out)
196
+ out_path.parent.mkdir(parents=True, exist_ok=True)
197
+ torch.save(
198
+ {
199
+ "model": model.state_dict(),
200
+ "config": {
201
+ "vocab_size": len(chars),
202
+ "block_size": args.block_size,
203
+ "n_embd": args.n_embd,
204
+ "n_head": args.n_head,
205
+ "n_layer": args.n_layer,
206
+ },
207
+ "stoi": stoi,
208
+ "itos": itos,
209
+ },
210
+ out_path,
211
+ )
212
+ print(f"saved {out_path}")
213
+
214
+
215
+ if __name__ == "__main__":
216
+ main()