jgalego commited on
Commit
476298f
·
verified ·
1 Parent(s): 823cd4f

promote aux0.1

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ results/map.png filter=lfs diff=lfs merge=lfs -text
bridge2vec.py ADDED
@@ -0,0 +1,605 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # requires-python = ">=3.12"
3
+ # dependencies = [
4
+ # "datasets",
5
+ # "endplay",
6
+ # "huggingface-hub",
7
+ # "jinja2",
8
+ # "matplotlib",
9
+ # "numpy",
10
+ # "safetensors",
11
+ # "torch",
12
+ # "umap-learn",
13
+ # ]
14
+ # ///
15
+ """Bridge2Vec: embed bridge hands by how they take tricks, learned from double-dummy tables."""
16
+
17
+ # pylint: disable=too-many-arguments,too-many-instance-attributes,too-many-locals
18
+ # pylint: disable=too-many-positional-arguments,too-many-statements
19
+
20
+ import argparse
21
+ import json
22
+ import math
23
+ import os
24
+ import random
25
+ import sys
26
+ import time
27
+ from multiprocessing import Pool
28
+ from pathlib import Path
29
+
30
+ import datasets
31
+ import huggingface_hub.utils
32
+ import matplotlib.pyplot as plt
33
+ import numpy as np
34
+ import torch
35
+ import umap
36
+ from datasets import Dataset, load_dataset
37
+ from endplay._dds import SetMaxThreads
38
+ from endplay.dds import calc_all_tables
39
+ from endplay.types import Deal
40
+ from huggingface_hub import (
41
+ HfApi,
42
+ ModelCard,
43
+ ModelCardData,
44
+ PyTorchModelHubMixin,
45
+ snapshot_download,
46
+ )
47
+ from huggingface_hub.errors import RepositoryNotFoundError
48
+ from torch import nn
49
+ from torch.nn import functional
50
+
51
+ REPO = "jgalego/bridge2vec"
52
+ DATA = "jgalego/bridge2vec-deals"
53
+ HERE = Path(__file__).parent
54
+ # Progress bars redraw in place, which shows up as garbage in HF Jobs logs.
55
+ PROGRESS = sys.stderr.isatty()
56
+ RANKS = "23456789TJQKA"
57
+ SEATS = "NESW"
58
+ # Suits in PBN order, then notrump: the rows of a double-dummy table.
59
+ STRAINS = "SHDCN"
60
+ # DDS solves at most 32 tables per call; a test hand gets one call's worth of deals.
61
+ CHUNK = 32
62
+
63
+
64
+ def cpus():
65
+ """Return the CPU quota visible to this process."""
66
+ try:
67
+ quota, period = Path("/sys/fs/cgroup/cpu.max").read_text(encoding="utf-8").split()
68
+ if quota != "max":
69
+ return max(1, int(quota) // int(period))
70
+ except OSError:
71
+ pass
72
+ return len(os.sched_getaffinity(0))
73
+
74
+
75
+ def hand_pbn(cards):
76
+ """Write card ids (13 * suit + rank) as a PBN hand, e.g. AKQ32.KJ4.T9.A87."""
77
+ suits = [sorted((c % 13 for c in cards if c // 13 == s), reverse=True) for s in range(4)]
78
+ return ".".join("".join(RANKS[r] for r in suit) for suit in suits)
79
+
80
+
81
+ def parse_hand(text):
82
+ """Card ids of a PBN hand."""
83
+ suits = text.split(".")
84
+ if len(suits) != 4:
85
+ raise ValueError(f"{text!r} needs four suits separated by dots")
86
+ cards = [13 * s + RANKS.index(r) for s, ranks in enumerate(suits) for r in ranks.upper()]
87
+ if len(set(cards)) != 13:
88
+ raise ValueError(f"{text!r} is not 13 different cards")
89
+ return cards
90
+
91
+
92
+ def parse_deal(text):
93
+ """Card ids of a PBN deal with North first, shape (4, 13)."""
94
+ hands = [parse_hand(h) for h in text.removeprefix("N:").split()]
95
+ if len(hands) != 4 or len({c for h in hands for c in h}) != 52:
96
+ raise ValueError(f"{text!r} is not four hands of one deck")
97
+ return hands
98
+
99
+
100
+ def tensors(rows):
101
+ """Cards (n, 4, 13) and double-dummy tables (n, 5, 4) of a split."""
102
+ cards = torch.tensor([parse_deal(d) for d in rows["deal"]])
103
+ return cards, torch.tensor(np.array(rows["dd"]))
104
+
105
+
106
+ def profile(cards):
107
+ """HCP and suit lengths of hands, shape (..., 5)."""
108
+ hcp = (cards % 13 - 8).clamp(min=0).sum(-1, keepdim=True)
109
+ return torch.cat([hcp, functional.one_hot(cards // 13, 4).sum(-2)], -1).float()
110
+
111
+
112
+ def permute_suits(cards, dd):
113
+ """Relabel the suits of each deal at random, moving the table rows with them."""
114
+ perm = torch.rand(len(cards), 4, device=cards.device).argsort(1)
115
+ cards = perm.gather(1, (cards // 13).flatten(1)).view_as(cards) * 13 + cards % 13
116
+ rows = perm.argsort(1)[:, :, None].expand(-1, -1, 4)
117
+ return cards, torch.cat([dd[:, :4].gather(1, rows), dd[:, 4:]], 1)
118
+
119
+
120
+ class Bridge2Vec(nn.Module, PyTorchModelHubMixin):
121
+ """Transformer over the 13 cards of a hand; an MLP on four hands predicts the table."""
122
+
123
+ def __init__(self, dim=256, depth=4, heads=8, embed_dim=128, hidden=1024):
124
+ super().__init__()
125
+ self.suits = nn.Embedding(4, dim)
126
+ self.ranks = nn.Embedding(13, dim)
127
+ layer = nn.TransformerEncoderLayer(
128
+ dim, heads, 4 * dim, dropout=0.0, batch_first=True, norm_first=True
129
+ )
130
+ self.encoder = nn.TransformerEncoder(layer, depth, enable_nested_tensor=False)
131
+ self.norm = nn.LayerNorm(dim)
132
+ self.project = nn.Sequential(nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, embed_dim))
133
+ self.table = nn.Sequential(
134
+ nn.Linear(4 * embed_dim, hidden),
135
+ nn.GELU(),
136
+ nn.Linear(hidden, hidden),
137
+ nn.GELU(),
138
+ nn.Linear(hidden, 5 * 14),
139
+ )
140
+ self.value = nn.Linear(embed_dim, 5 * 4)
141
+ nn.init.constant_(self.value.bias, 6.5)
142
+ self.shape = nn.Linear(embed_dim, 5)
143
+
144
+ def encode(self, cards):
145
+ """Unit-length embeddings of hands, cards shape (n, 13)."""
146
+ x = self.suits(cards // 13) + self.ranks(cards % 13)
147
+ pooled = self.norm(self.encoder(x)).mean(1)
148
+ return functional.normalize(self.project(pooled), dim=-1)
149
+
150
+ def forward(self, cards):
151
+ """Hand embeddings, table logits, expected tables and profiles of deals (n, 4, 13).
152
+
153
+ Row d of the table comes from the hands in the order declarer d, LHO, partner, RHO,
154
+ so rotating the seats rotates the table. Expected tables are per hand, with
155
+ declarers relative to it: itself, LHO, partner, RHO.
156
+ """
157
+ n = len(cards)
158
+ z = self.encode(cards.flatten(0, 1)).view(n, 4, -1)
159
+ views = torch.stack([z.roll(-d, 1).flatten(1) for d in range(4)], 1)
160
+ logits = self.table(views).view(n, 4, 5, 14).transpose(1, 2)
161
+ return z, logits, self.value(z).view(n, 4, 5, 4), self.shape(z)
162
+
163
+ @torch.no_grad()
164
+ def run(self, cards, batch_size=4096):
165
+ """Hand embeddings, double-dummy tables and expected tables for deals (n, 4, 13)."""
166
+ device = next(self.parameters()).device
167
+ starts = range(0, len(cards), batch_size)
168
+ parts = [self(cards[i : i + batch_size].to(device)) for i in starts]
169
+ z, logits, value, _ = (torch.cat(p).float().cpu() for p in zip(*parts))
170
+ return {"hands": z, "tricks": logits.argmax(-1), "value": value}
171
+
172
+ @torch.no_grad()
173
+ def embed(self, hands):
174
+ """Embeddings and expected tables for PBN hands."""
175
+ device = next(self.parameters()).device
176
+ z = self.encode(torch.tensor([parse_hand(h) for h in hands], device=device))
177
+ return z.float().cpu(), self.value(z).view(-1, 5, 4).float().cpu()
178
+
179
+
180
+ def solve(job):
181
+ """32 deals with their double-dummy tables. Test deals share North. Seeded per chunk."""
182
+ split, index, seed = job
183
+ rng = random.Random(f"{seed}:{split}:{index}")
184
+ north = rng.sample(range(52), 13) if split == "test" else []
185
+ deals = []
186
+ for _ in range(CHUNK):
187
+ rest = [c for c in range(52) if c not in north]
188
+ rng.shuffle(rest)
189
+ cards = north + rest
190
+ deals.append("N:" + " ".join(hand_pbn(cards[i : i + 13]) for i in range(0, 52, 13)))
191
+ tables = calc_all_tables([Deal(d) for d in deals])
192
+ return [{"deal": d, "dd": t.to_list()} for d, t in zip(deals, tables)]
193
+
194
+
195
+ def data(args):
196
+ """Deal random hands, solve them double dummy; save as parquet, optionally push."""
197
+ start = time.time()
198
+ jobs = [("test", i) for i in range(args.test_hands)]
199
+ jobs += [("train", i) for i in range(args.deals // CHUNK)]
200
+ jobs = jobs[args.shard :: args.shards]
201
+ rows = {"train": [], "test": []}
202
+ step = max(1, len(jobs) // 20)
203
+ with Pool(cpus(), initializer=SetMaxThreads, initargs=(1,)) as pool:
204
+ tasks = [(split, i, args.seed) for split, i in jobs]
205
+ for n, ((split, _), part) in enumerate(zip(jobs, pool.imap(solve, tasks)), 1):
206
+ rows[split] += part
207
+ if n % step == 0 or n == len(jobs):
208
+ rate = n * CHUNK / (time.time() - start)
209
+ print(json.dumps({"chunks": n, "of": len(jobs), "deals_per_s": round(rate, 1)}),
210
+ flush=True)
211
+ out = Path(args.output)
212
+ out.mkdir(parents=True, exist_ok=True)
213
+ for split, part in rows.items():
214
+ if not part:
215
+ continue
216
+ name = f"{split}-{args.shard:05d}-of-{args.shards:05d}.parquet"
217
+ Dataset.from_list(part).to_parquet(out / name)
218
+ if args.push:
219
+ HfApi().upload_file(
220
+ path_or_fileobj=out / name,
221
+ path_in_repo=f"data/{name}",
222
+ repo_id=args.repo,
223
+ repo_type="dataset",
224
+ commit_message=f"Add {name}",
225
+ )
226
+ stats = {split: len(part) for split, part in rows.items()}
227
+ print(json.dumps({**stats, "cpus": cpus(), "minutes": round((time.time() - start) / 60, 1)}))
228
+
229
+
230
+ def table(source, split):
231
+ """A split from the Hub or from a local data folder."""
232
+ if Path(source).is_dir():
233
+ files = str(Path(source) / f"{split}-*.parquet")
234
+ return load_dataset("parquet", data_files=files, split="train")
235
+ return load_dataset(source, split=split)
236
+
237
+
238
+ def schedule(step, warmup, total):
239
+ """Linear warmup, then cosine decay to zero."""
240
+ if step < warmup:
241
+ return (step + 1) / warmup
242
+ return 0.5 * (1 + math.cos(math.pi * (step - warmup) / max(1, total - warmup)))
243
+
244
+
245
+ def push_result(repo, name, result, revision=None):
246
+ """Upload a result as results/<name>.json in the model repo."""
247
+ HfApi().upload_file(
248
+ path_or_fileobj=json.dumps(result, indent=1).encode(),
249
+ path_in_repo=f"results/{name}.json",
250
+ repo_id=repo,
251
+ revision=revision,
252
+ commit_message=f"Add {name} results",
253
+ )
254
+
255
+
256
+ def train(args):
257
+ """Predict each deal's table from its four hand embeddings, plus per-hand heads."""
258
+ device = "cuda" if torch.cuda.is_available() else "cpu"
259
+ random.seed(args.seed)
260
+ torch.manual_seed(args.seed)
261
+ torch.set_num_threads(cpus())
262
+ cards, dd = (t.to(device) for t in tensors(table(args.data, "train")))
263
+ model = Bridge2Vec(
264
+ dim=args.dim, depth=args.depth, embed_dim=args.embed_dim, hidden=args.hidden
265
+ ).to(device)
266
+ scale = torch.tensor([10.0, 4, 4, 4, 4], device=device)
267
+ optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05)
268
+ steps = args.max_steps
269
+ lr_schedule = torch.optim.lr_scheduler.LambdaLR(
270
+ optimizer, lambda step: schedule(step, min(args.warmup, steps // 10 + 1), steps)
271
+ )
272
+ parameters = sum(p.numel() for p in model.parameters())
273
+ print(json.dumps({"deals": len(cards), "parameters": parameters}), flush=True)
274
+ start, log = time.time(), {}
275
+ model.train()
276
+ for step in range(steps):
277
+ batch = torch.randint(len(cards), (args.batch_size,), device=device)
278
+ hands, tricks = permute_suits(cards[batch], dd[batch])
279
+ with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"):
280
+ _, logits, value, shape = model(hands)
281
+ expected = torch.stack([tricks.roll(-d, 2) for d in range(4)], 1).float()
282
+ table_loss = functional.cross_entropy(logits.float().flatten(0, 2), tricks.flatten())
283
+ value_loss = functional.mse_loss(value.float(), expected)
284
+ shape_loss = functional.mse_loss(shape.float(), profile(hands) / scale)
285
+ loss = table_loss + args.value_weight * value_loss + args.aux_weight * shape_loss
286
+ optimizer.zero_grad(set_to_none=True)
287
+ loss.backward()
288
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
289
+ optimizer.step()
290
+ lr_schedule.step()
291
+ if step % 100 == 0 or step == steps - 1:
292
+ log = {
293
+ "step": step,
294
+ "loss": round(loss.item(), 4),
295
+ "table": round(table_loss.item(), 4),
296
+ "value": round(value_loss.item(), 4),
297
+ "shape": round(shape_loss.item(), 4),
298
+ "exact": round((logits.argmax(-1) == tricks).float().mean().item(), 3),
299
+ "minutes": round((time.time() - start) / 60, 1),
300
+ }
301
+ print(json.dumps(log), flush=True)
302
+
303
+ model.eval()
304
+ model.save_pretrained(args.output)
305
+ (Path(args.output) / "README.md").unlink(missing_ok=True)
306
+ if args.push:
307
+ api = HfApi()
308
+ if args.revision:
309
+ api.create_branch(args.repo, branch=args.revision, exist_ok=True)
310
+ api.upload_folder(
311
+ folder_path=args.output,
312
+ repo_id=args.repo,
313
+ revision=args.revision,
314
+ commit_message="Upload model",
315
+ )
316
+ api.upload_file(
317
+ path_or_fileobj=__file__,
318
+ path_in_repo="bridge2vec.py",
319
+ repo_id=args.repo,
320
+ revision=args.revision,
321
+ )
322
+ push_result(
323
+ args.repo,
324
+ "train",
325
+ {
326
+ **log,
327
+ "data": args.data,
328
+ "deals": len(cards),
329
+ "steps": steps,
330
+ "batch_size": args.batch_size,
331
+ "learning_rate": args.lr,
332
+ "value_weight": args.value_weight,
333
+ "aux_weight": args.aux_weight,
334
+ "dim": args.dim,
335
+ "depth": args.depth,
336
+ "embed_dim": args.embed_dim,
337
+ "hidden": args.hidden,
338
+ "parameters": parameters,
339
+ "runtime_s": round(time.time() - start),
340
+ "device": torch.cuda.get_device_name() if device == "cuda" else "cpu",
341
+ },
342
+ args.revision,
343
+ )
344
+
345
+
346
+ def nearest(score, groups, chunk=1024):
347
+ """For each row, the best-scoring row of another group; score(rows) gives a block."""
348
+ out = []
349
+ for start in range(0, len(groups), chunk):
350
+ rows = slice(start, start + chunk)
351
+ block = score(rows).float()
352
+ block[groups[rows, None] == groups[None]] = -math.inf
353
+ out.append(block.argmax(1))
354
+ return torch.cat(out)
355
+
356
+
357
+ def retrieval(tables, groups, embedding, features):
358
+ """Mean table distance to the nearest neighbour by embedding, HCP and shape, and chance."""
359
+ methods = {
360
+ "embedding": lambda rows: embedding[rows] @ embedding.T,
361
+ "hcp_shape": lambda rows: (torch.rand(len(groups))[None] * 1e-3
362
+ - torch.cdist(features[rows], features, p=1)),
363
+ "random": lambda rows: torch.rand(len(groups[rows]), len(groups)),
364
+ }
365
+ flat = tables.flatten(1).float()
366
+ return {
367
+ name: round((flat[nearest(score, groups)] - flat).abs().mean().item(), 3)
368
+ for name, score in methods.items()
369
+ }
370
+
371
+
372
+ def hard_pairs(embedding, expected, features, margin=0.5):
373
+ """Triplet accuracy among hands the HCP and shape heads cannot tell apart.
374
+
375
+ Hands share a bucket when they have the same suit-length pattern and HCP within a band
376
+ of three. For an anchor and two bucket mates whose expected tables differ from the
377
+ anchor's by more than the margin, the embedding must rank the closer table first.
378
+ Chance is 0.5, and so is anything that sees only HCP and shape.
379
+ """
380
+ keys = [(int(f[0]) // 3, *sorted(f[1:].int().tolist())) for f in features]
381
+ flat = expected.flatten(1)
382
+ right = total = 0
383
+ for key in set(keys):
384
+ mates = torch.tensor([i for i, k in enumerate(keys) if k == key])
385
+ if len(mates) < 3:
386
+ continue
387
+ near = torch.cdist(flat[mates], flat[mates], p=1) / flat.shape[1]
388
+ close = 1 - embedding[mates] @ embedding[mates].T
389
+ gap = near[:, :, None] - near[:, None, :]
390
+ same = torch.eye(len(mates), dtype=torch.bool)
391
+ different = (gap.abs() > margin) & ~same[:, :, None] & ~same[:, None, :]
392
+ agree = (close[:, :, None] - close[:, None, :]) * gap > 0
393
+ right += (agree & different).sum().item() / 2
394
+ total += different.sum().item() / 2
395
+ return {"triplets": int(total), "accuracy": round(right / max(total, 1), 3)}
396
+
397
+
398
+ def hand_map(embedding, expected, path):
399
+ """UMAP of the test hands' embeddings, coloured by expected notrump tricks."""
400
+ reducer = umap.UMAP(metric="cosine", n_neighbors=min(15, len(embedding) - 1), random_state=0)
401
+ xy = reducer.fit_transform(embedding.numpy())
402
+ fig, ax = plt.subplots(figsize=(8, 7))
403
+ points = ax.scatter(*xy.T, c=expected, cmap="viridis", s=4, linewidths=0)
404
+ fig.colorbar(points, ax=ax, label="Expected notrump tricks with North declaring", shrink=0.7)
405
+ ax.set_axis_off()
406
+ fig.savefig(path, dpi=150, bbox_inches="tight")
407
+ plt.close(fig)
408
+
409
+
410
+ def evaluate(args):
411
+ """Score predicted tables and expected tables; retrieve hands and deals that play alike."""
412
+ device = "cuda" if torch.cuda.is_available() else "cpu"
413
+ model = Bridge2Vec.from_pretrained(args.model, revision=args.revision).to(device).eval()
414
+ test = table(args.data, "test")
415
+ norths = [d.removeprefix("N:").split()[0] for d in test["deal"]]
416
+ _, groups = np.unique(norths, return_inverse=True)
417
+ if args.limit:
418
+ test = test.select(np.flatnonzero(groups < args.limit))
419
+ groups = groups[groups < args.limit]
420
+ groups = torch.tensor(groups)
421
+ cards, dd = tensors(test)
422
+ out = model.run(cards)
423
+ error = (out["tricks"] - dd).abs()
424
+ first = torch.tensor(np.unique(groups.numpy(), return_index=True)[1])
425
+ hands = len(first)
426
+ expected = torch.zeros(hands, 5, 4).index_add_(0, groups, dd.float())
427
+ expected /= torch.bincount(groups, minlength=hands)[:, None, None]
428
+ value = out["value"][first, 0]
429
+ result = {
430
+ "model": args.model,
431
+ "tables": {
432
+ "deals": len(dd),
433
+ "mae": round(error.float().mean().item(), 3),
434
+ "exact": round((error == 0).float().mean().item(), 3),
435
+ "within_one": round((error <= 1).float().mean().item(), 3),
436
+ "table_exact": round((error == 0).flatten(1).all(1).float().mean().item(), 3),
437
+ "mae_by_strain": {
438
+ s: round(error[:, i].float().mean().item(), 3) for i, s in enumerate(STRAINS)
439
+ },
440
+ },
441
+ "hands": {
442
+ "hands": hands,
443
+ "deals_per_hand": round(len(dd) / hands, 1),
444
+ "expected_mae": round((value - expected).abs().mean().item(), 3),
445
+ "constant_mae": round((expected.mean(0) - expected).abs().mean().item(), 3),
446
+ "retrieval": retrieval(
447
+ expected, torch.arange(hands), out["hands"][first, 0], profile(cards[first, 0])
448
+ ),
449
+ "hard_pairs": hard_pairs(
450
+ out["hands"][first, 0], expected, profile(cards[first, 0])
451
+ ),
452
+ },
453
+ "deal_retrieval": retrieval(
454
+ dd, groups, functional.normalize(out["hands"].flatten(1), dim=-1),
455
+ profile(cards).flatten(1),
456
+ ),
457
+ }
458
+ print(json.dumps(result, indent=1))
459
+ Path(args.output).mkdir(parents=True, exist_ok=True)
460
+ hand_map(out["hands"][first, 0], expected[:, 4, 0], Path(args.output) / "map.png")
461
+ if args.push:
462
+ push_result(args.repo, "eval", result, args.revision)
463
+ HfApi().upload_file(
464
+ path_or_fileobj=Path(args.output) / "map.png",
465
+ path_in_repo="results/map.png",
466
+ repo_id=args.repo,
467
+ revision=args.revision,
468
+ )
469
+
470
+
471
+ def strains(tricks):
472
+ """A table (5, 4) as {strain: {seat: tricks}}."""
473
+ return {s: dict(zip(SEATS, row)) for s, row in zip(STRAINS, tricks.tolist())}
474
+
475
+
476
+ def embed(args):
477
+ """Print a hand's embedding, expected tricks and look-alikes, or a deal's table."""
478
+ model = Bridge2Vec.from_pretrained(args.model).eval()
479
+ if args.deal:
480
+ cards = torch.tensor([parse_deal(args.deal)])
481
+ out = model.run(cards)
482
+ truth = calc_all_tables([Deal("N:" + args.deal.removeprefix("N:"))])[0].to_list()
483
+ result = {
484
+ "predicted": strains(out["tricks"][0]),
485
+ "double_dummy": strains(torch.tensor(truth)),
486
+ "embedding": [round(x, 4) for x in out["hands"][0].flatten().tolist()],
487
+ }
488
+ else:
489
+ z, value = model.embed([args.hand])
490
+ deals = table(args.data, "test")["deal"]
491
+ gallery = sorted({d.removeprefix("N:").split()[0] for d in deals})
492
+ similarity = z @ model.embed(gallery)[0].T
493
+ top = similarity[0].topk(min(args.top, similarity.shape[1]))
494
+ hcp, *lengths = profile(torch.tensor(parse_hand(args.hand))).int().tolist()
495
+ result = {
496
+ "hcp": hcp,
497
+ "lengths": dict(zip(STRAINS, lengths)),
498
+ "expected_tricks": {
499
+ who: {s: round(t, 1) for s, t in zip(STRAINS, value[0, :, j].tolist())}
500
+ for j, who in ((0, "this hand declares"), (2, "partner declares"))
501
+ },
502
+ "nearest": [
503
+ {"hand": gallery[i], "cosine": round(s, 3)}
504
+ for s, i in zip(top.values.tolist(), top.indices.tolist())
505
+ ],
506
+ "embedding": [round(x, 4) for x in z[0].tolist()],
507
+ }
508
+ print(json.dumps(result, indent=1))
509
+
510
+
511
+ def card(args):
512
+ """Render card.jinja into card/README.md with the results stored in the model repo."""
513
+ try:
514
+ folder = Path(snapshot_download(args.repo, allow_patterns="results/*.json"))
515
+ paths = folder.glob("results/*.json")
516
+ except RepositoryNotFoundError:
517
+ paths = []
518
+ results = {path.stem: json.loads(path.read_text(encoding="utf-8")) for path in paths}
519
+ meta = ModelCardData(
520
+ model_name=args.repo.split("/")[1],
521
+ datasets=[DATA],
522
+ license="mit",
523
+ library_name="pytorch",
524
+ pipeline_tag="feature-extraction",
525
+ tags=["contract-bridge", "double-dummy", "embeddings", "weird2vec"],
526
+ )
527
+ rendered = ModelCard.from_template(
528
+ meta,
529
+ template_path=HERE / "card.jinja",
530
+ repo=args.repo,
531
+ data=DATA,
532
+ train=results.get("train"),
533
+ eval=results.get("eval"),
534
+ )
535
+ (HERE / "card").mkdir(exist_ok=True)
536
+ rendered.save(HERE / "card" / "README.md")
537
+
538
+
539
+ def main():
540
+ """Parse arguments and run a command."""
541
+ parser = argparse.ArgumentParser(description=__doc__)
542
+ commands = parser.add_subparsers(dest="command", required=True)
543
+
544
+ data_parser = commands.add_parser("data")
545
+ data_parser.add_argument("--deals", type=int, default=200_000, help="train deals")
546
+ data_parser.add_argument("--test-hands", type=int, default=1000, help=f"{CHUNK} deals each")
547
+ data_parser.add_argument("--shard", type=int, default=0)
548
+ data_parser.add_argument("--shards", type=int, default=1)
549
+ data_parser.add_argument("--seed", type=int, default=0)
550
+ data_parser.add_argument("--output", default="out/data")
551
+ data_parser.add_argument("--repo", default=DATA)
552
+ data_parser.add_argument("--push", action="store_true")
553
+ data_parser.set_defaults(run=data)
554
+
555
+ train_parser = commands.add_parser("train")
556
+ train_parser.add_argument("--data", default=DATA, help="dataset repo or local data folder")
557
+ train_parser.add_argument("--max-steps", type=int, default=50_000)
558
+ train_parser.add_argument("--batch-size", type=int, default=1024, help="deals per step")
559
+ train_parser.add_argument("--lr", type=float, default=3e-4)
560
+ train_parser.add_argument("--warmup", type=int, default=1000)
561
+ train_parser.add_argument("--value-weight", type=float, default=0.1)
562
+ train_parser.add_argument("--aux-weight", type=float, default=0.1)
563
+ train_parser.add_argument("--dim", type=int, default=256)
564
+ train_parser.add_argument("--depth", type=int, default=4)
565
+ train_parser.add_argument("--embed-dim", type=int, default=128)
566
+ train_parser.add_argument("--hidden", type=int, default=1024)
567
+ train_parser.add_argument("--seed", type=int, default=0)
568
+ train_parser.add_argument("--output", default="out/model")
569
+ train_parser.add_argument("--repo", default=REPO)
570
+ train_parser.add_argument("--revision", default=None, help="branch to push to")
571
+ train_parser.add_argument("--push", action="store_true")
572
+ train_parser.set_defaults(run=train)
573
+
574
+ eval_parser = commands.add_parser("eval")
575
+ eval_parser.add_argument("--model", default=REPO, help="model repo or local folder")
576
+ eval_parser.add_argument("--data", default=DATA)
577
+ eval_parser.add_argument("--limit", type=int, default=None, help="number of test hands")
578
+ eval_parser.add_argument("--repo", default=REPO, help="where --push stores the result")
579
+ eval_parser.add_argument("--revision", default=None, help="model branch")
580
+ eval_parser.add_argument("--output", default="out/eval", help="where the hand map goes")
581
+ eval_parser.add_argument("--push", action="store_true")
582
+ eval_parser.set_defaults(run=evaluate)
583
+
584
+ embed_parser = commands.add_parser("embed")
585
+ target = embed_parser.add_mutually_exclusive_group(required=True)
586
+ target.add_argument("--hand", help="PBN hand, spades first, e.g. AKQ32.KJ4.T9.A87")
587
+ target.add_argument("--deal", help="PBN deal, North first, e.g. N:AKQ32.KJ4.T9.A87 ...")
588
+ embed_parser.add_argument("--model", default=REPO)
589
+ embed_parser.add_argument("--data", default=DATA, help="test hands to search")
590
+ embed_parser.add_argument("--top", type=int, default=5)
591
+ embed_parser.set_defaults(run=embed)
592
+
593
+ card_parser = commands.add_parser("card")
594
+ card_parser.add_argument("--repo", default=REPO)
595
+ card_parser.set_defaults(run=card)
596
+
597
+ args = parser.parse_args()
598
+ if not PROGRESS:
599
+ datasets.disable_progress_bars()
600
+ huggingface_hub.utils.disable_progress_bars()
601
+ args.run(args)
602
+
603
+
604
+ if __name__ == "__main__":
605
+ main()
config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "depth": 4,
3
+ "dim": 256,
4
+ "embed_dim": 128,
5
+ "heads": 8,
6
+ "hidden": 1024
7
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e7d5c6b469fe2e36cc221144fad60a221da1ca3e3f0d21eab41068607180cb90
3
+ size 19656236
results/eval.json ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model": "jgalego/bridge2vec",
3
+ "tables": {
4
+ "deals": 32000,
5
+ "mae": 0.366,
6
+ "exact": 0.654,
7
+ "within_one": 0.984,
8
+ "table_exact": 0.016,
9
+ "mae_by_strain": {
10
+ "S": 0.34,
11
+ "H": 0.343,
12
+ "D": 0.344,
13
+ "C": 0.342,
14
+ "N": 0.46
15
+ }
16
+ },
17
+ "hands": {
18
+ "hands": 1000,
19
+ "deals_per_hand": 32.0,
20
+ "expected_mae": 0.345,
21
+ "constant_mae": 1.251,
22
+ "retrieval": {
23
+ "embedding": 0.836,
24
+ "hcp_shape": 0.613,
25
+ "random": 1.752
26
+ },
27
+ "hard_pairs": {
28
+ "triplets": 98216,
29
+ "accuracy": 0.769
30
+ }
31
+ },
32
+ "deal_retrieval": {
33
+ "embedding": 1.788,
34
+ "hcp_shape": 1.411,
35
+ "random": 3.131
36
+ }
37
+ }
results/map.png ADDED

Git LFS Details

  • SHA256: b18752e9725632ed675c4c7fcc60ffccef7621a0eeff481b5351641208f5d98b
  • Pointer size: 131 Bytes
  • Size of remote file: 119 kB
results/train.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "step": 49999,
3
+ "loss": 1.2938,
4
+ "table": 0.7635,
5
+ "value": 5.3025,
6
+ "shape": 0.0006,
7
+ "exact": 0.667,
8
+ "minutes": 80.6,
9
+ "data": "jgalego/bridge2vec-deals",
10
+ "deals": 200000,
11
+ "steps": 50000,
12
+ "batch_size": 1024,
13
+ "learning_rate": 0.0003,
14
+ "value_weight": 0.1,
15
+ "aux_weight": 0.1,
16
+ "dim": 256,
17
+ "depth": 4,
18
+ "embed_dim": 128,
19
+ "hidden": 1024,
20
+ "parameters": 4912479,
21
+ "runtime_s": 4841,
22
+ "device": "NVIDIA A10G"
23
+ }