poincare-hyper / src /eval_hierarchy.py
DHDRL's picture
Rename eval_hierarchy.py to src/eval_hierarchy.py
4c4dbc9 verified
Raw
History Blame Contribute Delete
3.11 kB
"""
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(),
}