Download bridge2vec.py from jgalego/bridge2vec: direct link, hf CLI and curl.
- Browser
- Download file 26.6 kB
-
https://huggingface.co/jgalego/bridge2vec/resolve/main/bridge2vec.py
- Command line
-
hf download hf://jgalego/bridge2vec/bridge2vec.py
-
curl -L -o bridge2vec.py https://huggingface.co/jgalego/bridge2vec/resolve/main/bridge2vec.py
26.6 kB
| # /// script | |
| # requires-python = ">=3.12" | |
| # dependencies = [ | |
| # "datasets", | |
| # "endplay", | |
| # "huggingface-hub", | |
| # "jinja2", | |
| # "matplotlib", | |
| # "numpy", | |
| # "safetensors", | |
| # "torch", | |
| # "umap-learn", | |
| # ] | |
| # /// | |
| """Bridge2Vec: embed bridge hands by how they take tricks, learned from double-dummy tables.""" | |
| # pylint: disable=too-many-arguments,too-many-instance-attributes,too-many-locals | |
| # pylint: disable=too-many-positional-arguments,too-many-statements | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import random | |
| import sys | |
| import time | |
| from multiprocessing import Pool | |
| from pathlib import Path | |
| import datasets | |
| import huggingface_hub.utils | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import torch | |
| import umap | |
| from datasets import Dataset, load_dataset | |
| from endplay._dds import SetMaxThreads | |
| from endplay.dds import calc_all_tables | |
| from endplay.types import Deal | |
| from huggingface_hub import ( | |
| HfApi, | |
| ModelCard, | |
| ModelCardData, | |
| PyTorchModelHubMixin, | |
| snapshot_download, | |
| ) | |
| from huggingface_hub.errors import RepositoryNotFoundError | |
| from torch import nn | |
| from torch.nn import functional | |
| REPO = "jgalego/bridge2vec" | |
| DATA = "jgalego/bridge2vec-deals" | |
| HERE = Path(__file__).parent | |
| # Progress bars redraw in place, which shows up as garbage in HF Jobs logs. | |
| PROGRESS = sys.stderr.isatty() | |
| RANKS = "23456789TJQKA" | |
| SEATS = "NESW" | |
| # Suits in PBN order, then notrump: the rows of a double-dummy table. | |
| STRAINS = "SHDCN" | |
| # DDS solves at most 32 tables per call; a test hand gets one call's worth of deals. | |
| CHUNK = 32 | |
| def cpus(): | |
| """Return the CPU quota visible to this process.""" | |
| try: | |
| quota, period = Path("/sys/fs/cgroup/cpu.max").read_text(encoding="utf-8").split() | |
| if quota != "max": | |
| return max(1, int(quota) // int(period)) | |
| except OSError: | |
| pass | |
| return len(os.sched_getaffinity(0)) | |
| def hand_pbn(cards): | |
| """Write card ids (13 * suit + rank) as a PBN hand, e.g. AKQ32.KJ4.T9.A87.""" | |
| suits = [sorted((c % 13 for c in cards if c // 13 == s), reverse=True) for s in range(4)] | |
| return ".".join("".join(RANKS[r] for r in suit) for suit in suits) | |
| def parse_hand(text): | |
| """Card ids of a PBN hand.""" | |
| suits = text.split(".") | |
| if len(suits) != 4: | |
| raise ValueError(f"{text!r} needs four suits separated by dots") | |
| cards = [13 * s + RANKS.index(r) for s, ranks in enumerate(suits) for r in ranks.upper()] | |
| if len(set(cards)) != 13: | |
| raise ValueError(f"{text!r} is not 13 different cards") | |
| return cards | |
| def parse_deal(text): | |
| """Card ids of a PBN deal with North first, shape (4, 13).""" | |
| hands = [parse_hand(h) for h in text.removeprefix("N:").split()] | |
| if len(hands) != 4 or len({c for h in hands for c in h}) != 52: | |
| raise ValueError(f"{text!r} is not four hands of one deck") | |
| return hands | |
| def tensors(rows): | |
| """Cards (n, 4, 13) and double-dummy tables (n, 5, 4) of a split.""" | |
| cards = torch.tensor([parse_deal(d) for d in rows["deal"]]) | |
| return cards, torch.tensor(np.array(rows["dd"])) | |
| def profile(cards): | |
| """HCP and suit lengths of hands, shape (..., 5).""" | |
| hcp = (cards % 13 - 8).clamp(min=0).sum(-1, keepdim=True) | |
| return torch.cat([hcp, functional.one_hot(cards // 13, 4).sum(-2)], -1).float() | |
| def permute_suits(cards, dd): | |
| """Relabel the suits of each deal at random, moving the table rows with them.""" | |
| perm = torch.rand(len(cards), 4, device=cards.device).argsort(1) | |
| cards = perm.gather(1, (cards // 13).flatten(1)).view_as(cards) * 13 + cards % 13 | |
| rows = perm.argsort(1)[:, :, None].expand(-1, -1, 4) | |
| return cards, torch.cat([dd[:, :4].gather(1, rows), dd[:, 4:]], 1) | |
| class Bridge2Vec(nn.Module, PyTorchModelHubMixin): | |
| """Transformer over the 13 cards of a hand; an MLP on four hands predicts the table.""" | |
| def __init__(self, dim=256, depth=4, heads=8, embed_dim=128, hidden=1024): | |
| super().__init__() | |
| self.suits = nn.Embedding(4, dim) | |
| self.ranks = nn.Embedding(13, dim) | |
| layer = nn.TransformerEncoderLayer( | |
| dim, heads, 4 * dim, dropout=0.0, batch_first=True, norm_first=True | |
| ) | |
| self.encoder = nn.TransformerEncoder(layer, depth, enable_nested_tensor=False) | |
| self.norm = nn.LayerNorm(dim) | |
| self.project = nn.Sequential(nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, embed_dim)) | |
| self.table = nn.Sequential( | |
| nn.Linear(4 * embed_dim, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, 5 * 14), | |
| ) | |
| self.value = nn.Linear(embed_dim, 5 * 4) | |
| nn.init.constant_(self.value.bias, 6.5) | |
| self.shape = nn.Linear(embed_dim, 5) | |
| def encode(self, cards): | |
| """Unit-length embeddings of hands, cards shape (n, 13).""" | |
| x = self.suits(cards // 13) + self.ranks(cards % 13) | |
| pooled = self.norm(self.encoder(x)).mean(1) | |
| return functional.normalize(self.project(pooled), dim=-1) | |
| def forward(self, cards): | |
| """Hand embeddings, table logits, expected tables and profiles of deals (n, 4, 13). | |
| Row d of the table comes from the hands in the order declarer d, LHO, partner, RHO, | |
| so rotating the seats rotates the table. Expected tables are per hand, with | |
| declarers relative to it: itself, LHO, partner, RHO. | |
| """ | |
| n = len(cards) | |
| z = self.encode(cards.flatten(0, 1)).view(n, 4, -1) | |
| views = torch.stack([z.roll(-d, 1).flatten(1) for d in range(4)], 1) | |
| logits = self.table(views).view(n, 4, 5, 14).transpose(1, 2) | |
| return z, logits, self.value(z).view(n, 4, 5, 4), self.shape(z) | |
| def run(self, cards, batch_size=4096): | |
| """Hand embeddings, double-dummy tables and expected tables for deals (n, 4, 13).""" | |
| device = next(self.parameters()).device | |
| starts = range(0, len(cards), batch_size) | |
| parts = [self(cards[i : i + batch_size].to(device)) for i in starts] | |
| z, logits, value, _ = (torch.cat(p).float().cpu() for p in zip(*parts)) | |
| return {"hands": z, "tricks": logits.argmax(-1), "value": value} | |
| def embed(self, hands): | |
| """Embeddings and expected tables for PBN hands.""" | |
| device = next(self.parameters()).device | |
| z = self.encode(torch.tensor([parse_hand(h) for h in hands], device=device)) | |
| return z.float().cpu(), self.value(z).view(-1, 5, 4).float().cpu() | |
| def solve(job): | |
| """32 deals with their double-dummy tables. Test deals share North. Seeded per chunk.""" | |
| split, index, seed = job | |
| rng = random.Random(f"{seed}:{split}:{index}") | |
| north = rng.sample(range(52), 13) if split == "test" else [] | |
| deals = [] | |
| for _ in range(CHUNK): | |
| rest = [c for c in range(52) if c not in north] | |
| rng.shuffle(rest) | |
| cards = north + rest | |
| deals.append("N:" + " ".join(hand_pbn(cards[i : i + 13]) for i in range(0, 52, 13))) | |
| tables = calc_all_tables([Deal(d) for d in deals]) | |
| return [{"deal": d, "dd": t.to_list()} for d, t in zip(deals, tables)] | |
| def data(args): | |
| """Deal random hands, solve them double dummy; save as parquet, optionally push.""" | |
| start = time.time() | |
| jobs = [("test", i) for i in range(args.test_hands)] | |
| jobs += [("train", i) for i in range(args.deals // CHUNK)] | |
| jobs = jobs[args.shard :: args.shards] | |
| rows = {"train": [], "test": []} | |
| step = max(1, len(jobs) // 20) | |
| with Pool(cpus(), initializer=SetMaxThreads, initargs=(1,)) as pool: | |
| tasks = [(split, i, args.seed) for split, i in jobs] | |
| for n, ((split, _), part) in enumerate(zip(jobs, pool.imap(solve, tasks)), 1): | |
| rows[split] += part | |
| if n % step == 0 or n == len(jobs): | |
| rate = n * CHUNK / (time.time() - start) | |
| print(json.dumps({"chunks": n, "of": len(jobs), "deals_per_s": round(rate, 1)}), | |
| flush=True) | |
| out = Path(args.output) | |
| out.mkdir(parents=True, exist_ok=True) | |
| for split, part in rows.items(): | |
| if not part: | |
| continue | |
| name = f"{split}-{args.shard:05d}-of-{args.shards:05d}.parquet" | |
| Dataset.from_list(part).to_parquet(out / name) | |
| if args.push: | |
| HfApi().upload_file( | |
| path_or_fileobj=out / name, | |
| path_in_repo=f"data/{name}", | |
| repo_id=args.repo, | |
| repo_type="dataset", | |
| commit_message=f"Add {name}", | |
| ) | |
| stats = {split: len(part) for split, part in rows.items()} | |
| print(json.dumps({**stats, "cpus": cpus(), "minutes": round((time.time() - start) / 60, 1)})) | |
| def table(source, split): | |
| """A split from the Hub or from a local data folder.""" | |
| if Path(source).is_dir(): | |
| files = str(Path(source) / f"{split}-*.parquet") | |
| return load_dataset("parquet", data_files=files, split="train") | |
| return load_dataset(source, split=split) | |
| def schedule(step, warmup, total): | |
| """Linear warmup, then cosine decay to zero.""" | |
| if step < warmup: | |
| return (step + 1) / warmup | |
| return 0.5 * (1 + math.cos(math.pi * (step - warmup) / max(1, total - warmup))) | |
| def metric_loss(z, value, temperature, target_temperature): | |
| """Make hands that play alike close: soft targets from the distance between expected tables.""" | |
| z = z.float().flatten(0, 1) | |
| value = value.detach().float().flatten(0, 1).flatten(1) | |
| self_pair = torch.eye(len(z), dtype=torch.bool, device=z.device) | |
| distance = torch.cdist(value, value, p=1) / value.shape[1] | |
| target = (-distance / target_temperature).masked_fill(self_pair, -1e9).softmax(1) | |
| logits = (z @ z.T / temperature).masked_fill(self_pair, -1e9) | |
| return -(target * logits.log_softmax(1)).sum(1).mean() | |
| def push_result(repo, name, result, revision=None): | |
| """Upload a result as results/<name>.json in the model repo.""" | |
| HfApi().upload_file( | |
| path_or_fileobj=json.dumps(result, indent=1).encode(), | |
| path_in_repo=f"results/{name}.json", | |
| repo_id=repo, | |
| revision=revision, | |
| commit_message=f"Add {name} results", | |
| ) | |
| def train(args): | |
| """Predict each deal's table from its four hand embeddings, plus per-hand heads.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| random.seed(args.seed) | |
| torch.manual_seed(args.seed) | |
| torch.set_num_threads(cpus()) | |
| cards, dd = (t.to(device) for t in tensors(table(args.data, "train"))) | |
| model = Bridge2Vec( | |
| dim=args.dim, depth=args.depth, embed_dim=args.embed_dim, hidden=args.hidden | |
| ).to(device) | |
| scale = torch.tensor([10.0, 4, 4, 4, 4], device=device) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05) | |
| steps = args.max_steps | |
| lr_schedule = torch.optim.lr_scheduler.LambdaLR( | |
| optimizer, lambda step: schedule(step, min(args.warmup, steps // 10 + 1), steps) | |
| ) | |
| parameters = sum(p.numel() for p in model.parameters()) | |
| print(json.dumps({"deals": len(cards), "parameters": parameters}), flush=True) | |
| start, log = time.time(), {} | |
| model.train() | |
| for step in range(steps): | |
| batch = torch.randint(len(cards), (args.batch_size,), device=device) | |
| hands, tricks = permute_suits(cards[batch], dd[batch]) | |
| with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): | |
| z, logits, value, shape = model(hands) | |
| expected = torch.stack([tricks.roll(-d, 2) for d in range(4)], 1).float() | |
| table_loss = functional.cross_entropy(logits.float().flatten(0, 2), tricks.flatten()) | |
| value_loss = functional.mse_loss(value.float(), expected) | |
| shape_loss = functional.mse_loss(shape.float(), profile(hands) / scale) | |
| metric = metric_loss(z, value, args.temperature, args.target_temperature) | |
| loss = (table_loss + args.value_weight * value_loss + args.aux_weight * shape_loss | |
| + args.metric_weight * metric) | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| lr_schedule.step() | |
| if step % 100 == 0 or step == steps - 1: | |
| log = { | |
| "step": step, | |
| "loss": round(loss.item(), 4), | |
| "table": round(table_loss.item(), 4), | |
| "value": round(value_loss.item(), 4), | |
| "shape": round(shape_loss.item(), 4), | |
| "metric": round(metric.item(), 4), | |
| "exact": round((logits.argmax(-1) == tricks).float().mean().item(), 3), | |
| "minutes": round((time.time() - start) / 60, 1), | |
| } | |
| print(json.dumps(log), flush=True) | |
| model.eval() | |
| model.save_pretrained(args.output) | |
| (Path(args.output) / "README.md").unlink(missing_ok=True) | |
| if args.push: | |
| api = HfApi() | |
| if args.revision: | |
| api.create_branch(args.repo, branch=args.revision, exist_ok=True) | |
| api.upload_folder( | |
| folder_path=args.output, | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| commit_message="Upload model", | |
| ) | |
| api.upload_file( | |
| path_or_fileobj=__file__, | |
| path_in_repo="bridge2vec.py", | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| ) | |
| push_result( | |
| args.repo, | |
| "train", | |
| { | |
| **log, | |
| "data": args.data, | |
| "deals": len(cards), | |
| "steps": steps, | |
| "batch_size": args.batch_size, | |
| "learning_rate": args.lr, | |
| "value_weight": args.value_weight, | |
| "aux_weight": args.aux_weight, | |
| "metric_weight": args.metric_weight, | |
| "temperature": args.temperature, | |
| "target_temperature": args.target_temperature, | |
| "dim": args.dim, | |
| "depth": args.depth, | |
| "embed_dim": args.embed_dim, | |
| "hidden": args.hidden, | |
| "parameters": parameters, | |
| "runtime_s": round(time.time() - start), | |
| "device": torch.cuda.get_device_name() if device == "cuda" else "cpu", | |
| }, | |
| args.revision, | |
| ) | |
| def nearest(score, groups, chunk=1024): | |
| """For each row, the best-scoring row of another group; score(rows) gives a block.""" | |
| out = [] | |
| for start in range(0, len(groups), chunk): | |
| rows = slice(start, start + chunk) | |
| block = score(rows).float() | |
| block[groups[rows, None] == groups[None]] = -math.inf | |
| out.append(block.argmax(1)) | |
| return torch.cat(out) | |
| def retrieval(tables, groups, embedding, features): | |
| """Mean table distance to the nearest neighbour by embedding, HCP and shape, and chance.""" | |
| methods = { | |
| "embedding": lambda rows: embedding[rows] @ embedding.T, | |
| "hcp_shape": lambda rows: (torch.rand(len(groups))[None] * 1e-3 | |
| - torch.cdist(features[rows], features, p=1)), | |
| "random": lambda rows: torch.rand(len(groups[rows]), len(groups)), | |
| } | |
| flat = tables.flatten(1).float() | |
| return { | |
| name: round((flat[nearest(score, groups)] - flat).abs().mean().item(), 3) | |
| for name, score in methods.items() | |
| } | |
| def hard_pairs(embedding, expected, features, margin=0.5): | |
| """Triplet accuracy among hands the HCP and shape heads cannot tell apart. | |
| Hands share a bucket when they have the same suit-length pattern and HCP within a band | |
| of three. For an anchor and two bucket mates whose expected tables differ from the | |
| anchor's by more than the margin, the embedding must rank the closer table first. | |
| Chance is 0.5, and so is anything that sees only HCP and shape. | |
| """ | |
| keys = [(int(f[0]) // 3, *sorted(f[1:].int().tolist())) for f in features] | |
| flat = expected.flatten(1) | |
| right = total = 0 | |
| for key in set(keys): | |
| mates = torch.tensor([i for i, k in enumerate(keys) if k == key]) | |
| if len(mates) < 3: | |
| continue | |
| near = torch.cdist(flat[mates], flat[mates], p=1) / flat.shape[1] | |
| close = 1 - embedding[mates] @ embedding[mates].T | |
| gap = near[:, :, None] - near[:, None, :] | |
| same = torch.eye(len(mates), dtype=torch.bool) | |
| different = (gap.abs() > margin) & ~same[:, :, None] & ~same[:, None, :] | |
| agree = (close[:, :, None] - close[:, None, :]) * gap > 0 | |
| right += (agree & different).sum().item() / 2 | |
| total += different.sum().item() / 2 | |
| return {"triplets": int(total), "accuracy": round(right / max(total, 1), 3)} | |
| def hand_map(embedding, expected, path): | |
| """UMAP of the test hands' embeddings, coloured by expected notrump tricks.""" | |
| reducer = umap.UMAP(metric="cosine", n_neighbors=min(15, len(embedding) - 1), random_state=0) | |
| xy = reducer.fit_transform(embedding.numpy()) | |
| fig, ax = plt.subplots(figsize=(8, 7)) | |
| points = ax.scatter(*xy.T, c=expected, cmap="viridis", s=4, linewidths=0) | |
| fig.colorbar(points, ax=ax, label="Expected notrump tricks with North declaring", shrink=0.7) | |
| ax.set_axis_off() | |
| fig.savefig(path, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| def evaluate(args): | |
| """Score predicted tables and expected tables; retrieve hands and deals that play alike.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = Bridge2Vec.from_pretrained(args.model, revision=args.revision).to(device).eval() | |
| test = table(args.data, "test") | |
| norths = [d.removeprefix("N:").split()[0] for d in test["deal"]] | |
| _, groups = np.unique(norths, return_inverse=True) | |
| if args.limit: | |
| test = test.select(np.flatnonzero(groups < args.limit)) | |
| groups = groups[groups < args.limit] | |
| groups = torch.tensor(groups) | |
| cards, dd = tensors(test) | |
| out = model.run(cards) | |
| error = (out["tricks"] - dd).abs() | |
| first = torch.tensor(np.unique(groups.numpy(), return_index=True)[1]) | |
| hands = len(first) | |
| expected = torch.zeros(hands, 5, 4).index_add_(0, groups, dd.float()) | |
| expected /= torch.bincount(groups, minlength=hands)[:, None, None] | |
| value = out["value"][first, 0] | |
| result = { | |
| "model": args.model, | |
| "tables": { | |
| "deals": len(dd), | |
| "mae": round(error.float().mean().item(), 3), | |
| "exact": round((error == 0).float().mean().item(), 3), | |
| "within_one": round((error <= 1).float().mean().item(), 3), | |
| "table_exact": round((error == 0).flatten(1).all(1).float().mean().item(), 3), | |
| "mae_by_strain": { | |
| s: round(error[:, i].float().mean().item(), 3) for i, s in enumerate(STRAINS) | |
| }, | |
| }, | |
| "hands": { | |
| "hands": hands, | |
| "deals_per_hand": round(len(dd) / hands, 1), | |
| "expected_mae": round((value - expected).abs().mean().item(), 3), | |
| "constant_mae": round((expected.mean(0) - expected).abs().mean().item(), 3), | |
| "retrieval": retrieval( | |
| expected, torch.arange(hands), out["hands"][first, 0], profile(cards[first, 0]) | |
| ), | |
| "hard_pairs": hard_pairs( | |
| out["hands"][first, 0], expected, profile(cards[first, 0]) | |
| ), | |
| }, | |
| "deal_retrieval": retrieval( | |
| dd, groups, functional.normalize(out["hands"].flatten(1), dim=-1), | |
| profile(cards).flatten(1), | |
| ), | |
| } | |
| print(json.dumps(result, indent=1)) | |
| Path(args.output).mkdir(parents=True, exist_ok=True) | |
| hand_map(out["hands"][first, 0], expected[:, 4, 0], Path(args.output) / "map.png") | |
| if args.push: | |
| push_result(args.repo, "eval", result, args.revision) | |
| HfApi().upload_file( | |
| path_or_fileobj=Path(args.output) / "map.png", | |
| path_in_repo="results/map.png", | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| ) | |
| def strains(tricks): | |
| """A table (5, 4) as {strain: {seat: tricks}}.""" | |
| return {s: dict(zip(SEATS, row)) for s, row in zip(STRAINS, tricks.tolist())} | |
| def embed(args): | |
| """Print a hand's embedding, expected tricks and look-alikes, or a deal's table.""" | |
| model = Bridge2Vec.from_pretrained(args.model).eval() | |
| if args.deal: | |
| cards = torch.tensor([parse_deal(args.deal)]) | |
| out = model.run(cards) | |
| truth = calc_all_tables([Deal("N:" + args.deal.removeprefix("N:"))])[0].to_list() | |
| result = { | |
| "predicted": strains(out["tricks"][0]), | |
| "double_dummy": strains(torch.tensor(truth)), | |
| "embedding": [round(x, 4) for x in out["hands"][0].flatten().tolist()], | |
| } | |
| else: | |
| z, value = model.embed([args.hand]) | |
| deals = table(args.data, "test")["deal"] | |
| gallery = sorted({d.removeprefix("N:").split()[0] for d in deals}) | |
| similarity = z @ model.embed(gallery)[0].T | |
| top = similarity[0].topk(min(args.top, similarity.shape[1])) | |
| hcp, *lengths = profile(torch.tensor(parse_hand(args.hand))).int().tolist() | |
| result = { | |
| "hcp": hcp, | |
| "lengths": dict(zip(STRAINS, lengths)), | |
| "expected_tricks": { | |
| who: {s: round(t, 1) for s, t in zip(STRAINS, value[0, :, j].tolist())} | |
| for j, who in ((0, "this hand declares"), (2, "partner declares")) | |
| }, | |
| "nearest": [ | |
| {"hand": gallery[i], "cosine": round(s, 3)} | |
| for s, i in zip(top.values.tolist(), top.indices.tolist()) | |
| ], | |
| "embedding": [round(x, 4) for x in z[0].tolist()], | |
| } | |
| print(json.dumps(result, indent=1)) | |
| def card(args): | |
| """Render card.jinja into card/README.md with the results stored in the model repo.""" | |
| try: | |
| folder = Path(snapshot_download(args.repo, allow_patterns="results/*.json")) | |
| paths = folder.glob("results/*.json") | |
| except RepositoryNotFoundError: | |
| paths = [] | |
| results = {path.stem: json.loads(path.read_text(encoding="utf-8")) for path in paths} | |
| meta = ModelCardData( | |
| model_name=args.repo.split("/")[1], | |
| datasets=[DATA], | |
| license="mit", | |
| library_name="pytorch", | |
| pipeline_tag="feature-extraction", | |
| tags=["contract-bridge", "double-dummy", "embeddings", "weird2vec"], | |
| ) | |
| rendered = ModelCard.from_template( | |
| meta, | |
| template_path=HERE / "card.jinja", | |
| repo=args.repo, | |
| data=DATA, | |
| train=results.get("train"), | |
| eval=results.get("eval"), | |
| ) | |
| (HERE / "card").mkdir(exist_ok=True) | |
| rendered.save(HERE / "card" / "README.md") | |
| def main(): | |
| """Parse arguments and run a command.""" | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| commands = parser.add_subparsers(dest="command", required=True) | |
| data_parser = commands.add_parser("data") | |
| data_parser.add_argument("--deals", type=int, default=200_000, help="train deals") | |
| data_parser.add_argument("--test-hands", type=int, default=1000, help=f"{CHUNK} deals each") | |
| data_parser.add_argument("--shard", type=int, default=0) | |
| data_parser.add_argument("--shards", type=int, default=1) | |
| data_parser.add_argument("--seed", type=int, default=0) | |
| data_parser.add_argument("--output", default="out/data") | |
| data_parser.add_argument("--repo", default=DATA) | |
| data_parser.add_argument("--push", action="store_true") | |
| data_parser.set_defaults(run=data) | |
| train_parser = commands.add_parser("train") | |
| train_parser.add_argument("--data", default=DATA, help="dataset repo or local data folder") | |
| train_parser.add_argument("--max-steps", type=int, default=50_000) | |
| train_parser.add_argument("--batch-size", type=int, default=1024, help="deals per step") | |
| train_parser.add_argument("--lr", type=float, default=3e-4) | |
| train_parser.add_argument("--warmup", type=int, default=1000) | |
| train_parser.add_argument("--value-weight", type=float, default=0.1) | |
| train_parser.add_argument("--aux-weight", type=float, default=0.1) | |
| train_parser.add_argument("--metric-weight", type=float, default=0.0, | |
| help="weight of the loss that ties embedding to table distance") | |
| train_parser.add_argument("--temperature", type=float, default=0.1) | |
| train_parser.add_argument("--target-temperature", type=float, default=0.1) | |
| train_parser.add_argument("--dim", type=int, default=256) | |
| train_parser.add_argument("--depth", type=int, default=4) | |
| train_parser.add_argument("--embed-dim", type=int, default=128) | |
| train_parser.add_argument("--hidden", type=int, default=1024) | |
| train_parser.add_argument("--seed", type=int, default=0) | |
| train_parser.add_argument("--output", default="out/model") | |
| train_parser.add_argument("--repo", default=REPO) | |
| train_parser.add_argument("--revision", default=None, help="branch to push to") | |
| train_parser.add_argument("--push", action="store_true") | |
| train_parser.set_defaults(run=train) | |
| eval_parser = commands.add_parser("eval") | |
| eval_parser.add_argument("--model", default=REPO, help="model repo or local folder") | |
| eval_parser.add_argument("--data", default=DATA) | |
| eval_parser.add_argument("--limit", type=int, default=None, help="number of test hands") | |
| eval_parser.add_argument("--repo", default=REPO, help="where --push stores the result") | |
| eval_parser.add_argument("--revision", default=None, help="model branch") | |
| eval_parser.add_argument("--output", default="out/eval", help="where the hand map goes") | |
| eval_parser.add_argument("--push", action="store_true") | |
| eval_parser.set_defaults(run=evaluate) | |
| embed_parser = commands.add_parser("embed") | |
| target = embed_parser.add_mutually_exclusive_group(required=True) | |
| target.add_argument("--hand", help="PBN hand, spades first, e.g. AKQ32.KJ4.T9.A87") | |
| target.add_argument("--deal", help="PBN deal, North first, e.g. N:AKQ32.KJ4.T9.A87 ...") | |
| embed_parser.add_argument("--model", default=REPO) | |
| embed_parser.add_argument("--data", default=DATA, help="test hands to search") | |
| embed_parser.add_argument("--top", type=int, default=5) | |
| embed_parser.set_defaults(run=embed) | |
| card_parser = commands.add_parser("card") | |
| card_parser.add_argument("--repo", default=REPO) | |
| card_parser.set_defaults(run=card) | |
| args = parser.parse_args() | |
| if not PROGRESS: | |
| datasets.disable_progress_bars() | |
| huggingface_hub.utils.disable_progress_bars() | |
| args.run(args) | |
| if __name__ == "__main__": | |
| main() | |