File size: 3,779 Bytes
9e4b2c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
"""Retrain on 90% of edges and score the held-out 10%: does it generalise?

Same procedure as the released pyg/heldout_check.py, with the sample size scaled
down for this much smaller graph (848 nodes, ~359k possible node pairs)."""

import argparse
import time

import torch
from torch_geometric.utils import to_undirected

from config import Paths, add_source_arg
from node2vec_model import build_model, load_graph

_ap = argparse.ArgumentParser()
add_source_arg(_ap)
_ap.add_argument("--epochs", type=int, default=200,
                 help="match the budget used for the saved embeddings")
cli = _ap.parse_args()
paths = Paths(cli.source_table)

dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
g = torch.Generator(device="cpu").manual_seed(1)

data = load_graph(paths.graph)
N = data.num_nodes
ei = data.edge_index.to(dev)
u, v = ei.min(0).values.to(torch.int64), ei.max(0).values.to(torch.int64)
keys = torch.unique(u * N + v)

perm = torch.randperm(keys.numel(), generator=g).to(dev)
n_test = keys.numel() // 10
test_keys, train_keys = keys[perm[:n_test]], keys[perm[n_test:]]
train_ei = to_undirected(
    torch.stack([train_keys // N, train_keys % N]), num_nodes=N)
deg_tr = torch.bincount(train_ei[0].to(dev), minlength=N)
print(f"train edges {train_keys.numel():,}  held-out {test_keys.numel():,}  "
      f"nodes isolated by the split: {(deg_tr == 0).sum().item()}")

args = argparse.Namespace(
    embedding_dim=128, walk_length=20, context_size=10, walks_per_node=10,
    num_negative_samples=1, p=1.0, q=1.0)
model = build_model(type("D", (), {"edge_index": train_ei.cpu(), "num_nodes": N}),
                    args, dev)
loader = model.loader(batch_size=128, shuffle=True, num_workers=4)
opt = torch.optim.SparseAdam(list(model.parameters()), lr=0.01)

for epoch in range(1, cli.epochs + 1):
    model.train()
    t0, tot, n = time.perf_counter(), 0.0, 0
    for pos_rw, neg_rw in loader:
        opt.zero_grad()
        loss = model.loss(pos_rw.to(dev), neg_rw.to(dev))
        loss.backward()
        opt.step()
        tot, n = tot + loss.item(), n + 1
    if epoch % max(cli.epochs // 5, 1) == 0 or epoch == 1:
        print(f"  epoch {epoch:>3}  loss {tot / n:.4f}  {time.perf_counter() - t0:.1f}s")

model.eval()
with torch.no_grad():
    z = model()
zc = torch.nn.functional.normalize(z, dim=1)


def sample_non_edges(n):
    out, got = [], 0
    while got < n:
        c = torch.randint(N, (2, n), generator=g).to(dev)
        a, b = c.min(0).values, c.max(0).values
        ok = (a != b) & ~torch.isin(a.to(torch.int64) * N + b.to(torch.int64), keys)
        out.append(torch.stack([a[ok], b[ok]]))
        got += int(ok.sum())
    return torch.cat(out, dim=1)[:, :n]


def auc(pos, neg):
    s = torch.cat([pos, neg]).double()
    order = torch.argsort(s)
    ranks = torch.empty_like(s)
    ranks[order] = torch.arange(1, s.numel() + 1, dtype=torch.float64, device=s.device)
    np_, nn_ = pos.numel(), neg.numel()
    return ((ranks[:np_].sum() - np_ * (np_ + 1) / 2) / (np_ * nn_)).item()


M = 20_000
te = test_keys[torch.randperm(n_test, generator=g)[:M].to(dev)]
pe = torch.stack([te // N, te % N])
tr = train_keys[torch.randperm(train_keys.numel(), generator=g)[:M].to(dev)]
tre = torch.stack([tr // N, tr % N])
ne = sample_non_edges(M)

print("\nlink prediction AUC (cosine):")
for name, e in (("held-out edges (never seen)", pe), ("train edges (in-sample)", tre)):
    print(f"  {name:<28} {auc((zc[e[0]] * zc[e[1]]).sum(1), (zc[ne[0]] * zc[ne[1]]).sum(1)):.4f}")

# degree-only baseline: would preferential attachment alone explain it?
d = deg_tr.double()
print("\ndegree baseline (preferential attachment, d_u * d_v):")
print(f"  held-out edges               "
      f"{auc(d[pe[0]] * d[pe[1]], d[ne[0]] * d[ne[1]]):.4f}")