File size: 3,109 Bytes
ae73c7f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
"""
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(),
    }