Skip to content
12 changes: 6 additions & 6 deletions data_loading.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,8 @@
import torch
import logging
import itertools
from data_util import GraphData, HeteroData, z_norm, create_hetero_obj
from data_util import GraphData, HeteroData, create_hetero_obj
from edge_feature_utils import get_non_normalized_edge_feature_indices, z_norm, z_norm_except

def get_data(args, data_config):
'''Loads the AML transaction data.
Expand Down Expand Up @@ -121,10 +122,10 @@ def get_data(args, data_config):

#Normalize data
tr_data.x = val_data.x = te_data.x = z_norm(tr_data.x)
if not args.model == 'rgcn':
tr_data.edge_attr, val_data.edge_attr, te_data.edge_attr = z_norm(tr_data.edge_attr), z_norm(val_data.edge_attr), z_norm(te_data.edge_attr)
else:
tr_data.edge_attr[:, :-1], val_data.edge_attr[:, :-1], te_data.edge_attr[:, :-1] = z_norm(tr_data.edge_attr[:, :-1]), z_norm(val_data.edge_attr[:, :-1]), z_norm(te_data.edge_attr[:, :-1])
excluded_edge_feature_indices = get_non_normalized_edge_feature_indices(args.model)
tr_data.edge_attr = z_norm_except(tr_data.edge_attr, excluded_edge_feature_indices)
val_data.edge_attr = z_norm_except(val_data.edge_attr, excluded_edge_feature_indices)
te_data.edge_attr = z_norm_except(te_data.edge_attr, excluded_edge_feature_indices)

#Create heterogenous if reverese MP is enabled
#TODO: if I observe wierd behaviour, maybe add .detach.clone() to all torch tensors, but I don't think they're attached to any computation graph just yet
Expand All @@ -138,4 +139,3 @@ def get_data(args, data_config):
logging.info(f'test data object: {te_data}')

return tr_data, val_data, te_data, tr_inds, val_inds, te_inds

16 changes: 7 additions & 9 deletions data_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from torch_geometric.data import Data, HeteroData
from torch_geometric.typing import OptTensor
import numpy as np
from edge_feature_utils import get_port_feature_indices

def to_adj_nodes_with_times(data):
num_nodes = data.num_nodes
Expand Down Expand Up @@ -133,24 +134,21 @@ def add_time_deltas(self):
self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_tds], dim=1)
return self

def z_norm(data):
std = data.std(0).unsqueeze(0)
std = torch.where(std == 0, torch.tensor(1, dtype=torch.float32).cpu(), std)
return (data - data.mean(0).unsqueeze(0)) / std

def create_hetero_obj(x, y, edge_index, edge_attr, timestamps, args):
'''Creates a heterogenous graph object for reverse message passing'''
data = HeteroGraphData()

data['node'].x = x
data['node', 'to', 'node'].edge_index = edge_index
data['node', 'rev_to', 'node'].edge_index = edge_index.flipud()
data['node', 'to', 'node'].edge_attr = edge_attr
data['node', 'rev_to', 'node'].edge_attr = edge_attr
data['node', 'to', 'node'].edge_attr = edge_attr.clone()
data['node', 'rev_to', 'node'].edge_attr = edge_attr.clone()
if args.ports:
#swap the in- and outgoing port numberings for the reverse edges
data['node', 'rev_to', 'node'].edge_attr[:, [-1, -2]] = data['node', 'rev_to', 'node'].edge_attr[:, [-2, -1]]
port_feature_indices = get_port_feature_indices(args.ports)
left_idx, right_idx = port_feature_indices
data['node', 'rev_to', 'node'].edge_attr[:, [left_idx, right_idx]] = data['node', 'rev_to', 'node'].edge_attr[:, [right_idx, left_idx]]
data['node', 'to', 'node'].y = y
data['node', 'to', 'node'].timestamps = timestamps

return data
return data
50 changes: 50 additions & 0 deletions edge_feature_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
import torch

BASE_EDGE_FEATURE_COUNT = 4
PAYMENT_FORMAT_INDEX = 3


def average_residual_update(current, update):
return (current + update) / 2


def get_edge_type_index():
return PAYMENT_FORMAT_INDEX


def get_non_normalized_edge_feature_indices(model_name):
return (PAYMENT_FORMAT_INDEX,) if model_name == "rgcn" else ()


def get_port_feature_indices(use_ports):
if not use_ports:
return ()

return tuple(range(BASE_EDGE_FEATURE_COUNT, BASE_EDGE_FEATURE_COUNT + 2))


def get_time_delta_feature_indices(use_ports, use_tds):
if not use_tds:
return ()

start_idx = BASE_EDGE_FEATURE_COUNT + (2 if use_ports else 0)
return tuple(range(start_idx, start_idx + 2))


def z_norm(data):
std = data.std(0).unsqueeze(0)
std = torch.where(std == 0, torch.ones_like(std), std)
return (data - data.mean(0).unsqueeze(0)) / std


def z_norm_except(data, excluded_indices):
if not excluded_indices:
return z_norm(data)

normalized = data.clone()
excluded_indices = set(excluded_indices)
include_indices = [idx for idx in range(data.shape[1]) if idx not in excluded_indices]
if include_indices:
normalized[:, include_indices] = z_norm(data[:, include_indices])

return normalized
26 changes: 14 additions & 12 deletions models.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import torch.nn.functional as F
import torch
import logging
from edge_feature_utils import average_residual_update

class GINe(torch.nn.Module):
def __init__(self, num_features, num_gnn_layers, n_classes=2,
Expand Down Expand Up @@ -44,9 +45,9 @@ def forward(self, x, edge_index, edge_attr):
edge_attr = self.edge_emb(edge_attr)

for i in range(self.num_gnn_layers):
x = (x + F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr)))) / 2
x = average_residual_update(x, F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr))))
if self.edge_updates:
edge_attr = edge_attr + self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)) / 2
edge_attr = average_residual_update(edge_attr, self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)))

x = x[edge_index.T].reshape(-1, 2 * self.n_hidden).relu()
x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1)
Expand Down Expand Up @@ -90,9 +91,9 @@ def forward(self, x, edge_index, edge_attr):
edge_attr = self.edge_emb(edge_attr)

for i in range(self.num_gnn_layers):
x = (x + F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr)))) / 2
x = average_residual_update(x, F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr))))
if self.edge_updates:
edge_attr = edge_attr + self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)) / 2
edge_attr = average_residual_update(edge_attr, self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)))

logging.debug(f"x.shape = {x.shape}, x[edge_index.T].shape = {x[edge_index.T].shape}")
x = x[edge_index.T].reshape(-1, 2 * self.n_hidden).relu()
Expand Down Expand Up @@ -146,9 +147,9 @@ def forward(self, x, edge_index, edge_attr):
edge_attr = self.edge_emb(edge_attr)

for i in range(self.num_gnn_layers):
x = (x + F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr)))) / 2
x = average_residual_update(x, F.relu(self.batch_norms[i](self.convs[i](x, edge_index, edge_attr))))
if self.edge_updates:
edge_attr = edge_attr + self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)) / 2
edge_attr = average_residual_update(edge_attr, self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)))

logging.debug(f"x.shape = {x.shape}, x[edge_index.T].shape = {x[edge_index.T].shape}")
x = x[edge_index.T].reshape(-1, 2 * self.n_hidden).relu()
Expand All @@ -162,7 +163,7 @@ class RGCN(nn.Module):
def __init__(self, num_features, edge_dim, num_relations, num_gnn_layers, n_classes=2,
n_hidden=100, edge_update=False,
residual=True,
dropout=0.0, final_dropout=0.5, n_bases=-1):
dropout=0.0, final_dropout=0.5, n_bases=-1, edge_type_index=3):
super(RGCN, self).__init__()

self.num_features = num_features
Expand All @@ -175,6 +176,7 @@ def __init__(self, num_features, edge_dim, num_relations, num_gnn_layers, n_clas
self.edge_update = edge_update
self.num_relations = num_relations
self.n_bases = n_bases
self.edge_type_index = edge_type_index

self.node_emb = nn.Linear(num_features, n_hidden)
self.edge_emb = nn.Linear(edge_dim, n_hidden)
Expand Down Expand Up @@ -216,21 +218,21 @@ def reset_parameters(self):
m.reset_parameters()

def forward(self, x, edge_index, edge_attr):
edge_type = edge_attr[:, -1].long()
#edge_attr = edge_attr[:, :-1]
edge_type = edge_attr[:, self.edge_type_index].long()
edge_attr = torch.cat((edge_attr[:, :self.edge_type_index], edge_attr[:, self.edge_type_index + 1:]), dim=-1)
src, dst = edge_index

x = self.node_emb(x)
edge_attr = self.edge_emb(edge_attr)

for i in range(self.num_gnn_layers):
x = (x + F.relu(self.bns[i](self.convs[i](x, edge_index, edge_type)))) / 2
x = average_residual_update(x, F.relu(self.bns[i](self.convs[i](x, edge_index, edge_type))))
if self.edge_update:
edge_attr = (edge_attr + F.relu(self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1)))) / 2
edge_attr = average_residual_update(edge_attr, F.relu(self.emlps[i](torch.cat([x[src], x[dst], edge_attr], dim=-1))))

x = x[edge_index.T].reshape(-1, 2 * self.n_hidden).relu()
x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1)
x = self.mlp(x)
out = x

return x
return x
127 changes: 127 additions & 0 deletions tests/test_edge_feature_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
import importlib.util
import sys
import types
import unittest
from types import SimpleNamespace

import torch

from edge_feature_utils import (
average_residual_update,
get_edge_type_index,
get_non_normalized_edge_feature_indices,
get_port_feature_indices,
get_time_delta_feature_indices,
z_norm_except,
)


if importlib.util.find_spec("torch_geometric") is None:
fake_data_module = types.ModuleType("torch_geometric.data")
fake_typing_module = types.ModuleType("torch_geometric.typing")

class _FakeData:
def __init__(self, *args, **kwargs):
pass

class _FakeStore:
pass

class _FakeHeteroData(dict):
def __getitem__(self, key):
if key not in self:
dict.__setitem__(self, key, _FakeStore())
return dict.__getitem__(self, key)

fake_data_module.Data = _FakeData
fake_data_module.HeteroData = _FakeHeteroData
fake_typing_module.OptTensor = object

sys.modules.setdefault("torch_geometric", types.ModuleType("torch_geometric"))
sys.modules["torch_geometric.data"] = fake_data_module
sys.modules["torch_geometric.typing"] = fake_typing_module

from data_util import create_hetero_obj


class EdgeFeatureUtilsTests(unittest.TestCase):
def test_average_residual_update_returns_mean(self):
current = torch.tensor([[2.0, 4.0]])
update = torch.tensor([[6.0, 8.0]])

result = average_residual_update(current, update)

self.assertTrue(torch.equal(result, torch.tensor([[4.0, 6.0]])))

def test_rgcn_keeps_payment_format_column_unnormalized(self):
edge_attr = torch.tensor(
[
[1.0, 10.0, 100.0, 0.0, 7.0, 70.0, 700.0, 7000.0],
[2.0, 20.0, 200.0, 1.0, 8.0, 80.0, 800.0, 8000.0],
[3.0, 30.0, 300.0, 2.0, 9.0, 90.0, 900.0, 9000.0],
]
)

normalized = z_norm_except(edge_attr, get_non_normalized_edge_feature_indices("rgcn"))

self.assertTrue(torch.equal(normalized[:, get_edge_type_index()], edge_attr[:, get_edge_type_index()]))
self.assertAlmostEqual(float(normalized[:, 0].mean()), 0.0, places=6)
self.assertAlmostEqual(float(normalized[:, -1].mean()), 0.0, places=6)

def test_feature_indices_stay_stable_when_ports_and_tds_are_enabled(self):
self.assertEqual(get_port_feature_indices(True), (4, 5))
self.assertEqual(get_time_delta_feature_indices(True, True), (6, 7))

def test_create_hetero_obj_swaps_only_port_columns_when_ports_and_tds_are_enabled(self):
edge_attr = torch.tensor(
[
[1.0, 10.0, 100.0, 0.0, 7.0, 8.0, 70.0, 80.0],
[2.0, 20.0, 200.0, 1.0, 9.0, 10.0, 90.0, 100.0],
]
)
data = create_hetero_obj(
x=torch.ones((3, 1)),
y=torch.tensor([0, 1]),
edge_index=torch.tensor([[0, 1], [1, 2]]),
edge_attr=edge_attr,
timestamps=edge_attr[:, 0],
args=SimpleNamespace(ports=True, tds=True),
)

port_indices = list(get_port_feature_indices(True))
td_indices = list(get_time_delta_feature_indices(True, True))
expected_reverse = edge_attr.clone()
expected_reverse[:, port_indices] = expected_reverse[:, list(reversed(port_indices))]

self.assertTrue(torch.equal(data["node", "to", "node"].edge_attr, edge_attr))
self.assertTrue(torch.equal(data["node", "rev_to", "node"].edge_attr, expected_reverse))
self.assertTrue(
torch.equal(
data["node", "rev_to", "node"].edge_attr[:, td_indices],
edge_attr[:, td_indices],
)
)

def test_create_hetero_obj_clones_forward_and_reverse_edge_attributes(self):
edge_attr = torch.tensor(
[
[1.0, 10.0, 100.0, 0.0, 7.0, 8.0],
[2.0, 20.0, 200.0, 1.0, 9.0, 10.0],
]
)
data = create_hetero_obj(
x=torch.ones((3, 1)),
y=torch.tensor([0, 1]),
edge_index=torch.tensor([[0, 1], [1, 2]]),
edge_attr=edge_attr,
timestamps=edge_attr[:, 0],
args=SimpleNamespace(ports=True, tds=False),
)

data["node", "rev_to", "node"].edge_attr[:, 0] = -1

self.assertTrue(torch.equal(data["node", "to", "node"].edge_attr, edge_attr))


if __name__ == "__main__":
unittest.main()
6 changes: 4 additions & 2 deletions training.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from torch_geometric.utils import degree
import wandb
import logging
from edge_feature_utils import get_edge_type_index

def train_homo(tr_loader, val_loader, te_loader, tr_inds, val_inds, te_inds, model, optimizer, loss_fn, args, config, device, val_data, te_data, data_config):
#training
Expand Down Expand Up @@ -158,7 +159,8 @@ def get_model(sample_batch, config, args):
model = RGCN(
num_features=n_feats, edge_dim=e_dim, num_relations=8, num_gnn_layers=round(config.n_gnn_layers),
n_classes=2, n_hidden=round(config.n_hidden),
edge_update=args.emlps, dropout=config.dropout, final_dropout=config.final_dropout, n_bases=None #(maybe)
edge_update=args.emlps, dropout=config.dropout, final_dropout=config.final_dropout, n_bases=None,
edge_type_index=get_edge_type_index() #(maybe)
)

return model
Expand Down Expand Up @@ -234,4 +236,4 @@ def train_gnn(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, args, data
else:
model = train_homo(tr_loader, val_loader, te_loader, tr_inds, val_inds, te_inds, model, optimizer, loss_fn, args, config, device, val_data, te_data, data_config)

wandb.finish()
wandb.finish()