| """ |
| Track B driver: train Poincaré vs Euclidean hierarchy embeddings under |
| identical conditions and report the comparison table. |
| |
| python -m src.run_hierarchy_embed --synthetic --n-nodes 500 --dim 8 |
| python -m src.run_hierarchy_embed --pbdb-taxa Dinosauria Mammalia --dim 8 |
| """ |
| from __future__ import annotations |
| import argparse |
| import random |
| from typing import List, Tuple |
|
|
| import numpy as np |
| import torch |
|
|
| from .embed_hierarchy import train_hierarchy_embedding |
| from .eval_hierarchy import compare_geometries, radius_diagnostics |
| from .synthetic_tree import get_synthetic_tree_dataset |
| from .data_pbdb_taxonomy import get_pbdb_taxonomy_dataset |
| from .provenance import CheckpointStore, hash_code |
| import os |
|
|
|
|
| def set_all_seeds(seed: int): |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
|
|
|
|
| def split_edges(dataset, test_frac: float = 0.2, seed: int = 0) -> Tuple[List, List]: |
| """ |
| WARNING, found by actually running this on a real tree (not assumed): |
| for a TREE, every non-root node has EXACTLY ONE parent edge (verified: |
| max count of any node as a 'child' across all edges is 1). Holding out |
| that single edge removes 100% of that node's positive training signal |
| -- it has nothing left pulling it toward its true parent. This makes a |
| held-out-edge split fundamentally unsuited to evaluating a tree |
| embedding's generalization; it doesn't test what you think it tests. |
| Use eval_mode='reconstruction' (the default) instead, which matches |
| Nickel & Kiela's own protocol: train on all edges, evaluate rank |
| recovery on those same edges (measuring embedding CAPACITY/fidelity at |
| a given dimension, not generalization to unseen relations -- which is |
| what "does hyperbolic geometry need fewer dimensions" actually means). |
| This function is kept only for the (explicitly discouraged) held_out |
| eval mode. |
| """ |
| edges = list(dataset.edge_idx) |
| rng = random.Random(seed) |
| rng.shuffle(edges) |
| n_test = max(1, int(len(edges) * test_frac)) |
| return edges[n_test:], edges[:n_test] |
|
|
|
|
| def main(): |
| p = argparse.ArgumentParser() |
| src = p.add_mutually_exclusive_group(required=True) |
| src.add_argument("--synthetic", action="store_true", |
| help="Use a synthetic random tree (runs anywhere, no network).") |
| src.add_argument("--pbdb-taxa", nargs="+", default=None, |
| help="Real PBDB taxon groups (requires network to paleobiodb.org).") |
| p.add_argument("--tree-type", choices=["random_recursive", "balanced"], default="random_recursive", |
| help="'random_recursive': shallow/bushy, depth~ln(n). " |
| "'balanced': fixed branching_factor/depth, genuinely deep -- " |
| "the regime the hyperbolic-advantage literature targets.") |
| p.add_argument("--branching-factor", type=int, default=3) |
| p.add_argument("--tree-depth", type=int, default=6) |
| p.add_argument("--n-nodes", type=int, default=500, help="Synthetic tree size (random_recursive only).") |
| p.add_argument("--dims", type=int, nargs="+", default=[8], |
| help="One or more embedding dimensions to sweep.") |
| p.add_argument("--epochs", type=int, default=200) |
| p.add_argument("--lr", type=float, default=0.01) |
| p.add_argument("--margin", type=float, default=1.0) |
| p.add_argument("--neg-samples", type=int, default=10) |
| p.add_argument("--seed", type=int, default=0) |
| p.add_argument("--seeds", type=int, nargs="+", default=None, |
| help="If given, run all (dim, loss_type) combos across these " |
| "seeds and report mean +/- std, instead of a single seed.") |
| p.add_argument("--loss-types", choices=["margin", "softmax"], nargs="+", default=["margin"], |
| help="One or both losses to compare under identical conditions.") |
| p.add_argument("--burn-in-epochs", type=int, default=0) |
| p.add_argument("--burn-in-lr-mult", type=float, default=0.1) |
| p.add_argument("--learnable-c", action="store_true", |
| help="Let curvature be learned per model instead of fixed at --c. " |
| "Applies to the Poincare model only (a fixed scalar for the " |
| "Euclidean baseline has no equivalent meaning). Guarded by " |
| "--c-min/--c-max and a separate curvature learning rate -- " |
| "see embed_hierarchy.py for why those guards are required.") |
| p.add_argument("--c", type=float, default=1.0, help="Fixed curvature when --learnable-c is not set.") |
| p.add_argument("--c-values", type=float, nargs="+", default=None, |
| help="Sweep a fixed-c grid instead of a single --c (recommended " |
| "primary methodology: interpretable, no coupled instability, " |
| "matches how most hyperbolic embedding papers report results). " |
| "Mutually exclusive with --learnable-c.") |
| p.add_argument("--c-min", type=float, default=0.1) |
| p.add_argument("--c-max", type=float, default=3.0) |
| p.add_argument("--curvature-lr-mult", type=float, default=0.1) |
| p.add_argument("--optimizer", choices=["radam", "rsgd"], default="radam", |
| help="radam (RiemannianAdam, used throughout so far) or " |
| "rsgd (RiemannianSGD -- what Nickel & Kiela's original " |
| "2017 paper actually used; never isolated as a variable " |
| "in this project before now).") |
| p.add_argument("--eval-mode", choices=["reconstruction", "held_out"], default="reconstruction", |
| help="'reconstruction' (default, matches Nickel & Kiela's protocol): " |
| "train on all edges, evaluate rank-recovery on those same edges " |
| "-- measures embedding capacity/fidelity at a given dimension. " |
| "'held_out': held-out edge split -- WARNING, verified broken for " |
| "tree data (every node has exactly one parent edge, so holding " |
| "it out removes 100%% of that node's training signal). Kept only " |
| "for illustration of that failure mode.") |
| args = p.parse_args() |
|
|
| set_all_seeds(args.seed) |
|
|
| if args.synthetic: |
| dataset, meta, provenance = get_synthetic_tree_dataset( |
| n_nodes=args.n_nodes, seed=args.seed, tree_type=args.tree_type, |
| branching_factor=args.branching_factor, depth=args.tree_depth, |
| ) |
| node_depths = meta["node_depths_by_idx"] |
| else: |
| dataset, meta, provenance = get_pbdb_taxonomy_dataset(base_names=args.pbdb_taxa) |
| node_depths = None |
|
|
| if args.eval_mode == "reconstruction": |
| train_edges = list(dataset.edge_idx) |
| test_edges = train_edges |
| else: |
| print("\n*** WARNING: --eval-mode held_out is known-broken for tree data. ***") |
| print("*** Every node has exactly one parent edge; holding it out removes ***") |
| print("*** 100% of that node's training signal. Numbers below will look ***") |
| print("*** close to random and should NOT be read as 'geometry doesn't help'.***\n") |
| train_edges, test_edges = split_edges(dataset, test_frac=0.2, seed=args.seed) |
| print(f"[data] provenance={provenance} nodes={dataset.num_nodes} " |
| f"train_edges={len(train_edges)} test_edges={len(test_edges)}") |
|
|
| class _EdgeSubset: |
| def __init__(self, edge_idx, num_nodes): |
| self.edge_idx = edge_idx |
| self.num_nodes = num_nodes |
| train_ds = _EdgeSubset(train_edges, dataset.num_nodes) |
|
|
| if args.learnable_c and args.c_values: |
| p.error("--learnable-c and --c-values are mutually exclusive") |
|
|
| results_table = [] |
| seeds = args.seeds or [args.seed] |
| c_grid = args.c_values or [args.c] |
| for dim in args.dims: |
| for loss_type in args.loss_types: |
| for c_val in c_grid: |
| label = f"dim={dim}, loss={loss_type}, " + ( |
| "learnable_c" if args.learnable_c else f"c={c_val}" |
| ) |
| print(f"\n{'='*64}\n{label}\n{'='*64}") |
| deltas_mrr, deltas_rank = [], [] |
| per_seed = [] |
| for sd in seeds: |
| p_model, p_metrics = train_hierarchy_embedding( |
| train_ds, geometry="poincare", dim=dim, epochs=args.epochs, |
| lr=args.lr, margin=args.margin, neg_samples=args.neg_samples, |
| loss_type=loss_type, burn_in_epochs=args.burn_in_epochs, |
| burn_in_lr_mult=args.burn_in_lr_mult, c=c_val, |
| learnable_c=args.learnable_c, curvature_lr_mult=args.curvature_lr_mult, |
| optimizer_type=args.optimizer, |
| c_min=args.c_min, c_max=args.c_max, seed=sd, |
| ) |
| e_model, e_metrics = train_hierarchy_embedding( |
| train_ds, geometry="euclidean", dim=dim, epochs=args.epochs, |
| lr=args.lr, margin=args.margin, neg_samples=args.neg_samples, |
| loss_type=loss_type, burn_in_epochs=args.burn_in_epochs, |
| burn_in_lr_mult=args.burn_in_lr_mult, seed=sd, |
| ) |
| comparison = compare_geometries(p_model, e_model, test_edges) |
| deltas_mrr.append(comparison["delta_mrr"]) |
| deltas_rank.append(comparison["delta_mean_rank"]) |
| per_seed.append(comparison) |
| print(f" seed={sd:<5} P_mrr={comparison['poincare']['mrr']:.4f} " |
| f"E_mrr={comparison['euclidean']['mrr']:.4f} " |
| f"delta_mrr={comparison['delta_mrr']:+.4f}") |
|
|
| deltas_mrr_t = torch.tensor(deltas_mrr) |
| mean_delta = deltas_mrr_t.mean().item() |
| std_delta = deltas_mrr_t.std().item() if len(seeds) > 1 else 0.0 |
| wins = sum(1 for d in deltas_mrr if d > 0) |
| print(f" --> mean delta_mrr = {mean_delta:+.4f} +/- {std_delta:.4f} " |
| f"(Poincare wins {wins}/{len(seeds)} seeds)") |
|
|
| if node_depths is not None: |
| p_radius = radius_diagnostics(p_model, node_depths) |
| e_radius = radius_diagnostics(e_model, node_depths) |
| print(f" radius-depth correlation poincare={p_radius['radius_depth_correlation']:+.4f}" |
| f" euclidean={e_radius['radius_depth_correlation']:+.4f}") |
|
|
| if args.learnable_c: |
| print(f" learned curvature (last seed): {p_model.manifold.c.item():.4f} " |
| f"(started at {c_val})") |
|
|
| results_table.append({ |
| "dim": dim, "loss_type": loss_type, "c": c_val, |
| "learnable_c": args.learnable_c, "seeds": seeds, |
| "mean_delta_mrr": mean_delta, "std_delta_mrr": std_delta, |
| "wins": wins, "n_seeds": len(seeds), |
| }) |
|
|
| print(f"\n{'='*64}\nSummary across dimensions, losses, and curvatures\n{'='*64}") |
| print(f"{'dim':>5} {'loss':>10} {'c':>8} {'mean_delta_mrr':>16} {'std':>8} {'wins':>8}") |
| for r in results_table: |
| c_label = "learned" if r["learnable_c"] else f"{r['c']:.3f}" |
| print(f"{r['dim']:>5} {r['loss_type']:>10} {c_label:>8} {r['mean_delta_mrr']:>+16.4f} " |
| f"{r['std_delta_mrr']:>8.4f} {r['wins']}/{r['n_seeds']:>6}") |
|
|
| store = CheckpointStore(checkpoints_dir="checkpoints") |
| code_hash = hash_code(os.path.dirname(os.path.abspath(__file__))) |
| store.save( |
| model_state={"results_table": results_table}, |
| config={"dims": args.dims, "epochs": args.epochs, "lr": args.lr, |
| "margin": args.margin, "neg_samples": args.neg_samples, "seed": args.seed}, |
| dataset_hash=meta.get("edge_hash") or f"synthetic_{meta.get('seed')}_{dataset.num_nodes}", |
| code_hash=code_hash, |
| data_provenance=provenance, |
| extra={"meta": {k: v for k, v in meta.items() if k != "node_to_idx"}}, |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|