# /// 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) @torch.no_grad() 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/.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