ParallaxOpen commited on
Commit
e861cd4
·
verified ·
1 Parent(s): 912e0dd

Upload train_v3.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train_v3.py +288 -0
train_v3.py ADDED
@@ -0,0 +1,288 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Parallax-Chess v3: Predict move as single index (from_sq * 64 + to_sq).
3
+ Fixes the v2 bug where only from-square was supervised.
4
+ """
5
+ import sys, math, time, json, random
6
+ from pathlib import Path
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ from torch.utils.data import Dataset, DataLoader
11
+ import chess
12
+
13
+ ROOT = Path(__file__).parent.parent
14
+ sys.path.insert(0, str(ROOT / "scripts"))
15
+ sys.path.insert(0, str(ROOT / "src"))
16
+ sys.path.insert(0, str(ROOT.parent))
17
+ from centauri.models.base import SmallLMConfig
18
+ from centauri.models.torch.small_lm import SmallLM
19
+ from chess_encoding import encode_board, VOCAB_SIZE
20
+
21
+ random.seed(42)
22
+
23
+ # Move index: from_sq * 64 + to_sq for non-promo, then promo moves
24
+ # Total: 64*64 = 4096 normal + 4*8*8 = 256 promo = 4352 total
25
+ MOVE_INDEX_SIZE = 4352 # 64*64 + 4*64 (promotions)
26
+
27
+
28
+ def move_to_index(move):
29
+ """Convert chess.Move to index 0..4351."""
30
+ if move.promotion:
31
+ promo_map = {chess.QUEEN: 0, chess.ROOK: 1, chess.BISHOP: 2, chess.KNIGHT: 3}
32
+ return 4096 + promo_map[move.promotion] * 64 + move.to_square
33
+ return move.from_square * 64 + move.to_square
34
+
35
+
36
+ def index_to_move(idx, board):
37
+ """Convert index back to chess.Move, only if legal."""
38
+ if idx >= 4096:
39
+ promo_offset = idx - 4096
40
+ promo_type = promo_offset // 64
41
+ to_sq = promo_offset % 64
42
+ promo_map = {0: chess.QUEEN, 1: chess.ROOK, 2: chess.BISHOP, 3: chess.KNIGHT}
43
+ promotion = promo_map[promo_type]
44
+ # Find from_square (must be a pawn on correct file)
45
+ for from_sq in chess.SQUARES:
46
+ piece = board.piece_at(from_sq)
47
+ if piece and piece.piece_type == chess.PAWN and piece.color == board.turn:
48
+ move = chess.Move(from_sq, to_sq, promotion=promotion)
49
+ if move in board.legal_moves:
50
+ return move
51
+ return None
52
+ else:
53
+ from_sq = idx // 64
54
+ to_sq = idx % 64
55
+ move = chess.Move(from_sq, to_sq)
56
+ if move in board.legal_moves:
57
+ return move
58
+ return None
59
+
60
+
61
+ class ChessDataset(Dataset):
62
+ """Load SF-evaluated positions."""
63
+
64
+ def __init__(self, data_paths):
65
+ self.samples = []
66
+ for path in data_paths:
67
+ path = Path(path)
68
+ if not path.exists():
69
+ continue
70
+ with open(path, "r", encoding="utf-8") as f:
71
+ for line in f:
72
+ try:
73
+ item = json.loads(line)
74
+ board = chess.Board(item["fen"])
75
+ move_key = "best_move" if "best_move" in item else "move"
76
+ move = chess.Move.from_uci(item[move_key])
77
+ if move not in board.legal_moves:
78
+ continue
79
+
80
+ board_tokens = encode_board(board)
81
+ move_idx = move_to_index(move)
82
+
83
+ if "eval_cp" in item:
84
+ value = max(-1.0, min(1.0, item["eval_cp"] / 1000.0))
85
+ elif "result" in item:
86
+ r = item["result"]
87
+ value = 1.0 if r == "1-0" else -1.0 if r == "0-1" else 0.0
88
+ else:
89
+ value = 0.0
90
+
91
+ self.samples.append((board_tokens, move_idx, value))
92
+
93
+ if random.random() < 0.5:
94
+ flipped = flip_board(board_tokens)
95
+ flipped_move = flip_move_index(move_idx)
96
+ self.samples.append((flipped, flipped_move, value))
97
+ except:
98
+ pass
99
+ print("Loaded %d samples" % len(self.samples))
100
+
101
+ def __len__(self):
102
+ return len(self.samples)
103
+
104
+ def __getitem__(self, i):
105
+ board_tokens, move_idx, value = self.samples[i]
106
+ return (
107
+ torch.tensor(board_tokens, dtype=torch.long),
108
+ torch.tensor(move_idx, dtype=torch.long),
109
+ torch.tensor(value, dtype=torch.float),
110
+ )
111
+
112
+
113
+ def flip_board(board_tokens):
114
+ flipped = list(board_tokens)
115
+ for r in range(8):
116
+ for f in range(8):
117
+ flipped[r * 8 + (7 - f)] = board_tokens[r * 8 + f]
118
+ return flipped
119
+
120
+
121
+ def flip_sq(sq):
122
+ """Mirror a square horizontally: file f -> 7-f."""
123
+ rank = sq // 8
124
+ file = sq % 8
125
+ return rank * 8 + (7 - file)
126
+
127
+
128
+ def flip_move_index(idx):
129
+ """Flip move index when board is mirrored horizontally."""
130
+ if idx >= 4096:
131
+ promo_offset = idx - 4096
132
+ promo_type = promo_offset // 64
133
+ to_sq = promo_offset % 64
134
+ return 4096 + promo_type * 64 + flip_sq(to_sq)
135
+ from_sq = idx // 64
136
+ to_sq = idx % 64
137
+ return flip_sq(from_sq) * 64 + flip_sq(to_sq)
138
+
139
+
140
+ class ParallaxChessV3(nn.Module):
141
+ """Board -> move index (4352 classes) + value."""
142
+
143
+ def __init__(self, cfg):
144
+ super().__init__()
145
+ self.cfg = cfg
146
+ self.board_encoder = nn.Embedding(VOCAB_SIZE, cfg.d_model)
147
+ from centauri.models.torch.small_lm import SmallLM
148
+ self.backbone = SmallLM(cfg)
149
+ self.policy_head = nn.Linear(cfg.d_model, MOVE_INDEX_SIZE)
150
+ self.value_head = nn.Sequential(
151
+ nn.Linear(cfg.d_model, 256), nn.ReLU(), nn.Dropout(0.1),
152
+ nn.Linear(256, 1), nn.Tanh())
153
+
154
+ def forward(self, board_tokens):
155
+ x = self.board_encoder(board_tokens)
156
+ freq, _ = __import__('centauri.models.torch.small_lm', fromlist=['precompute_rope']).precompute_rope(
157
+ self.cfg, x.device, x.dtype)
158
+ for layer in self.backbone.layers:
159
+ x, _ = layer(x, freq, None)
160
+ pooled = x.mean(dim=1)
161
+ return {"logits": self.policy_head(pooled), "value": self.value_head(pooled).squeeze(-1)}
162
+
163
+ def predict_move(self, board, device=None):
164
+ if device is None:
165
+ device = next(self.parameters()).device
166
+ board_tokens = encode_board(board)
167
+ inp = torch.tensor([board_tokens], dtype=torch.long).to(device)
168
+ with torch.no_grad():
169
+ out = self(inp)
170
+ probs = torch.softmax(out["logits"], dim=-1)
171
+
172
+ # Score only legal moves
173
+ best_move, best_score = None, -float("inf")
174
+ for move in board.legal_moves:
175
+ idx = move_to_index(move)
176
+ score = probs[0, idx].item()
177
+ if score > best_score:
178
+ best_score = score
179
+ best_move = move
180
+ return best_move
181
+
182
+ def evaluate(self, board, device=None):
183
+ if device is None:
184
+ device = next(self.parameters()).device
185
+ board_tokens = encode_board(board)
186
+ inp = torch.tensor([board_tokens], dtype=torch.long).to(device)
187
+ with torch.no_grad():
188
+ return self(inp)["value"].item() * 1000
189
+
190
+
191
+ def main():
192
+ import argparse
193
+ parser = argparse.ArgumentParser()
194
+ parser.add_argument("--batch_size", type=int, default=64)
195
+ parser.add_argument("--lr", type=float, default=1e-3)
196
+ parser.add_argument("--max_steps", type=int, default=80000)
197
+ parser.add_argument("--save_dir", default=str(ROOT / "checkpoints" / "parallax_chess_v3"))
198
+ args = parser.parse_args()
199
+
200
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
201
+ print("Device:", device)
202
+
203
+ cfg_dict = {
204
+ "vocab_size": VOCAB_SIZE, "d_model": 512, "n_heads": 8, "n_kv_heads": 4,
205
+ "n_layers": 8, "intermediate_size": 2048, "max_seq_len": 128,
206
+ "norm_type": "rms", "rope_type": "neox", "n_experts": 0,
207
+ "n_loops": 1, "loop_mode": "per_layer",
208
+ }
209
+ cfg = SmallLMConfig(**cfg_dict)
210
+ model = ParallaxChessV3(cfg).to(device)
211
+ n_params = sum(p.numel() for p in model.parameters())
212
+ print("Params: %d (%.1fM)" % (n_params, n_params / 1e6))
213
+
214
+ data_paths = [
215
+ ROOT / "data" / "chess_train_sf.jsonl",
216
+ ROOT / "data" / "chess_train_large.jsonl",
217
+ ]
218
+ ds = ChessDataset(data_paths)
219
+ loader = DataLoader(ds, batch_size=args.batch_size, shuffle=True, num_workers=0, pin_memory=True)
220
+
221
+ save_dir = Path(args.save_dir)
222
+ save_dir.mkdir(parents=True, exist_ok=True)
223
+
224
+ step = 0
225
+ for f in sorted(save_dir.glob("step_*.pt")):
226
+ try:
227
+ ckpt = torch.load(str(f), map_location=device)
228
+ model.load_state_dict(ckpt["model"])
229
+ step = ckpt.get("step", 0)
230
+ print("Resumed from %s at step %d" % (f.name, step))
231
+ break
232
+ except:
233
+ pass
234
+
235
+ total_steps = min(args.max_steps, len(ds) // args.batch_size * 5)
236
+ opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01, betas=(0.9, 0.98))
237
+
238
+ def lr_schedule(s):
239
+ warmup = 2000
240
+ if s < warmup:
241
+ return s / warmup
242
+ progress = (s - warmup) / max(1, total_steps - warmup)
243
+ return max(0.05, 0.5 * (1.0 + math.cos(math.pi * progress)))
244
+
245
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, lr_schedule)
246
+
247
+ model.train()
248
+ t0 = time.time()
249
+
250
+ for epoch in range(999):
251
+ for board_tokens, move_indices, values in loader:
252
+ board_tokens = board_tokens.to(device)
253
+ move_indices = move_indices.to(device)
254
+ values = values.to(device)
255
+
256
+ out = model(board_tokens)
257
+
258
+ policy_loss = F.cross_entropy(out["logits"], move_indices)
259
+ value_loss = F.mse_loss(out["value"], values)
260
+ loss = policy_loss + 1.0 * value_loss
261
+
262
+ loss.backward()
263
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
264
+ opt.step()
265
+ opt.zero_grad()
266
+ sched.step()
267
+
268
+ if step % 100 == 0:
269
+ lr = sched.get_last_lr()[0]
270
+ print("Step %6d | P:%.4f V:%.4f | Loss:%.4f | LR %.1e" % (
271
+ step, policy_loss.item(), value_loss.item(), loss.item(), lr))
272
+
273
+ if step > 0 and step % 10000 == 0:
274
+ ckpt = save_dir / ("step_%d.pt" % step)
275
+ torch.save({"model": model.state_dict(), "step": step, "config": cfg_dict, "n_params": n_params}, str(ckpt))
276
+ print("Saved %s" % ckpt.name)
277
+
278
+ step += 1
279
+ if step >= total_steps:
280
+ break
281
+
282
+ final = save_dir / "final.pt"
283
+ torch.save({"model": model.state_dict(), "step": step, "config": cfg_dict, "n_params": n_params}, str(final))
284
+ print("Done! %s (%d steps)" % (final, step))
285
+
286
+
287
+ if __name__ == "__main__":
288
+ main()