|
| 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) |
0 commit comments