""" Evaluation for hierarchy embeddings: for each held-out (child, parent) edge, rank the true parent among ALL nodes by distance from the child. Perfect recovery -> rank 1 for every edge. This is the standard protocol from Nickel & Kiela (2017) and is what makes the Poincaré-vs-Euclidean comparison a checkable number rather than a narrative claim. """ from __future__ import annotations from typing import Dict, List, Tuple import torch from .embed_hierarchy import HierarchyEmbedding @torch.no_grad() def mean_rank_and_map( model: HierarchyEmbedding, edges: List[Tuple[int, int]], device: str = "cpu", ) -> Dict[str, float]: model.eval() all_idx = torch.arange(model.num_nodes, device=device) all_pts = model.points(all_idx) # (N, D) ranks = [] for child_idx, parent_idx in edges: child_pt = model.points(torch.tensor([child_idx], device=device)) # (1, D) dists = model.distance(child_pt.expand(model.num_nodes, -1), all_pts) # (N,) true_dist = dists[parent_idx] closer = (dists < true_dist).sum().item() if child_idx != parent_idx and dists[child_idx] < true_dist: closer -= 1 rank = closer + 1 ranks.append(rank) ranks_t = torch.tensor(ranks, dtype=torch.float) return { "mean_rank": ranks_t.mean().item(), "mrr": (1.0 / ranks_t).mean().item(), "map": (1.0 / ranks_t).mean().item(), # single relevant item per query "hits@1": (ranks_t <= 1).float().mean().item(), "hits@3": (ranks_t <= 3).float().mean().item(), "hits@10": (ranks_t <= 10).float().mean().item(), "n_queries": len(ranks), } def compare_geometries( poincare_model: HierarchyEmbedding, euclidean_model: HierarchyEmbedding, test_edges: List[Tuple[int, int]], device: str = "cpu", ) -> Dict[str, Dict[str, float]]: p_metrics = mean_rank_and_map(poincare_model, test_edges, device=device) e_metrics = mean_rank_and_map(euclidean_model, test_edges, device=device) return { "poincare": p_metrics, "euclidean": e_metrics, "delta_mrr": p_metrics["mrr"] - e_metrics["mrr"], "delta_mean_rank": e_metrics["mean_rank"] - p_metrics["mean_rank"], # positive = Poincare better (lower rank) "delta_hits@10": p_metrics["hits@10"] - e_metrics["hits@10"], } @torch.no_grad() def radius_diagnostics(model: HierarchyEmbedding, node_depths: Dict[int, int]) -> Dict[str, float]: radii = model.radii() idxs = sorted(node_depths.keys()) r = torch.tensor([radii[i].item() for i in idxs]) d = torch.tensor([float(node_depths[i]) for i in idxs]) r_centered = r - r.mean() d_centered = d - d.mean() denom = (r_centered.pow(2).sum().sqrt() * d_centered.pow(2).sum().sqrt()) corr = (r_centered * d_centered).sum() / denom if denom > 0 else torch.tensor(0.0) return { "radius_depth_correlation": corr.item(), "mean_radius": r.mean().item(), "std_radius": r.std().item(), "min_radius": r.min().item(), "max_radius": r.max().item(), }