Download knot2vec.py from jgalego/knot2vec: direct link, hf CLI and curl.
- Browser
- Download file 25.9 kB
-
https://huggingface.co/jgalego/knot2vec/resolve/main/knot2vec.py
- Command line
-
hf download hf://jgalego/knot2vec/knot2vec.py
-
curl -L -o knot2vec.py https://huggingface.co/jgalego/knot2vec/resolve/main/knot2vec.py
25.9 kB
| # /// script | |
| # requires-python = ">=3.12" | |
| # dependencies = [ | |
| # "database-knotinfo", | |
| # "datasets", | |
| # "huggingface-hub", | |
| # "jinja2", | |
| # "matplotlib", | |
| # "numpy", | |
| # "safetensors", | |
| # "snappy", | |
| # "torch", | |
| # "umap-learn", | |
| # ] | |
| # /// | |
| """Knot2Vec: embed knot diagrams so that diagrams of the same knot land together.""" | |
| # 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 snappy | |
| import torch | |
| import umap | |
| from database_knotinfo import link_list | |
| from datasets import Dataset, load_dataset | |
| from huggingface_hub import ( | |
| DatasetCard, | |
| DatasetCardData, | |
| HfApi, | |
| ModelCard, | |
| ModelCardData, | |
| PyTorchModelHubMixin, | |
| snapshot_download, | |
| ) | |
| from huggingface_hub.errors import EntryNotFoundError, RepositoryNotFoundError | |
| from torch import nn | |
| from torch.nn import functional | |
| REPO = "jgalego/knot2vec" | |
| DATA = "jgalego/knot2vec-diagrams" | |
| LANDMARKS = { | |
| "3_1": "trefoil", | |
| "4_1": "figure-eight", | |
| "10_124": "T(3,5)", | |
| "11n_34": "Conway", | |
| } | |
| HERE = Path(__file__).parent | |
| # Progress bars redraw in place, which shows up as garbage in HF Jobs logs. | |
| PROGRESS = sys.stderr.isatty() | |
| # Knots up to 13 crossings have |signature| <= 12, and signatures are even. | |
| MAX_SIGNATURE = 12 | |
| PAD = 4 | |
| 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 gauss(pd): | |
| """Walk the knot along its edges: the crossing each edge enters, and how. | |
| The kind is 2 * positive + over, so 0..3. Needs a PD code whose under-strand goes | |
| from edge a to edge a + 1, as SnapPy writes them. | |
| """ | |
| n = 2 * len(pd) | |
| if not n: | |
| raise ValueError("the diagram has no crossings") | |
| crossing, kind = [None] * n, [None] * n | |
| for i, (a, b, c, d) in enumerate(pd): | |
| if c != (a + 1) % n: | |
| raise ValueError(f"crossing {i} is not oriented along the edges") | |
| positive = (b - d) % n == 1 | |
| over_in = b if (d - b) % n == 1 else d | |
| crossing[a], kind[a] = i, 2 * positive | |
| crossing[over_in], kind[over_in] = i, 2 * positive + 1 | |
| if None in crossing: | |
| raise ValueError("the PD code does not describe a knot") | |
| return crossing, kind | |
| def tokens(crossing, kind, shift=0): | |
| """Start the walk at another edge and name crossings in order of first visit.""" | |
| crossing, kind = crossing[shift:] + crossing[:shift], kind[shift:] + kind[:shift] | |
| names = {} | |
| return [names.setdefault(c, len(names)) for c in crossing], kind | |
| def collate(sequences, device): | |
| """Pad (ids, kinds) pairs into tensors and a padding mask.""" | |
| length = max(len(ids) for ids, _ in sequences) | |
| ids = torch.zeros(len(sequences), length, dtype=torch.long) | |
| kinds = torch.full((len(sequences), length), PAD, dtype=torch.long) | |
| for row, (i, k) in enumerate(sequences): | |
| ids[row, : len(i)] = torch.tensor(i) | |
| kinds[row, : len(k)] = torch.tensor(k) | |
| return ids.to(device), kinds.to(device), (kinds == PAD).to(device) | |
| def normalize(pd): | |
| """Relabel any PD code the way SnapPy writes it.""" | |
| return snappy.Link([tuple(c) for c in pd]).PD_code() | |
| class Knot2Vec(nn.Module, PyTorchModelHubMixin): | |
| """Transformer over the Gauss sequence of a diagram, mean-pooled into a unit vector.""" | |
| def __init__( | |
| self, | |
| dim=256, | |
| depth=6, | |
| heads=8, | |
| max_crossings=64, | |
| embed_dim=256, | |
| target_mean=(0.0, 0.0), | |
| target_std=(1.0, 1.0), | |
| ): | |
| super().__init__() | |
| self.target_mean, self.target_std = list(target_mean), list(target_std) | |
| self.ids = nn.Embedding(max_crossings, dim) | |
| self.kinds = nn.Embedding(PAD + 1, dim, padding_idx=PAD) | |
| self.positions = nn.Embedding(2 * max_crossings, dim) | |
| layer = nn.TransformerEncoderLayer( | |
| dim, heads, 4 * dim, dropout=0.1, 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.signature = nn.Linear(dim, MAX_SIGNATURE + 1) | |
| self.regress = nn.Linear(dim, 2) | |
| def forward(self, ids, kinds, pad): | |
| """Embedding, signature logits and standardized log-determinant and volume.""" | |
| positions = torch.arange(ids.shape[1], device=ids.device) | |
| x = self.ids(ids) + self.kinds(kinds) + self.positions(positions) | |
| h = self.norm(self.encoder(x, src_key_padding_mask=pad)) | |
| keep = (~pad).unsqueeze(-1).to(h.dtype) | |
| pooled = (h * keep).sum(1) / keep.sum(1) | |
| z = functional.normalize(self.project(pooled), dim=-1) | |
| return z, self.signature(pooled), self.regress(pooled) | |
| def run(self, pds, batch_size=1024): | |
| """Embeddings and invariant guesses for PD codes in SnapPy's labeling.""" | |
| device = next(self.parameters()).device | |
| out = [] | |
| for start in range(0, len(pds), batch_size): | |
| chunk = [tokens(*gauss(pd)) for pd in pds[start : start + batch_size]] | |
| out.append(self(*collate(chunk, device))) | |
| z, logits, reg = (torch.cat(parts).float() for parts in zip(*out)) | |
| mean, std = torch.tensor(self.target_mean), torch.tensor(self.target_std) | |
| reg = reg.cpu() * std + mean | |
| return { | |
| "embedding": z, | |
| "signature": logits.argmax(-1).cpu() * 2 - MAX_SIGNATURE, | |
| "determinant": reg[:, 0].exp().round().long(), | |
| "volume": reg[:, 1].clamp(min=0), | |
| } | |
| def embed(self, pds, batch_size=1024): | |
| """Unit-length embeddings for PD codes in SnapPy's labeling.""" | |
| return self.run(pds, batch_size)["embedding"] | |
| def scramble(link, max_crossings, max_moves=40): | |
| """Another diagram of the same knot, from random Reidemeister moves.""" | |
| while True: | |
| view = link.copy() | |
| view.backtrack(random.randint(2, max_moves)) | |
| if random.random() < 0.5: | |
| view.simplify("basic") | |
| if 0 < len(view.crossings) <= max_crossings: | |
| return view.PD_code() | |
| def diagrams(job): | |
| """Train and test diagrams of one knot. Seeded per knot, so reruns match.""" | |
| name, pd, seen, args = job | |
| random.seed(f"{args.seed}:{name}") | |
| link = snappy.Link(pd) | |
| rows = [{"name": name, "pd": pd, "split": "train"}] if seen else [] | |
| splits = ["train"] * (args.views if seen else 0) + ["test"] * args.test_views | |
| rows += [ | |
| {"name": name, "pd": scramble(link, args.max_crossings, args.max_moves), "split": split} | |
| for split in splits | |
| ] | |
| for row in rows: | |
| gauss(row["pd"]) | |
| return rows | |
| def knot_table(limit, holdout, seed): | |
| """Prime knots from KnotInfo with their invariants; a random share is held out.""" | |
| rows = [r for r in link_list()[1:] if r["name"].count("_") == 1 and r["pd_notation"]] | |
| if limit: | |
| rows = random.Random(seed).sample(rows, limit) | |
| unseen = set(random.Random(seed).sample([r["name"] for r in rows], round(holdout * len(rows)))) | |
| knots = [] | |
| for r in rows: | |
| signature = int(r["signature"]) | |
| if abs(signature) > MAX_SIGNATURE: | |
| raise ValueError(f"{r['name']} has signature {signature}") | |
| knots.append( | |
| { | |
| "name": r["name"], | |
| "crossing_number": int(r["crossing_number"]), | |
| "pd": normalize(json.loads(r["pd_notation"])), | |
| "braid": r["braid_notation"], | |
| "signature": signature, | |
| "determinant": int(r["determinant"]), | |
| "volume": float(r["volume"] or 0), | |
| "seen": r["name"] not in unseen, | |
| } | |
| ) | |
| return knots | |
| def data(args): | |
| """Build the knot table and scrambled diagrams; save as parquet, optionally push.""" | |
| start = time.time() | |
| knots = knot_table(args.limit, args.holdout, args.seed) | |
| jobs = [(k["name"], k["pd"], k["seen"], args) for k in knots] | |
| with Pool(cpus()) as pool: | |
| rows = [row for part in pool.imap(diagrams, jobs, chunksize=16) for row in part] | |
| tables = {("knots", "train"): Dataset.from_list(knots)} | |
| for split in ("train", "test"): | |
| tables["diagrams", split] = Dataset.from_list( | |
| [{"name": r["name"], "pd": r["pd"]} for r in rows if r["split"] == split] | |
| ) | |
| out = Path(args.output) | |
| out.mkdir(parents=True, exist_ok=True) | |
| for (config, split), part in tables.items(): | |
| part.to_parquet(out / f"{config}-{split}.parquet") | |
| stats = {f"{config}-{split}": len(part) for (config, split), part in tables.items()} | |
| print(json.dumps({**stats, "minutes": round((time.time() - start) / 60, 1)})) | |
| if args.push: | |
| for (config, split), part in tables.items(): | |
| part.push_to_hub(args.repo, config_name=config, split=split, private=True) | |
| def table(source, config, split): | |
| """A config and split from the Hub or from a local data folder.""" | |
| if Path(source).is_dir(): | |
| files = str(Path(source) / f"{config}-{split}.parquet") | |
| return load_dataset("parquet", data_files=files, split="train") | |
| return load_dataset(source, config, 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 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): | |
| """Contrastive training on pairs of diagrams of the same knot, plus invariant heads.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| random.seed(args.seed) | |
| torch.manual_seed(args.seed) | |
| knots = {k["name"]: k for k in table(args.data, "knots", "train") if k["seen"]} | |
| views = {} | |
| for row in table(args.data, "diagrams", "train"): | |
| views.setdefault(row["name"], []).append(gauss(row["pd"])) | |
| names = sorted(views) | |
| targets = np.array( | |
| [[math.log(knots[n]["determinant"]), knots[n]["volume"]] for n in names], dtype=np.float32 | |
| ) | |
| model = Knot2Vec( | |
| dim=args.dim, | |
| depth=args.depth, | |
| max_crossings=args.max_crossings, | |
| target_mean=targets.mean(0).tolist(), | |
| target_std=(targets.std(0) + 1e-6).tolist(), | |
| ).to(device) | |
| if args.init: | |
| # Copy the weights; embeddings for longer diagrams keep their fresh rows. | |
| weights = model.state_dict() | |
| for key, value in Knot2Vec.from_pretrained(args.init).state_dict().items(): | |
| weights[key][: len(value)] = value | |
| model.load_state_dict(weights) | |
| targets = torch.tensor((targets - targets.mean(0)) / (targets.std(0) + 1e-6), device=device) | |
| signatures = torch.tensor( | |
| [(knots[n]["signature"] + MAX_SIGNATURE) // 2 for n in names], 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) | |
| ) | |
| batch_size = min(args.batch_size, len(names)) | |
| print(json.dumps({"knots": len(names), "diagrams": sum(map(len, views.values())), | |
| "parameters": sum(p.numel() for p in model.parameters())}), flush=True) | |
| start, log = time.time(), {} | |
| model.train() | |
| for step in range(steps): | |
| batch = random.sample(range(len(names)), batch_size) | |
| pairs = [ | |
| random.sample(views[names[i]], 2) if len(views[names[i]]) > 1 else views[names[i]] * 2 | |
| for i in batch | |
| ] | |
| sequences = [ | |
| tokens(c, k, random.randrange(len(c))) for pair in zip(*pairs) for c, k in pair | |
| ] | |
| sig, reg = signatures[batch], targets[batch] | |
| with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): | |
| z, logits, guess = model(*collate(sequences, device)) | |
| z, logits, guess = z.float(), logits.float(), guess.float() | |
| z1, z2 = z[:batch_size], z[batch_size:] | |
| similarity = z1 @ z2.T / args.temperature | |
| labels = torch.arange(batch_size, device=device) | |
| contrastive = (functional.cross_entropy(similarity, labels) | |
| + functional.cross_entropy(similarity.T, labels)) / 2 | |
| invariants = functional.cross_entropy(logits, sig.repeat(2)) + functional.mse_loss( | |
| guess, reg.repeat(2, 1) | |
| ) | |
| loss = contrastive + args.aux_weight * invariants | |
| 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), | |
| "contrastive": round(contrastive.item(), 4), | |
| "invariants": round(invariants.item(), 4), | |
| "pair_top1": round((similarity.argmax(1) == labels).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="knot2vec.py", | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| ) | |
| push_result( | |
| args.repo, | |
| "train", | |
| { | |
| **log, | |
| "data": args.data, | |
| "steps": steps, | |
| "batch_size": batch_size, | |
| "learning_rate": args.lr, | |
| "temperature": args.temperature, | |
| "aux_weight": args.aux_weight, | |
| "init": args.init, | |
| "dim": args.dim, | |
| "depth": args.depth, | |
| "parameters": sum(p.numel() for p in model.parameters()), | |
| "runtime_s": round(time.time() - start), | |
| "device": torch.cuda.get_device_name() if device == "cuda" else "cpu", | |
| }, | |
| args.revision, | |
| ) | |
| def nearest(queries, gallery, k=5, chunk=4096): | |
| """Indices of the k most similar gallery rows for each query.""" | |
| starts = range(0, len(queries), chunk) | |
| return torch.cat([(queries[i : i + chunk] @ gallery.T).topk(k).indices for i in starts]) | |
| def knot_map(gallery, knots, path): | |
| """UMAP of the catalogue's embeddings, coloured by signature, with a few landmarks.""" | |
| xy = umap.UMAP(metric="cosine", random_state=0).fit_transform(gallery.cpu().numpy()) | |
| fig, ax = plt.subplots(figsize=(8, 7)) | |
| points = ax.scatter( | |
| *xy.T, c=knots["signature"], cmap="Spectral", vmin=-MAX_SIGNATURE, vmax=MAX_SIGNATURE, | |
| s=2, linewidths=0, | |
| ) | |
| fig.colorbar(points, ax=ax, label="Signature", shrink=0.7) | |
| for i, name in enumerate(knots["name"]): | |
| if name in LANDMARKS: | |
| ax.annotate( | |
| f"{name} {LANDMARKS[name]}", xy[i], xytext=(6, 6), textcoords="offset points", | |
| fontsize=9, arrowprops={"arrowstyle": "-", "color": "0.3"}, | |
| ) | |
| ax.set_axis_off() | |
| fig.savefig(path, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| def evaluate(args): | |
| """Name held-out diagrams by nearest canonical diagram; check the invariant heads.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = Knot2Vec.from_pretrained(args.model, revision=args.revision).to(device).eval() | |
| knots = table(args.data, "knots", "train").to_dict() | |
| index = {n: i for i, n in enumerate(knots["name"])} | |
| gallery = model.embed(knots["pd"]).to(device) | |
| test = table(args.data, "diagrams", "test") | |
| longest = args.eval_max_crossings or model.positions.num_embeddings // 2 | |
| test = test.filter(lambda row: len(row["pd"]) <= longest) | |
| if args.limit: | |
| test = test.select(random.Random(0).sample(range(len(test)), min(args.limit, len(test)))) | |
| out = model.run(test["pd"]) | |
| top = nearest(out["embedding"].to(device), gallery).cpu() | |
| gold = torch.tensor([index[n] for n in test["name"]]) | |
| truth = {key: torch.tensor(knots[key])[gold] for key in | |
| ("seen", "crossing_number", "signature", "determinant", "volume")} | |
| hits1, hits5 = top[:, 0] == gold, (top == gold[:, None]).any(1) | |
| result = {"model": args.model, "data": args.data, "max_crossings": longest, | |
| "gallery": len(index), "chance_top1": round(1 / len(index), 6)} | |
| for group, mask in (("seen", truth["seen"]), ("unseen", ~truth["seen"])): | |
| if not mask.any(): | |
| continue | |
| crossings = truth["crossing_number"] | |
| result[group] = { | |
| "n": int(mask.sum()), | |
| "top1": round(hits1[mask].float().mean().item(), 3), | |
| "top5": round(hits5[mask].float().mean().item(), 3), | |
| "signature_accuracy": round( | |
| (out["signature"] == truth["signature"])[mask].float().mean().item(), 3 | |
| ), | |
| "determinant_within_10pct": round( | |
| (abs(out["determinant"] - truth["determinant"]) <= 0.1 * truth["determinant"])[mask] | |
| .float().mean().item(), 3 | |
| ), | |
| "volume_mae": round((out["volume"] - truth["volume"])[mask].abs().mean().item(), 3), | |
| "top1_by_crossings": { | |
| int(c): round(hits1[mask & (crossings == c)].float().mean().item(), 3) | |
| for c in crossings[mask].unique() | |
| }, | |
| } | |
| print(json.dumps(result, indent=1)) | |
| Path(args.output).mkdir(parents=True, exist_ok=True) | |
| knot_map(gallery, knots, Path(args.output) / "map.png") | |
| if args.push: | |
| push_result(args.repo, args.result, 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 embed(args): | |
| """Print the nearest catalogued knots and invariant guesses for one PD code.""" | |
| model = Knot2Vec.from_pretrained(args.model).eval() | |
| knots = table(args.data, "knots", "train") | |
| out = model.run([normalize(json.loads(args.pd))]) | |
| gallery = model.embed(knots["pd"]) | |
| top = nearest(out["embedding"], gallery, args.top)[0] | |
| similarity = out["embedding"] @ gallery[top].T | |
| print( | |
| json.dumps( | |
| { | |
| "nearest": [ | |
| {"knot": knots[int(i)]["name"], "cosine": round(s.item(), 3)} | |
| for i, s in zip(top, similarity[0]) | |
| ], | |
| "signature": int(out["signature"][0]), | |
| "determinant": int(out["determinant"][0]), | |
| "volume": round(out["volume"][0].item(), 3), | |
| "embedding": [round(x, 4) for x in out["embedding"][0].tolist()], | |
| }, | |
| indent=1, | |
| ) | |
| ) | |
| def dataset_card(args): | |
| """Render dataset.jinja into dataset/README.md, keeping the metadata the data push wrote.""" | |
| try: | |
| meta = DatasetCard.load(DATA).data.to_dict() | |
| except (RepositoryNotFoundError, EntryNotFoundError): | |
| meta = {} | |
| sizes = { | |
| f"{info['config_name']}-{split['name']}": split["num_examples"] | |
| for info in meta.get("dataset_info", []) | |
| for split in info["splits"] | |
| } | |
| meta.update( | |
| license="gpl-3.0", | |
| pretty_name="Knot2Vec diagrams", | |
| task_categories=["feature-extraction"], | |
| size_categories=["100K<n<1M"], | |
| tags=["knot-theory", "mathematics", "weird2vec"], | |
| ) | |
| rendered = DatasetCard.from_template( | |
| DatasetCardData(**meta), | |
| template_path=HERE / "dataset.jinja", | |
| repo=args.repo, | |
| data=DATA, | |
| sizes=sizes, | |
| ) | |
| (HERE / "dataset").mkdir(exist_ok=True) | |
| rendered.save(HERE / "dataset" / "README.md") | |
| def card(args): | |
| """Render the model card from card.jinja and the model repo's results, then the dataset card.""" | |
| 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=["knot-theory", "embeddings", "contrastive-learning", "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") | |
| dataset_card(args) | |
| 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("--limit", type=int, default=None, help="number of knots") | |
| data_parser.add_argument("--views", type=int, default=48, help="train diagrams per seen knot") | |
| data_parser.add_argument("--test-views", type=int, default=4, help="test diagrams per knot") | |
| data_parser.add_argument("--holdout", type=float, default=0.1, help="share of unseen knots") | |
| data_parser.add_argument("--max-crossings", type=int, default=64) | |
| data_parser.add_argument("--max-moves", type=int, default=40, help="Reidemeister moves") | |
| 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=20_000) | |
| train_parser.add_argument("--batch-size", type=int, default=512, help="knots 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("--temperature", type=float, default=0.05) | |
| train_parser.add_argument("--aux-weight", type=float, default=1.0) | |
| train_parser.add_argument("--dim", type=int, default=256) | |
| train_parser.add_argument("--depth", type=int, default=6) | |
| train_parser.add_argument("--max-crossings", type=int, default=64) | |
| train_parser.add_argument("--init", default=None, help="model repo or folder to start from") | |
| 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 diagrams") | |
| 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("--result", default="eval", help="results/<result>.json for --push") | |
| eval_parser.add_argument("--eval-max-crossings", type=int, default=None) | |
| eval_parser.add_argument("--output", default="out/eval", help="where the knot map goes") | |
| eval_parser.add_argument("--push", action="store_true") | |
| eval_parser.set_defaults(run=evaluate) | |
| embed_parser = commands.add_parser("embed") | |
| embed_parser.add_argument("--pd", required=True, help="PD code as JSON, e.g. [[1,5,2,4],...]") | |
| embed_parser.add_argument("--model", default=REPO) | |
| embed_parser.add_argument("--data", default=DATA) | |
| 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() | |