knot2vec / knot2vec.py
jgalego's picture
Add eval cap option
687bf73 verified
Raw History Blame Contribute Delete
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)
@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/<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()