Skip to content

Commit 7f0d2f0

Browse files
authored
CTAN (#338)
* getting ctan forward * getting ctan forward * Upd * Upd * Add delta t norm * Add delta t norm * Add sum merge op to link predictor * Add init time to memory * Expose epsilon/gamma diffusion params * Expose epsilon/gamma diffusion params * Expose epsilon/gamma diffusion params * Expose epsilon/gamma diffusion params * Add tanh final activation * drill down further * overhaul wip * wip * wip * >0.5mrr * Revert * Upd * Add integration test * upd * Port decoder * Port memory * update readme * add tests * uipd * update docs * Simplify forward * fix docs
1 parent 66a2125 commit 7f0d2f0

11 files changed

Lines changed: 557 additions & 18 deletions

File tree

README.md

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -37,10 +37,10 @@ It provides a unified abstraction for both discrete and continuous-time graphs,
3737

3838
To request a method for prioritization, please [open an issue](https://github.com/tgm-team/tgm/issues) or [join the discussion](https://github.com/tgm-team/tgm/discussions).
3939

40-
| Status | Methods |
41-
| ----------- | --------------------------------------------------------------------------------------------------------------------------------------------------- |
42-
| Implemented | EdgeBank[^1], GCN[^2], GC-LSTM[^3], GraphMixer[^4], TGAT[^5], TGN[^6], DygFormer[^7], TPNet[^8], ROLAND [^13], PopTrack [^14], TNCN[^9], Base3[^15] |
43-
| Planned | DyGMamba[^10], NAT[^11] |
40+
| Status | Methods |
41+
| ----------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------- |
42+
| Implemented | EdgeBank[^1], GCN[^2], GC-LSTM[^3], GraphMixer[^4], TGAT[^5], TGN[^6], DygFormer[^7], TPNet[^8], ROLAND [^13], PopTrack [^14], TNCN[^9], Base3[^15] CTAN[^16] |
43+
| Planned | DyGMamba[^10], NAT[^11] |
4444

4545
## Installation
4646

@@ -228,6 +228,8 @@ We welcome contributions. If you encounter problems or would like to propose a n
228228
229229
[^15]: [Base3: a simple interpolation-based ensemble method for robust dynamic link prediction](https://www.arxiv.org/abs/2506.12764)
230230
231+
[^16]: [Long Range Propagation on Continuous-Time Dynamic Graphs](https://arxiv.org/abs/2406.02740)
232+
231233
[^10]: [DyGMamba: Efficiently Modeling Long-Term Temporal Dependency on Continuous-Time Dynamic Graphs with State Space Models](https://arxiv.org/abs/2408.04713)
232234
233235
[^11]: [Neighborhood-aware Scalable Temporal Network Representation Learning](https://arxiv.org/abs/2209.01084)

docs/api/nn.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
## Encoders
22

3+
::: tgm.nn.encoder.ctan
34
::: tgm.nn.encoder.dygformer
45
::: tgm.nn.encoder.tpnet
56
::: tgm.nn.encoder.gclstm

examples/linkproppred/ctan.py

Lines changed: 280 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,280 @@
1+
import argparse
2+
from typing import Tuple
3+
4+
import numpy as np
5+
import torch
6+
import torch.nn as nn
7+
import torch.nn.functional as F
8+
from tgb.linkproppred.evaluate import Evaluator
9+
from tqdm import tqdm
10+
11+
from tgm import DGraph
12+
from tgm.constants import (
13+
METRIC_TGB_LINKPROPPRED,
14+
PADDED_NODE_ID,
15+
RECIPE_TGB_LINK_PRED,
16+
)
17+
from tgm.data import DGData, DGDataLoader
18+
from tgm.hooks import DeduplicationHook, RecencyNeighborHook, RecipeRegistry
19+
from tgm.nn import LinkPredictor
20+
from tgm.nn.encoder import CTAN, CTANMemory, LastAggregator
21+
from tgm.util.logging import enable_logging, log_gpu, log_latency, log_metric
22+
from tgm.util.seed import seed_everything
23+
24+
parser = argparse.ArgumentParser(
25+
description='CTAN LinkPropPred Example',
26+
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
27+
)
28+
parser.add_argument('--seed', type=int, default=1337, help='random seed to use')
29+
parser.add_argument('--dataset', type=str, default='tgbl-wiki', help='Dataset name')
30+
parser.add_argument('--bsize', type=int, default=200, help='batch size')
31+
parser.add_argument('--device', type=str, default='cpu', help='torch device')
32+
parser.add_argument('--epochs', type=int, default=200, help='number of epochs')
33+
parser.add_argument('--n-layers', type=int, default=3, help='number of GNN layers')
34+
parser.add_argument(
35+
'--epsilon', type=float, default=0.5, help='discretization step size'
36+
)
37+
parser.add_argument('--gamma', type=float, default=0.1, help='diffusion strength')
38+
parser.add_argument('--lr', type=float, default=0.0001, help='learning rate')
39+
parser.add_argument(
40+
'--n-nbrs',
41+
type=int,
42+
nargs='+',
43+
default=[32],
44+
help='num sampled nbrs at each hop',
45+
)
46+
parser.add_argument('--time-dim', type=int, default=256, help='time encoding dimension')
47+
parser.add_argument('--embed-dim', type=int, default=256, help='attention dimension')
48+
parser.add_argument('--memory-dim', type=int, default=256, help='memory dimension')
49+
parser.add_argument(
50+
'--log-file-path', type=str, default=None, help='Optional path to write logs'
51+
)
52+
53+
args = parser.parse_args()
54+
enable_logging(log_file_path=args.log_file_path)
55+
56+
57+
@log_gpu
58+
@log_latency
59+
def train(
60+
loader: DGDataLoader,
61+
memory: nn.Module,
62+
encoder: nn.Module,
63+
decoder: nn.Module,
64+
opt: torch.optim.Optimizer,
65+
) -> float:
66+
memory.train()
67+
encoder.train()
68+
decoder.train()
69+
total_loss = 0
70+
static_node_x = loader.dgraph.static_node_x
71+
72+
memory.reset_state()
73+
74+
for batch in tqdm(loader):
75+
opt.zero_grad()
76+
77+
nbr_nodes = batch.nbr_nids[0].flatten()
78+
nbr_mask = nbr_nodes != PADDED_NODE_ID
79+
80+
num_nbrs = len(nbr_nodes) // (
81+
len(batch.edge_src) + len(batch.edge_dst) + len(batch.neg)
82+
)
83+
src_nodes = torch.cat(
84+
[
85+
batch.edge_src.repeat_interleave(num_nbrs),
86+
batch.edge_dst.repeat_interleave(num_nbrs),
87+
batch.neg.repeat_interleave(num_nbrs),
88+
]
89+
)
90+
nbr_edge_index = torch.stack(
91+
[
92+
batch.global_to_local(src_nodes[nbr_mask]),
93+
batch.global_to_local(nbr_nodes[nbr_mask]),
94+
]
95+
).to(dtype=torch.int64)
96+
97+
nbr_edge_time = batch.nbr_edge_time[0].flatten()[nbr_mask]
98+
nbr_edge_x = batch.nbr_edge_x[0].flatten(0, -2).float()[nbr_mask]
99+
100+
z, last_update = memory(batch.unique_nids)
101+
z = torch.cat([z, static_node_x[batch.unique_nids]], dim=-1)
102+
z = encoder(z, last_update, nbr_edge_index, nbr_edge_time, nbr_edge_x)
103+
104+
inv_src = batch.global_to_local(batch.edge_src)
105+
inv_dst = batch.global_to_local(batch.edge_dst)
106+
inv_neg = batch.global_to_local(batch.neg)
107+
pos_out = decoder(z[inv_src], z[inv_dst])
108+
neg_out = decoder(z[inv_src], z[inv_neg])
109+
110+
loss = F.binary_cross_entropy_with_logits(pos_out, torch.ones_like(pos_out))
111+
loss += F.binary_cross_entropy_with_logits(neg_out, torch.zeros_like(neg_out))
112+
113+
# Update memory with ground-truth state.
114+
memory.update_state(
115+
batch.edge_src, batch.edge_dst, batch.edge_time, z[inv_src], z[inv_dst]
116+
)
117+
118+
loss.backward()
119+
opt.step()
120+
total_loss += float(loss)
121+
122+
memory.detach()
123+
124+
return total_loss
125+
126+
127+
@log_gpu
128+
@log_latency
129+
@torch.no_grad()
130+
def eval(
131+
loader: DGDataLoader,
132+
memory: nn.Module,
133+
encoder: nn.Module,
134+
decoder: nn.Module,
135+
evaluator: Evaluator,
136+
) -> float:
137+
memory.eval()
138+
encoder.eval()
139+
decoder.eval()
140+
perf_list = []
141+
static_node_x = loader.dgraph.static_node_x
142+
143+
for batch in tqdm(loader):
144+
nbr_nodes = batch.nbr_nids[0].flatten()
145+
nbr_mask = nbr_nodes != PADDED_NODE_ID
146+
147+
num_nbrs = len(nbr_nodes) // (
148+
len(batch.edge_src) + len(batch.edge_dst) + len(batch.neg)
149+
)
150+
src_nodes = torch.cat(
151+
[
152+
batch.edge_src.repeat_interleave(num_nbrs),
153+
batch.edge_dst.repeat_interleave(num_nbrs),
154+
batch.neg.repeat_interleave(num_nbrs),
155+
]
156+
)
157+
nbr_edge_index = torch.stack(
158+
[
159+
batch.global_to_local(src_nodes[nbr_mask]),
160+
batch.global_to_local(nbr_nodes[nbr_mask]),
161+
]
162+
).to(dtype=torch.int64)
163+
nbr_edge_time = batch.nbr_edge_time[0].flatten()[nbr_mask]
164+
nbr_edge_x = batch.nbr_edge_x[0].flatten(0, -2).float()[nbr_mask]
165+
166+
z, last_update = memory(batch.unique_nids)
167+
z = torch.cat([z, static_node_x[batch.unique_nids]], dim=-1)
168+
z = encoder(z, last_update, nbr_edge_index, nbr_edge_time, nbr_edge_x)
169+
170+
for idx, neg_batch in enumerate(batch.neg_batch_list):
171+
dst_ids = torch.cat([batch.edge_dst[idx].unsqueeze(0), neg_batch])
172+
src_ids = batch.edge_src[idx].repeat(len(dst_ids))
173+
174+
inv_src = batch.global_to_local(src_ids)
175+
inv_dst = batch.global_to_local(dst_ids)
176+
y_pred = decoder(z[inv_src], z[inv_dst]).sigmoid()
177+
178+
input_dict = {
179+
'y_pred_pos': y_pred[0],
180+
'y_pred_neg': y_pred[1:],
181+
'eval_metric': [METRIC_TGB_LINKPROPPRED],
182+
}
183+
perf_list.append(evaluator.eval(input_dict)[METRIC_TGB_LINKPROPPRED])
184+
185+
# Update memory with ground-truth state.
186+
memory.update_state(
187+
batch.edge_src, batch.edge_dst, batch.edge_time, z[inv_src], z[inv_dst]
188+
)
189+
190+
return float(np.mean(perf_list))
191+
192+
193+
seed_everything(args.seed)
194+
evaluator = Evaluator(name=args.dataset)
195+
196+
full_data = DGData.from_tgb(args.dataset)
197+
if full_data.static_node_x is None:
198+
full_data.static_node_x = torch.randn((full_data.num_nodes, 1), device=args.device)
199+
200+
train_data, val_data, test_data = full_data.split()
201+
train_dg = DGraph(train_data, device=args.device)
202+
val_dg = DGraph(val_data, device=args.device)
203+
test_dg = DGraph(test_data, device=args.device)
204+
205+
206+
def compute_delta_t_stats(train_dg: DGraph) -> Tuple[float, float]:
207+
last_timestamp = {}
208+
delta_times = []
209+
210+
for src, dst, t in zip(train_dg.edge_src, train_dg.edge_dst, train_dg.edge_time):
211+
src, dst, t = src.item(), dst.item(), t.item()
212+
213+
dt_src = t - last_timestamp.get(src, train_dg.start_time)
214+
dt_dst = t - last_timestamp.get(dst, train_dg.start_time)
215+
delta_times.extend([dt_src, dt_dst])
216+
217+
last_timestamp[src] = t
218+
last_timestamp[dst] = t
219+
220+
return np.mean(delta_times), np.std(delta_times)
221+
222+
223+
mean_delta_t, std_delta_t = compute_delta_t_stats(train_dg)
224+
225+
nbr_hook = RecencyNeighborHook(
226+
num_nbrs=args.n_nbrs,
227+
num_nodes=full_data.num_nodes,
228+
seed_nodes_keys=['edge_src', 'edge_dst', 'neg'],
229+
seed_times_keys=['edge_time', 'edge_time', 'neg_time'],
230+
)
231+
232+
hm = RecipeRegistry.build(
233+
RECIPE_TGB_LINK_PRED, dataset_name=args.dataset, train_dg=train_dg
234+
)
235+
train_key, val_key, test_key = hm.keys
236+
hm.register_shared(nbr_hook)
237+
hm.register_shared(DeduplicationHook())
238+
239+
train_loader = DGDataLoader(train_dg, args.bsize, hook_manager=hm)
240+
val_loader = DGDataLoader(val_dg, args.bsize, hook_manager=hm)
241+
test_loader = DGDataLoader(test_dg, args.bsize, hook_manager=hm)
242+
243+
memory = CTANMemory(
244+
num_nodes=test_dg.num_nodes,
245+
memory_dim=args.memory_dim,
246+
aggr_module=LastAggregator(),
247+
init_time=train_dg.start_time,
248+
).to(args.device)
249+
encoder = CTAN(
250+
node_dim=train_dg.static_node_x_dim,
251+
edge_dim=train_dg.edge_x_dim,
252+
time_dim=args.time_dim,
253+
memory_dim=args.memory_dim,
254+
num_iters=args.n_layers,
255+
mean_delta_t=mean_delta_t,
256+
std_delta_t=std_delta_t,
257+
epsilon=args.epsilon,
258+
gamma=args.gamma,
259+
).to(args.device)
260+
decoder = LinkPredictor(node_dim=args.memory_dim, merge_op='sum').to(args.device)
261+
opt = torch.optim.Adam(
262+
set(memory.parameters()) | set(encoder.parameters()) | set(decoder.parameters()),
263+
lr=args.lr,
264+
)
265+
266+
for epoch in range(1, args.epochs + 1):
267+
with hm.activate(train_key):
268+
loss = train(train_loader, memory, encoder, decoder, opt)
269+
270+
with hm.activate(val_key):
271+
val_mrr = eval(val_loader, memory, encoder, decoder, evaluator)
272+
log_metric('Loss', loss, epoch=epoch)
273+
log_metric(f'Validation {METRIC_TGB_LINKPROPPRED}', val_mrr, epoch=epoch)
274+
275+
if epoch < args.epochs: # Reset hooks after each epoch, except last epoch
276+
hm.reset_state()
277+
278+
with hm.activate(test_key):
279+
test_mrr = eval(test_loader, memory, encoder, decoder, evaluator)
280+
log_metric(f'Test {METRIC_TGB_LINKPROPPRED}', test_mrr, epoch=args.epochs)

test/integration/test_ctan.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
import pytest
2+
3+
4+
@pytest.mark.integration
5+
@pytest.mark.parametrize('dataset', ['tgbl-wiki'])
6+
@pytest.mark.slurm(
7+
resources=[
8+
'--partition=main',
9+
'--cpus-per-task=2',
10+
'--mem=8G',
11+
'--time=1:00:00',
12+
'--gres=gpu:a100l:1',
13+
]
14+
)
15+
def test_ctan_linkprop_pred(slurm_job_runner, dataset):
16+
cmd = f"""
17+
python "$ROOT_DIR/examples/linkproppred/ctan.py" \
18+
--dataset {dataset} \
19+
--device cuda \
20+
--epochs 1"""
21+
state = slurm_job_runner(cmd)
22+
assert state == 'COMPLETED'

test/unit/test_nn/test_ctan.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
import torch
2+
3+
from tgm.nn.encoder import CTAN, CTANMemory, LastAggregator
4+
5+
6+
def test_ctan_last_aggre():
7+
B, S = 10, 1
8+
E, M, T = 7, 5, 2
9+
10+
edge_index = torch.randint(0, B, size=(2, B))
11+
edge_time = torch.randint(0, B, size=(B,))
12+
edge_feat = torch.randint(0, B, size=(B, E))
13+
memory = CTANMemory(
14+
B,
15+
M,
16+
aggr_module=LastAggregator(),
17+
)
18+
encoder = CTAN(
19+
edge_dim=E,
20+
memory_dim=M,
21+
time_dim=T,
22+
node_dim=S,
23+
)
24+
memory.train()
25+
encoder.train()
26+
z, last_update = memory(torch.unique(edge_index))
27+
z = torch.cat([z, torch.rand((len(z), 1))], dim=-1) # Dummy random node feats
28+
z = encoder(z, last_update, edge_index, edge_time, edge_feat)
29+
memory.detach()
30+
memory.reset_parameters()
31+
32+
assert z.shape == (B, M)
33+
assert not torch.isnan(z).any()
34+
35+
memory.eval()
36+
encoder.eval()
37+
z, last_update = memory(torch.unique(edge_index))
38+
z = torch.cat([z, torch.rand((len(z), 1))], dim=-1) # Dummy random node feats
39+
z = encoder(z, last_update, edge_index, edge_time, edge_feat)
40+
41+
memory.update_state(
42+
src=edge_index[0],
43+
pos_dst=edge_index[1],
44+
t=edge_time,
45+
src_emb=z[edge_index[0]],
46+
pos_dst_emb=z[edge_index[1]],
47+
)
48+
memory.detach()
49+
50+
assert z.shape == (B, M)
51+
assert not torch.isnan(z).any()

0 commit comments

Comments
 (0)