Skip to content

Commit 72c8bf9

Browse files
ntgbaooBảo Ngô
andauthored
Added tkgl i/o (#371)
* Added tkgl i/o * added edgebank example for tkgl * Updated test * Fixed bug on edge_event_idx from DGData * Fixed unit tests failed * Updated unit test for node event and egde event idx * Cleaned up * Updated time granularities * Updated unit test for time * Pump timeout for edgebank for tkgl --------- Co-authored-by: Bảo Ngô <ngot1@myumanitoba.ca>
1 parent 7691642 commit 72c8bf9

11 files changed

Lines changed: 636 additions & 12 deletions

File tree

Lines changed: 119 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,119 @@
1+
import argparse
2+
3+
import numpy as np
4+
import torch
5+
from tgb.linkproppred.evaluate import Evaluator
6+
from tqdm import tqdm
7+
8+
from tgm import DGraph
9+
from tgm.constants import METRIC_TGB_LINKPROPPRED
10+
from tgm.data import DGData, DGDataLoader
11+
from tgm.hooks import HookManager, TGBTKGNegativeEdgeSamplerHook
12+
from tgm.nn import EdgeBankPredictor
13+
from tgm.util.logging import enable_logging, log_latency, log_metric
14+
from tgm.util.seed import seed_everything
15+
16+
parser = argparse.ArgumentParser(
17+
description='EdgeBank LinkPropPred Example for knowledge graph',
18+
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
19+
)
20+
parser.add_argument('--seed', type=int, default=1337, help='random seed to use')
21+
parser.add_argument(
22+
'--dataset', type=str, default='tkgl-smallpedia', help='Dataset name'
23+
)
24+
parser.add_argument('--bsize', type=int, default=200, help='batch size')
25+
parser.add_argument('--window-ratio', type=float, default=0.15, help='Window ratio')
26+
parser.add_argument('--pos-prob', type=float, default=1.0, help='Positive edge prob')
27+
parser.add_argument(
28+
'--memory-mode',
29+
type=str,
30+
default='unlimited',
31+
choices=['unlimited', 'fixed'],
32+
help='Memory mode',
33+
)
34+
parser.add_argument(
35+
'--log-file-path', type=str, default=None, help='Optional path to write logs'
36+
)
37+
38+
args = parser.parse_args()
39+
enable_logging(log_file_path=args.log_file_path)
40+
41+
42+
@log_latency
43+
def eval(
44+
loader: DGDataLoader,
45+
model: EdgeBankPredictor,
46+
evaluator: Evaluator,
47+
) -> float:
48+
perf_list = []
49+
for batch in tqdm(loader):
50+
for idx, neg_batch in enumerate(batch.neg_batch_list):
51+
query_src = batch.edge_src[idx].repeat(len(neg_batch) + 1)
52+
query_dst = torch.cat([batch.edge_dst[idx].unsqueeze(0), neg_batch])
53+
54+
y_pred = model(query_src, query_dst)
55+
input_dict = {
56+
'y_pred_pos': y_pred[0],
57+
'y_pred_neg': y_pred[1:],
58+
'eval_metric': [METRIC_TGB_LINKPROPPRED],
59+
}
60+
perf_list.append(evaluator.eval(input_dict)[METRIC_TGB_LINKPROPPRED])
61+
model.update(batch.edge_src, batch.edge_dst, batch.edge_time)
62+
63+
return float(np.mean(perf_list))
64+
65+
66+
seed_everything(args.seed)
67+
evaluator = Evaluator(name=args.dataset)
68+
69+
data = DGData.from_tgb(args.dataset)
70+
min_dst_node = data.edge_index[:, 1].min().int()
71+
max_dst_node = data.edge_index[:, 1].max().int()
72+
73+
train_data, val_data, test_data = data.split()
74+
train_dg = DGraph(train_data)
75+
val_dg = DGraph(val_data)
76+
test_dg = DGraph(test_data)
77+
78+
train_data = train_dg.materialize(materialize_features=False)
79+
80+
81+
hm = HookManager(keys=['val', 'test'])
82+
hm.register(
83+
'val',
84+
TGBTKGNegativeEdgeSamplerHook(
85+
args.dataset,
86+
split_mode='val',
87+
first_dst_id=min_dst_node,
88+
last_dst_id=max_dst_node,
89+
),
90+
)
91+
hm.register(
92+
'test',
93+
TGBTKGNegativeEdgeSamplerHook(
94+
args.dataset,
95+
split_mode='test',
96+
first_dst_id=min_dst_node,
97+
last_dst_id=max_dst_node,
98+
),
99+
)
100+
101+
val_loader = DGDataLoader(val_dg, args.bsize, hook_manager=hm)
102+
test_loader = DGDataLoader(test_dg, args.bsize, hook_manager=hm)
103+
104+
model = EdgeBankPredictor(
105+
train_data.edge_src,
106+
train_data.edge_dst,
107+
train_data.edge_time,
108+
memory_mode=args.memory_mode,
109+
window_ratio=args.window_ratio,
110+
pos_prob=args.pos_prob,
111+
)
112+
113+
with hm.activate('val'):
114+
val_mrr = eval(val_loader, model, evaluator)
115+
log_metric(f'Validation {METRIC_TGB_LINKPROPPRED}', val_mrr)
116+
117+
with hm.activate('test'):
118+
test_mrr = eval(test_loader, model, evaluator)
119+
log_metric(f'Test {METRIC_TGB_LINKPROPPRED}', test_mrr)

scripts/download_tgb_datasets.sh

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ DATASETS=(
1111
"tgbl_wiki"
1212
"tgbn_trade"
1313
"thgl_software"
14+
"tkgl-smallpedia"
1415
#"tgbn_genre"
1516
#"tgbl_coin"
1617
#"tgbl_flight" TODO: Start working with the large graphs
@@ -67,7 +68,7 @@ download_dataset() {
6768
local dataset_name="${dataset//_/-}" # 'tgbl_wiki' -> 'tgbl-wiki'
6869
echo "Downloading dataset: $dataset_name"
6970

70-
if [[ "$dataset" == tgbl_* || "$dataset" == thgl_* ]]; then
71+
if [[ "$dataset" == tgbl_* || "$dataset" == thgl_* || "$dataset" == tkgl_* ]]; then
7172
.venv/bin/python -c "from tgb.linkproppred.dataset import LinkPropPredDataset as DS; DS(name='$dataset_name')"
7273
elif [[ "$dataset" == tgbn_* ]]; then
7374
.venv/bin/python -c "from tgb.nodeproppred.dataset import NodePropPredDataset as DS; DS(name='$dataset_name')"

test/integration/test_edgebank.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,3 +89,39 @@ def test_edgebank_linkprop_pred_fixed_memory_thgl(slurm_job_runner, dataset):
8989
--dataset {dataset} --memory-mode fixed"""
9090
state = slurm_job_runner(cmd)
9191
assert state == 'COMPLETED'
92+
93+
94+
@pytest.mark.integration
95+
@pytest.mark.parametrize('dataset', ['tkgl-smallpedia'])
96+
@pytest.mark.slurm(
97+
resources=[
98+
'--partition=main',
99+
'--cpus-per-task=2',
100+
'--mem=8G',
101+
'--time=1:15:00',
102+
]
103+
)
104+
def test_edgebank_linkprop_pred_unlimited_memory_tkgl(slurm_job_runner, dataset):
105+
cmd = f"""
106+
python "$ROOT_DIR/examples/linkproppred/tkgl/edgebank.py" \
107+
--dataset {dataset}"""
108+
state = slurm_job_runner(cmd)
109+
assert state == 'COMPLETED'
110+
111+
112+
@pytest.mark.integration
113+
@pytest.mark.parametrize('dataset', ['tkgl-smallpedia'])
114+
@pytest.mark.slurm(
115+
resources=[
116+
'--partition=main',
117+
'--cpus-per-task=2',
118+
'--mem=8G',
119+
'--time=2:00:00',
120+
]
121+
)
122+
def test_edgebank_linkprop_pred_fixed_memory_tkgl(slurm_job_runner, dataset):
123+
cmd = f"""
124+
python "$ROOT_DIR/examples/linkproppred/tkgl/edgebank.py" \
125+
--dataset {dataset} --memory-mode fixed"""
126+
state = slurm_job_runner(cmd)
127+
assert state == 'COMPLETED'

test/unit/test_core/test_timedelta.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -192,6 +192,10 @@ def test_tgb_native_time_deltas():
192192
'thgl-forum': TimeDeltaDG('s'),
193193
'thgl-github': TimeDeltaDG('s'),
194194
'thgl-myket': TimeDeltaDG('s'),
195+
'tkgl-smallpedia': TimeDeltaDG('Y'),
196+
'tkgl-polecat': TimeDeltaDG('D'),
197+
'tkgl-icews': TimeDeltaDG('D'),
198+
'tkgl-wikidata': TimeDeltaDG('Y'),
195199
}
196200
assert TGB_TIME_DELTAS == exp_dict
197201

test/unit/test_data/test_data.py

Lines changed: 124 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1196,7 +1196,13 @@ def test_from_pandas_bad_node_cols_not_specified():
11961196

11971197
@pytest.fixture
11981198
def tgb_dataset_factory():
1199-
def _make_dataset(split: str = 'all', with_node_feats: bool = False, thgl=False):
1199+
def _make_dataset(
1200+
split: str = 'all',
1201+
with_node_feats: bool = False,
1202+
with_edge_feats: bool = False,
1203+
thgl=False,
1204+
tkgl=False,
1205+
):
12001206
num_events, num_train, num_val = 10, 7, 2
12011207
train_indices = np.arange(0, num_train)
12021208
val_indices = np.arange(num_train, num_train + num_val)
@@ -1205,7 +1211,12 @@ def _make_dataset(split: str = 'all', with_node_feats: bool = False, thgl=False)
12051211
sources = np.random.randint(0, 1000, size=num_events)
12061212
destinations = np.random.randint(0, 1000, size=num_events)
12071213
timestamps = np.arange(num_events)
1208-
edge_feat = None
1214+
if not with_edge_feats:
1215+
edge_feat = None
1216+
elif with_edge_feats and not tkgl:
1217+
edge_feat = np.random.rand(num_events, 10)
1218+
elif with_edge_feats and tkgl:
1219+
edge_feat = np.random.rand(num_events // 2, 10)
12091220

12101221
train_mask = np.zeros(num_events, dtype=bool)
12111222
val_mask = np.zeros(num_events, dtype=bool)
@@ -1249,6 +1260,9 @@ def _make_dataset(split: str = 'all', with_node_feats: bool = False, thgl=False)
12491260
max(sources.max(), destinations.max()) + 1
12501261
)
12511262

1263+
if tkgl:
1264+
mock_dataset.full_data['edge_type'] = np.arange(num_events)
1265+
12521266
return mock_dataset
12531267

12541268
return _make_dataset
@@ -1316,6 +1330,59 @@ def _make_dataset(split: str = 'all', with_edge_type=True, with_node_type=False)
13161330
return _make_dataset
13171331

13181332

1333+
@pytest.fixture
1334+
def bad_tkgl_dataset_factory(): # Missing edge_type
1335+
def _make_dataset(split: str = 'all'):
1336+
num_events, num_train, num_val = 10, 7, 2
1337+
train_indices = np.arange(0, num_train)
1338+
val_indices = np.arange(num_train, num_train + num_val)
1339+
test_indices = np.arange(num_train + num_val, num_events)
1340+
1341+
sources = np.random.randint(0, 1000, size=num_events)
1342+
destinations = np.random.randint(0, 1000, size=num_events)
1343+
timestamps = np.arange(num_events)
1344+
edge_feat = None
1345+
w = np.random.rand(num_events, 10)
1346+
1347+
train_mask = np.zeros(num_events, dtype=bool)
1348+
val_mask = np.zeros(num_events, dtype=bool)
1349+
test_mask = np.zeros(num_events, dtype=bool)
1350+
1351+
train_mask[train_indices] = True
1352+
val_mask[val_indices] = True
1353+
test_mask[test_indices] = True
1354+
1355+
mock_dataset = MagicMock()
1356+
mock_dataset.train_mask = train_mask
1357+
mock_dataset.val_mask = val_mask
1358+
mock_dataset.test_mask = test_mask
1359+
mock_dataset.num_edges = num_events
1360+
mock_dataset.full_data = {
1361+
'sources': sources,
1362+
'destinations': destinations,
1363+
'timestamps': timestamps,
1364+
'edge_feat': edge_feat,
1365+
'w': w,
1366+
}
1367+
1368+
if split == 'all':
1369+
1 + max(np.max(sources), np.max(destinations))
1370+
else:
1371+
mask = {'train': train_mask, 'val': val_mask, 'test': test_mask}[split]
1372+
valid_src, valid_dst = sources[mask], destinations[mask]
1373+
1 + max(np.max(valid_src), np.max(valid_dst))
1374+
1375+
mock_dataset.node_feat = None
1376+
1377+
mock_dataset.full_data['node_label_dict'] = {}
1378+
for i in range(5):
1379+
mock_dataset.full_data['node_label_dict'][i] = {i: np.zeros(10)}
1380+
1381+
return mock_dataset
1382+
1383+
return _make_dataset
1384+
1385+
13191386
@pytest.fixture
13201387
def tgb_seq_dataset_factory():
13211388
def _make_dataset(
@@ -1369,11 +1436,6 @@ def _make_dataset(
13691436
return _make_dataset
13701437

13711438

1372-
def test_from_tkgl():
1373-
with pytest.raises(NotImplementedError):
1374-
DGData.from_tgb('tkgl-foo')
1375-
1376-
13771439
def test_from_bad_tgb_name():
13781440
with pytest.raises(ValueError):
13791441
DGData.from_tgb('foo')
@@ -2258,3 +2320,58 @@ def test_from_pandas_with_static_node_type():
22582320
)
22592321
assert isinstance(data, DGData)
22602322
torch.testing.assert_close(data.node_type.tolist(), node_dict['node_type'])
2323+
2324+
2325+
@pytest.mark.parametrize('with_node_feats', [True, False])
2326+
@pytest.mark.parametrize('with_edge_feats', [True, False])
2327+
@pytest.mark.parametrize('tkgl', [True])
2328+
@patch('tgb.linkproppred.dataset.LinkPropPredDataset')
2329+
@patch.dict('tgm.core.timedelta.TGB_TIME_DELTAS', {'tkgl-smallpedia': TimeDeltaDG('D')})
2330+
def test_from_tkgl(
2331+
mock_dataset_cls, tgb_dataset_factory, with_node_feats, with_edge_feats, tkgl
2332+
):
2333+
dataset = tgb_dataset_factory(
2334+
with_node_feats=with_node_feats, with_edge_feats=with_edge_feats, tkgl=tkgl
2335+
)
2336+
mock_dataset_cls.return_value = dataset
2337+
2338+
mock_native_time_delta = TimeDeltaDG('D') # Patched value
2339+
2340+
def _get_exp_edges():
2341+
src, dst = dataset.full_data['sources'], dataset.full_data['destinations']
2342+
return np.stack([src, dst], axis=1)
2343+
2344+
def _get_exp_times():
2345+
return dataset.full_data['timestamps']
2346+
2347+
def _get_exp_edge_type():
2348+
return dataset.full_data['edge_type']
2349+
2350+
def _get_exp_edge_feat():
2351+
edge_feat_np = dataset.full_data['edge_feat']
2352+
return np.concatenate((edge_feat_np, edge_feat_np))
2353+
2354+
data = DGData.from_tgb(name='tkgl-smallpedia')
2355+
assert isinstance(data, DGData)
2356+
assert data.time_delta == mock_native_time_delta
2357+
np.testing.assert_allclose(data.edge_index.numpy(), _get_exp_edges())
2358+
np.testing.assert_allclose(data.time.numpy(), _get_exp_times())
2359+
np.testing.assert_allclose(data.edge_type.numpy(), _get_exp_edge_type())
2360+
if with_edge_feats:
2361+
np.testing.assert_allclose(data.edge_x.numpy(), _get_exp_edge_feat())
2362+
2363+
# Confirm correct dataset instantiation
2364+
mock_dataset_cls.assert_called_once_with(name='tkgl-smallpedia')
2365+
2366+
if with_node_feats:
2367+
torch.testing.assert_close(data.static_node_x, torch.Tensor(dataset.node_feat))
2368+
else:
2369+
assert data.static_node_x is None
2370+
2371+
2372+
@patch('tgb.linkproppred.dataset.LinkPropPredDataset')
2373+
def test_from_bad_thgl(mock_dataset_cls, bad_tkgl_dataset_factory):
2374+
dataset = bad_tkgl_dataset_factory()
2375+
mock_dataset_cls.return_value = dataset
2376+
with pytest.raises(ValueError):
2377+
data = DGData.from_tgb(name='tkgl-smallpedia')

test/unit/test_hooks/test_device_transfer_hook.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ def test_device_transfer_hook_gpu_gpu(dg):
6161
batch = dg.materialize()
6262
batch.edge_src = batch.edge_src.to('cuda')
6363
batch.edge_dst = batch.edge_dst.to('cuda')
64-
batch.time = batch.time.to('cuda')
64+
batch.edge_time = batch.edge_time.to('cuda')
6565

6666
# Add a custom field and ensure it's also moved
6767
batch.foo = torch.rand(1, 2, device='cuda')

0 commit comments

Comments
 (0)