From 81c34ed4fc64f418d28bd8d4ffeae72c52bcb993 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:39:41 +0800 Subject: [PATCH 01/10] Update edge_feature_utils.py for edge feature fix Signed-off-by: benben951 --- edge_feature_utils.py | 50 +++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) create mode 100644 edge_feature_utils.py diff --git a/edge_feature_utils.py b/edge_feature_utils.py new file mode 100644 index 0000000..b57a738 --- /dev/null +++ b/edge_feature_utils.py @@ -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 From f411c03f3892e2b538323c6097f8a1ece8cb7185 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:41:11 +0800 Subject: [PATCH 02/10] Update tests/test_edge_feature_utils.py for edge feature fix Signed-off-by: benben951 --- tests/test_edge_feature_utils.py | 45 ++++++++++++++++++++++++++++++++ 1 file changed, 45 insertions(+) create mode 100644 tests/test_edge_feature_utils.py diff --git a/tests/test_edge_feature_utils.py b/tests/test_edge_feature_utils.py new file mode 100644 index 0000000..8bda82f --- /dev/null +++ b/tests/test_edge_feature_utils.py @@ -0,0 +1,45 @@ +import unittest + +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, +) + + +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)) + + +if __name__ == "__main__": + unittest.main() From f8e77757b6fe85f92ef69835a9e5e1413a3b82db Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:43:55 +0800 Subject: [PATCH 03/10] Update data_loading.py for edge feature fix Signed-off-by: benben951 --- data_loading.py | 271 ++++++++++++++++++++++++------------------------ 1 file changed, 136 insertions(+), 135 deletions(-) diff --git a/data_loading.py b/data_loading.py index 850dc31..3531181 100644 --- a/data_loading.py +++ b/data_loading.py @@ -1,141 +1,142 @@ -import pandas as pd -import numpy as np +import pandas as pd +import numpy as np import torch import logging import itertools -from data_util import GraphData, HeteroData, z_norm, create_hetero_obj - -def get_data(args, data_config): - '''Loads the AML transaction data. - - 1. The data is loaded from the csv and the necessary features are chosen. - 2. The data is split into training, validation and test data. - 3. PyG Data objects are created with the respective data splits. - ''' - - transaction_file = f"{data_config['paths']['aml_data']}/{args.data}/formatted_transactions.csv" #replace this with your path to the respective AML data objects - df_edges = pd.read_csv(transaction_file) - - logging.info(f'Available Edge Features: {df_edges.columns.tolist()}') - - df_edges['Timestamp'] = df_edges['Timestamp'] - df_edges['Timestamp'].min() - - max_n_id = df_edges.loc[:, ['from_id', 'to_id']].to_numpy().max() + 1 - df_nodes = pd.DataFrame({'NodeID': np.arange(max_n_id), 'Feature': np.ones(max_n_id)}) - timestamps = torch.Tensor(df_edges['Timestamp'].to_numpy()) - y = torch.LongTensor(df_edges['Is Laundering'].to_numpy()) - - logging.info(f"Illicit ratio = {sum(y)} / {len(y)} = {sum(y) / len(y) * 100:.2f}%") - logging.info(f"Number of nodes (holdings doing transcations) = {df_nodes.shape[0]}") - logging.info(f"Number of transactions = {df_edges.shape[0]}") - - edge_features = ['Timestamp', 'Amount Received', 'Received Currency', 'Payment Format'] - node_features = ['Feature'] - - logging.info(f'Edge features being used: {edge_features}') - logging.info(f'Node features being used: {node_features} ("Feature" is a placeholder feature of all 1s)') - - x = torch.tensor(df_nodes.loc[:, node_features].to_numpy()).float() - edge_index = torch.LongTensor(df_edges.loc[:, ['from_id', 'to_id']].to_numpy().T) - edge_attr = torch.tensor(df_edges.loc[:, edge_features].to_numpy()).float() - - n_days = int(timestamps.max() / (3600 * 24) + 1) - n_samples = y.shape[0] - logging.info(f'number of days and transactions in the data: {n_days} days, {n_samples} transactions') - - #data splitting - daily_irs, weighted_daily_irs, daily_inds, daily_trans = [], [], [], [] #irs = illicit ratios, inds = indices, trans = transactions - for day in range(n_days): - l = day * 24 * 3600 - r = (day + 1) * 24 * 3600 - day_inds = torch.where((timestamps >= l) & (timestamps < r))[0] - daily_irs.append(y[day_inds].float().mean()) - weighted_daily_irs.append(y[day_inds].float().mean() * day_inds.shape[0] / n_samples) - daily_inds.append(day_inds) - daily_trans.append(day_inds.shape[0]) - - split_per = [0.6, 0.2, 0.2] - daily_totals = np.array(daily_trans) - d_ts = daily_totals - I = list(range(len(d_ts))) - split_scores = dict() - for i,j in itertools.combinations(I, 2): - if j >= i: - split_totals = [d_ts[:i].sum(), d_ts[i:j].sum(), d_ts[j:].sum()] - split_totals_sum = np.sum(split_totals) - split_props = [v/split_totals_sum for v in split_totals] - split_error = [abs(v-t)/t for v,t in zip(split_props, split_per)] - score = max(split_error) #- (split_totals_sum/total) + 1 - split_scores[(i,j)] = score - else: - continue - - i,j = min(split_scores, key=split_scores.get) - #split contains a list for each split (train, validation and test) and each list contains the days that are part of the respective split - split = [list(range(i)), list(range(i, j)), list(range(j, len(daily_totals)))] - logging.info(f'Calculate split: {split}') - - #Now, we seperate the transactions based on their indices in the timestamp array - split_inds = {k: [] for k in range(3)} - for i in range(3): - for day in split[i]: - split_inds[i].append(daily_inds[day]) #split_inds contains a list for each split (tr,val,te) which contains the indices of each day seperately - - tr_inds = torch.cat(split_inds[0]) - val_inds = torch.cat(split_inds[1]) - te_inds = torch.cat(split_inds[2]) - - logging.info(f"Total train samples: {tr_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[tr_inds].float().mean() * 100 :.2f}% || Train days: {split[0][:5]}") - logging.info(f"Total val samples: {val_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[val_inds].float().mean() * 100:.2f}% || Val days: {split[1][:5]}") - logging.info(f"Total test samples: {te_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[te_inds].float().mean() * 100:.2f}% || Test days: {split[2][:5]}") - - #Creating the final data objects - tr_x, val_x, te_x = x, x, x - e_tr = tr_inds.numpy() - e_val = np.concatenate([tr_inds, val_inds]) - - tr_edge_index, tr_edge_attr, tr_y, tr_edge_times = edge_index[:,e_tr], edge_attr[e_tr], y[e_tr], timestamps[e_tr] - val_edge_index, val_edge_attr, val_y, val_edge_times = edge_index[:,e_val], edge_attr[e_val], y[e_val], timestamps[e_val] - te_edge_index, te_edge_attr, te_y, te_edge_times = edge_index, edge_attr, y, timestamps - - tr_data = GraphData (x=tr_x, y=tr_y, edge_index=tr_edge_index, edge_attr=tr_edge_attr, timestamps=tr_edge_times ) - val_data = GraphData(x=val_x, y=val_y, edge_index=val_edge_index, edge_attr=val_edge_attr, timestamps=val_edge_times) - te_data = GraphData (x=te_x, y=te_y, edge_index=te_edge_index, edge_attr=te_edge_attr, timestamps=te_edge_times ) - - #Adding ports and time-deltas if applicable - if args.ports: - logging.info(f"Start: adding ports") - tr_data.add_ports() - val_data.add_ports() - te_data.add_ports() - logging.info(f"Done: adding ports") - if args.tds: - logging.info(f"Start: adding time-deltas") - tr_data.add_time_deltas() - val_data.add_time_deltas() - te_data.add_time_deltas() - logging.info(f"Done: adding time-deltas") - +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. + + 1. The data is loaded from the csv and the necessary features are chosen. + 2. The data is split into training, validation and test data. + 3. PyG Data objects are created with the respective data splits. + ''' + + transaction_file = f"{data_config['paths']['aml_data']}/{args.data}/formatted_transactions.csv" #replace this with your path to the respective AML data objects + df_edges = pd.read_csv(transaction_file) + + logging.info(f'Available Edge Features: {df_edges.columns.tolist()}') + + df_edges['Timestamp'] = df_edges['Timestamp'] - df_edges['Timestamp'].min() + + max_n_id = df_edges.loc[:, ['from_id', 'to_id']].to_numpy().max() + 1 + df_nodes = pd.DataFrame({'NodeID': np.arange(max_n_id), 'Feature': np.ones(max_n_id)}) + timestamps = torch.Tensor(df_edges['Timestamp'].to_numpy()) + y = torch.LongTensor(df_edges['Is Laundering'].to_numpy()) + + logging.info(f"Illicit ratio = {sum(y)} / {len(y)} = {sum(y) / len(y) * 100:.2f}%") + logging.info(f"Number of nodes (holdings doing transcations) = {df_nodes.shape[0]}") + logging.info(f"Number of transactions = {df_edges.shape[0]}") + + edge_features = ['Timestamp', 'Amount Received', 'Received Currency', 'Payment Format'] + node_features = ['Feature'] + + logging.info(f'Edge features being used: {edge_features}') + logging.info(f'Node features being used: {node_features} ("Feature" is a placeholder feature of all 1s)') + + x = torch.tensor(df_nodes.loc[:, node_features].to_numpy()).float() + edge_index = torch.LongTensor(df_edges.loc[:, ['from_id', 'to_id']].to_numpy().T) + edge_attr = torch.tensor(df_edges.loc[:, edge_features].to_numpy()).float() + + n_days = int(timestamps.max() / (3600 * 24) + 1) + n_samples = y.shape[0] + logging.info(f'number of days and transactions in the data: {n_days} days, {n_samples} transactions') + + #data splitting + daily_irs, weighted_daily_irs, daily_inds, daily_trans = [], [], [], [] #irs = illicit ratios, inds = indices, trans = transactions + for day in range(n_days): + l = day * 24 * 3600 + r = (day + 1) * 24 * 3600 + day_inds = torch.where((timestamps >= l) & (timestamps < r))[0] + daily_irs.append(y[day_inds].float().mean()) + weighted_daily_irs.append(y[day_inds].float().mean() * day_inds.shape[0] / n_samples) + daily_inds.append(day_inds) + daily_trans.append(day_inds.shape[0]) + + split_per = [0.6, 0.2, 0.2] + daily_totals = np.array(daily_trans) + d_ts = daily_totals + I = list(range(len(d_ts))) + split_scores = dict() + for i,j in itertools.combinations(I, 2): + if j >= i: + split_totals = [d_ts[:i].sum(), d_ts[i:j].sum(), d_ts[j:].sum()] + split_totals_sum = np.sum(split_totals) + split_props = [v/split_totals_sum for v in split_totals] + split_error = [abs(v-t)/t for v,t in zip(split_props, split_per)] + score = max(split_error) #- (split_totals_sum/total) + 1 + split_scores[(i,j)] = score + else: + continue + + i,j = min(split_scores, key=split_scores.get) + #split contains a list for each split (train, validation and test) and each list contains the days that are part of the respective split + split = [list(range(i)), list(range(i, j)), list(range(j, len(daily_totals)))] + logging.info(f'Calculate split: {split}') + + #Now, we seperate the transactions based on their indices in the timestamp array + split_inds = {k: [] for k in range(3)} + for i in range(3): + for day in split[i]: + split_inds[i].append(daily_inds[day]) #split_inds contains a list for each split (tr,val,te) which contains the indices of each day seperately + + tr_inds = torch.cat(split_inds[0]) + val_inds = torch.cat(split_inds[1]) + te_inds = torch.cat(split_inds[2]) + + logging.info(f"Total train samples: {tr_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[tr_inds].float().mean() * 100 :.2f}% || Train days: {split[0][:5]}") + logging.info(f"Total val samples: {val_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[val_inds].float().mean() * 100:.2f}% || Val days: {split[1][:5]}") + logging.info(f"Total test samples: {te_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[te_inds].float().mean() * 100:.2f}% || Test days: {split[2][:5]}") + + #Creating the final data objects + tr_x, val_x, te_x = x, x, x + e_tr = tr_inds.numpy() + e_val = np.concatenate([tr_inds, val_inds]) + + tr_edge_index, tr_edge_attr, tr_y, tr_edge_times = edge_index[:,e_tr], edge_attr[e_tr], y[e_tr], timestamps[e_tr] + val_edge_index, val_edge_attr, val_y, val_edge_times = edge_index[:,e_val], edge_attr[e_val], y[e_val], timestamps[e_val] + te_edge_index, te_edge_attr, te_y, te_edge_times = edge_index, edge_attr, y, timestamps + + tr_data = GraphData (x=tr_x, y=tr_y, edge_index=tr_edge_index, edge_attr=tr_edge_attr, timestamps=tr_edge_times ) + val_data = GraphData(x=val_x, y=val_y, edge_index=val_edge_index, edge_attr=val_edge_attr, timestamps=val_edge_times) + te_data = GraphData (x=te_x, y=te_y, edge_index=te_edge_index, edge_attr=te_edge_attr, timestamps=te_edge_times ) + + #Adding ports and time-deltas if applicable + if args.ports: + logging.info(f"Start: adding ports") + tr_data.add_ports() + val_data.add_ports() + te_data.add_ports() + logging.info(f"Done: adding ports") + if args.tds: + logging.info(f"Start: adding time-deltas") + tr_data.add_time_deltas() + val_data.add_time_deltas() + te_data.add_time_deltas() + logging.info(f"Done: adding time-deltas") + #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]) - - #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 - if args.reverse_mp: - tr_data = create_hetero_obj(tr_data.x, tr_data.y, tr_data.edge_index, tr_data.edge_attr, tr_data.timestamps, args) - val_data = create_hetero_obj(val_data.x, val_data.y, val_data.edge_index, val_data.edge_attr, val_data.timestamps, args) - te_data = create_hetero_obj(te_data.x, te_data.y, te_data.edge_index, te_data.edge_attr, te_data.timestamps, args) + 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 + if args.reverse_mp: + tr_data = create_hetero_obj(tr_data.x, tr_data.y, tr_data.edge_index, tr_data.edge_attr, tr_data.timestamps, args) + val_data = create_hetero_obj(val_data.x, val_data.y, val_data.edge_index, val_data.edge_attr, val_data.timestamps, args) + te_data = create_hetero_obj(te_data.x, te_data.y, te_data.edge_index, te_data.edge_attr, te_data.timestamps, args) + + logging.info(f'train data object: {tr_data}') + logging.info(f'validation data object: {val_data}') + logging.info(f'test data object: {te_data}') + + return tr_data, val_data, te_data, tr_inds, val_inds, te_inds - logging.info(f'train data object: {tr_data}') - logging.info(f'validation data object: {val_data}') - logging.info(f'test data object: {te_data}') - - return tr_data, val_data, te_data, tr_inds, val_inds, te_inds - \ No newline at end of file From 1317e241dd59d0182e56cd73fecfb150b3138ed1 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:52:02 +0800 Subject: [PATCH 04/10] Update edge_feature_utils.py for edge feature fix Signed-off-by: benben951 From 78530cc78de75915e0dbbbbfb3480f1740611da3 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:52:03 +0800 Subject: [PATCH 05/10] Update data_util.py for edge feature fix Signed-off-by: benben951 --- data_util.py | 280 +++++++++++++++++++++++++-------------------------- 1 file changed, 139 insertions(+), 141 deletions(-) diff --git a/data_util.py b/data_util.py index 22dcfaa..954f6e7 100644 --- a/data_util.py +++ b/data_util.py @@ -2,155 +2,153 @@ from torch_geometric.data import Data, HeteroData from torch_geometric.typing import OptTensor import numpy as np - -def to_adj_nodes_with_times(data): - num_nodes = data.num_nodes - timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) - edges = torch.cat((data.edge_index.T, timestamps), dim=1) if not isinstance(data, HeteroData) else torch.cat((data['node', 'to', 'node'].edge_index.T, timestamps), dim=1) - adj_list_out = dict([(i, []) for i in range(num_nodes)]) - adj_list_in = dict([(i, []) for i in range(num_nodes)]) - for u,v,t in edges: - u,v,t = int(u), int(v), int(t) - adj_list_out[u] += [(v, t)] - adj_list_in[v] += [(u, t)] - return adj_list_in, adj_list_out - -def to_adj_edges_with_times(data): - num_nodes = data.num_nodes - timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) - edges = torch.cat((data.edge_index.T, timestamps), dim=1) - # calculate adjacent edges with times per node - adj_edges_out = dict([(i, []) for i in range(num_nodes)]) - adj_edges_in = dict([(i, []) for i in range(num_nodes)]) - for i, (u,v,t) in enumerate(edges): - u,v,t = int(u), int(v), int(t) - adj_edges_out[u] += [(i, v, t)] - adj_edges_in[v] += [(i, u, t)] - return adj_edges_in, adj_edges_out - -def ports(edge_index, adj_list): - ports = torch.zeros(edge_index.shape[1], 1) - ports_dict = {} - for v, nbs in adj_list.items(): - if len(nbs) < 1: continue - a = np.array(nbs) - a = a[a[:, -1].argsort()] - _, idx = np.unique(a[:,[0]],return_index=True,axis=0) - nbs_unique = a[np.sort(idx)][:,0] - for i, u in enumerate(nbs_unique): - ports_dict[(u,v)] = i - for i, e in enumerate(edge_index.T): - ports[i] = ports_dict[tuple(e.numpy())] - return ports - -def time_deltas(data, adj_edges_list): - time_deltas = torch.zeros(data.edge_index.shape[1], 1) - if data.timestamps is None: - return time_deltas - for v, edges in adj_edges_list.items(): - if len(edges) < 1: continue - a = np.array(edges) - a = a[a[:, -1].argsort()] - a_tds = [0] + [a[i+1,-1] - a[i,-1] for i in range(a.shape[0]-1)] - tds = np.hstack((a[:,0].reshape(-1,1), np.array(a_tds).reshape(-1,1))) - for i,td in tds: - time_deltas[i] = td - return time_deltas - -class GraphData(Data): - '''This is the homogenous graph object we use for GNN training if reverse MP is not enabled''' - def __init__( - self, x: OptTensor = None, edge_index: OptTensor = None, edge_attr: OptTensor = None, y: OptTensor = None, pos: OptTensor = None, - readout: str = 'edge', - num_nodes: int = None, - timestamps: OptTensor = None, - node_timestamps: OptTensor = None, - **kwargs - ): - super().__init__(x, edge_index, edge_attr, y, pos, **kwargs) - self.readout = readout - self.loss_fn = 'ce' - self.num_nodes = int(self.x.shape[0]) - self.node_timestamps = node_timestamps - if timestamps is not None: - self.timestamps = timestamps - elif edge_attr is not None: - self.timestamps = edge_attr[:,0].clone() - else: - self.timestamps = None - - def add_ports(self): - '''Adds port numberings to the edge features''' - reverse_ports = True - adj_list_in, adj_list_out = to_adj_nodes_with_times(self) - in_ports = ports(self.edge_index, adj_list_in) - out_ports = [ports(self.edge_index.flipud(), adj_list_out)] if reverse_ports else [] - self.edge_attr = torch.cat([self.edge_attr, in_ports] + out_ports, dim=1) - return self - - def add_time_deltas(self): - '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' - reverse_tds = True - adj_list_in, adj_list_out = to_adj_edges_with_times(self) - in_tds = time_deltas(self, adj_list_in) - out_tds = [time_deltas(self, adj_list_out)] if reverse_tds else [] - self.edge_attr = torch.cat([self.edge_attr, in_tds] + out_tds, dim=1) - return self - -class HeteroGraphData(HeteroData): - '''This is the heterogenous graph object we use for GNN training if reverse MP is enabled''' - def __init__( - self, - readout: str = 'edge', - **kwargs - ): - super().__init__(**kwargs) - self.readout = readout - - @property - def num_nodes(self): - return self['node'].x.shape[0] - - @property - def timestamps(self): - return self['node', 'to', 'node'].timestamps - - def add_ports(self): - '''Adds port numberings to the edge features''' - adj_list_in, adj_list_out = to_adj_nodes_with_times(self) - in_ports = ports(self['node', 'to', 'node'].edge_index, adj_list_in) - out_ports = ports(self['node', 'rev_to', 'node'].edge_index, adj_list_out) - self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_ports], dim=1) - self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_ports], dim=1) - return self - - def add_time_deltas(self): - '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' - adj_list_in, adj_list_out = to_adj_edges_with_times(self) - in_tds = time_deltas(self, adj_list_in) - out_tds = time_deltas(self, adj_list_out) - self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_tds], dim=1) - 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 - +from edge_feature_utils import get_port_feature_indices + +def to_adj_nodes_with_times(data): + num_nodes = data.num_nodes + timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) + edges = torch.cat((data.edge_index.T, timestamps), dim=1) if not isinstance(data, HeteroData) else torch.cat((data['node', 'to', 'node'].edge_index.T, timestamps), dim=1) + adj_list_out = dict([(i, []) for i in range(num_nodes)]) + adj_list_in = dict([(i, []) for i in range(num_nodes)]) + for u,v,t in edges: + u,v,t = int(u), int(v), int(t) + adj_list_out[u] += [(v, t)] + adj_list_in[v] += [(u, t)] + return adj_list_in, adj_list_out + +def to_adj_edges_with_times(data): + num_nodes = data.num_nodes + timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) + edges = torch.cat((data.edge_index.T, timestamps), dim=1) + # calculate adjacent edges with times per node + adj_edges_out = dict([(i, []) for i in range(num_nodes)]) + adj_edges_in = dict([(i, []) for i in range(num_nodes)]) + for i, (u,v,t) in enumerate(edges): + u,v,t = int(u), int(v), int(t) + adj_edges_out[u] += [(i, v, t)] + adj_edges_in[v] += [(i, u, t)] + return adj_edges_in, adj_edges_out + +def ports(edge_index, adj_list): + ports = torch.zeros(edge_index.shape[1], 1) + ports_dict = {} + for v, nbs in adj_list.items(): + if len(nbs) < 1: continue + a = np.array(nbs) + a = a[a[:, -1].argsort()] + _, idx = np.unique(a[:,[0]],return_index=True,axis=0) + nbs_unique = a[np.sort(idx)][:,0] + for i, u in enumerate(nbs_unique): + ports_dict[(u,v)] = i + for i, e in enumerate(edge_index.T): + ports[i] = ports_dict[tuple(e.numpy())] + return ports + +def time_deltas(data, adj_edges_list): + time_deltas = torch.zeros(data.edge_index.shape[1], 1) + if data.timestamps is None: + return time_deltas + for v, edges in adj_edges_list.items(): + if len(edges) < 1: continue + a = np.array(edges) + a = a[a[:, -1].argsort()] + a_tds = [0] + [a[i+1,-1] - a[i,-1] for i in range(a.shape[0]-1)] + tds = np.hstack((a[:,0].reshape(-1,1), np.array(a_tds).reshape(-1,1))) + for i,td in tds: + time_deltas[i] = td + return time_deltas + +class GraphData(Data): + '''This is the homogenous graph object we use for GNN training if reverse MP is not enabled''' + def __init__( + self, x: OptTensor = None, edge_index: OptTensor = None, edge_attr: OptTensor = None, y: OptTensor = None, pos: OptTensor = None, + readout: str = 'edge', + num_nodes: int = None, + timestamps: OptTensor = None, + node_timestamps: OptTensor = None, + **kwargs + ): + super().__init__(x, edge_index, edge_attr, y, pos, **kwargs) + self.readout = readout + self.loss_fn = 'ce' + self.num_nodes = int(self.x.shape[0]) + self.node_timestamps = node_timestamps + if timestamps is not None: + self.timestamps = timestamps + elif edge_attr is not None: + self.timestamps = edge_attr[:,0].clone() + else: + self.timestamps = None + + def add_ports(self): + '''Adds port numberings to the edge features''' + reverse_ports = True + adj_list_in, adj_list_out = to_adj_nodes_with_times(self) + in_ports = ports(self.edge_index, adj_list_in) + out_ports = [ports(self.edge_index.flipud(), adj_list_out)] if reverse_ports else [] + self.edge_attr = torch.cat([self.edge_attr, in_ports] + out_ports, dim=1) + return self + + def add_time_deltas(self): + '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' + reverse_tds = True + adj_list_in, adj_list_out = to_adj_edges_with_times(self) + in_tds = time_deltas(self, adj_list_in) + out_tds = [time_deltas(self, adj_list_out)] if reverse_tds else [] + self.edge_attr = torch.cat([self.edge_attr, in_tds] + out_tds, dim=1) + return self + +class HeteroGraphData(HeteroData): + '''This is the heterogenous graph object we use for GNN training if reverse MP is enabled''' + def __init__( + self, + readout: str = 'edge', + **kwargs + ): + super().__init__(**kwargs) + self.readout = readout + + @property + def num_nodes(self): + return self['node'].x.shape[0] + + @property + def timestamps(self): + return self['node', 'to', 'node'].timestamps + + def add_ports(self): + '''Adds port numberings to the edge features''' + adj_list_in, adj_list_out = to_adj_nodes_with_times(self) + in_ports = ports(self['node', 'to', 'node'].edge_index, adj_list_in) + out_ports = ports(self['node', 'rev_to', 'node'].edge_index, adj_list_out) + self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_ports], dim=1) + self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_ports], dim=1) + return self + + def add_time_deltas(self): + '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' + adj_list_in, adj_list_out = to_adj_edges_with_times(self) + in_tds = time_deltas(self, adj_list_in) + out_tds = time_deltas(self, adj_list_out) + self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_tds], dim=1) + self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_tds], dim=1) + return self + 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 \ No newline at end of file + return data From 88c88b0058bf06cbdbe5612a25c5427322a18f0f Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:52:04 +0800 Subject: [PATCH 06/10] Update models.py for edge feature fix Signed-off-by: benben951 --- models.py | 398 +++++++++++++++++++++++++++--------------------------- 1 file changed, 200 insertions(+), 198 deletions(-) diff --git a/models.py b/models.py index 568b29e..cee87d7 100644 --- a/models.py +++ b/models.py @@ -1,236 +1,238 @@ -import torch.nn as nn -from torch_geometric.nn import GINEConv, BatchNorm, Linear, GATConv, PNAConv, RGCNConv +import torch.nn as nn +from torch_geometric.nn import GINEConv, BatchNorm, Linear, GATConv, PNAConv, RGCNConv import torch.nn.functional as F import torch import logging - -class GINe(torch.nn.Module): - def __init__(self, num_features, num_gnn_layers, n_classes=2, - n_hidden=100, edge_updates=False, residual=True, - edge_dim=None, dropout=0.0, final_dropout=0.5): - super().__init__() - self.n_hidden = n_hidden - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.final_dropout = final_dropout - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - for _ in range(self.num_gnn_layers): - conv = GINEConv(nn.Sequential( - nn.Linear(self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden) - ), edge_dim=self.n_hidden) - if self.edge_updates: self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - +from edge_feature_utils import average_residual_update + +class GINe(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, + n_hidden=100, edge_updates=False, residual=True, + edge_dim=None, dropout=0.0, final_dropout=0.5): + super().__init__() + self.n_hidden = n_hidden + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.final_dropout = final_dropout + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + for _ in range(self.num_gnn_layers): + conv = GINEConv(nn.Sequential( + nn.Linear(self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden) + ), edge_dim=self.n_hidden) + if self.edge_updates: self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + 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.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 - - 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) - out = x - - return self.mlp(out) - -class GATe(torch.nn.Module): - def __init__(self, num_features, num_gnn_layers, n_classes=2, n_hidden=100, n_heads=4, edge_updates=False, edge_dim=None, dropout=0.0, final_dropout=0.5): - super().__init__() - # GAT specific code - tmp_out = n_hidden // n_heads - n_hidden = tmp_out * n_heads - - self.n_hidden = n_hidden - self.n_heads = n_heads - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.dropout = dropout - self.final_dropout = final_dropout - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - - for _ in range(self.num_gnn_layers): - conv = GATConv(self.n_hidden, tmp_out, self.n_heads, concat = True, dropout = self.dropout, add_self_loops = True, edge_dim=self.n_hidden) - if self.edge_updates: self.emlps.append(nn.Sequential(nn.Linear(3 * self.n_hidden, self.n_hidden),nn.ReLU(),nn.Linear(self.n_hidden, self.n_hidden),)) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - + 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) + out = x + + return self.mlp(out) + +class GATe(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, n_hidden=100, n_heads=4, edge_updates=False, edge_dim=None, dropout=0.0, final_dropout=0.5): + super().__init__() + # GAT specific code + tmp_out = n_hidden // n_heads + n_hidden = tmp_out * n_heads + + self.n_hidden = n_hidden + self.n_heads = n_heads + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.dropout = dropout + self.final_dropout = final_dropout + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + + for _ in range(self.num_gnn_layers): + conv = GATConv(self.n_hidden, tmp_out, self.n_heads, concat = True, dropout = self.dropout, add_self_loops = True, edge_dim=self.n_hidden) + if self.edge_updates: self.emlps.append(nn.Sequential(nn.Linear(3 * self.n_hidden, self.n_hidden),nn.ReLU(),nn.Linear(self.n_hidden, self.n_hidden),)) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + 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.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 - - 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() - logging.debug(f"x.shape = {x.shape}") - x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) - logging.debug(f"x.shape = {x.shape}") - out = x - - return self.mlp(out) - -class PNA(torch.nn.Module): - def __init__(self, num_features, num_gnn_layers, n_classes=2, - n_hidden=100, edge_updates=True, - edge_dim=None, dropout=0.0, final_dropout=0.5, deg=None): - super().__init__() - n_hidden = int((n_hidden // 5) * 5) - self.n_hidden = n_hidden - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.final_dropout = final_dropout - - aggregators = ['mean', 'min', 'max', 'std'] - scalers = ['identity', 'amplification', 'attenuation'] - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - for _ in range(self.num_gnn_layers): - conv = PNAConv(in_channels=n_hidden, out_channels=n_hidden, - aggregators=aggregators, scalers=scalers, deg=deg, - edge_dim=n_hidden, towers=5, pre_layers=1, post_layers=1, - divide_input=False) - if self.edge_updates: self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - + 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() + logging.debug(f"x.shape = {x.shape}") + x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) + logging.debug(f"x.shape = {x.shape}") + out = x + + return self.mlp(out) + +class PNA(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, + n_hidden=100, edge_updates=True, + edge_dim=None, dropout=0.0, final_dropout=0.5, deg=None): + super().__init__() + n_hidden = int((n_hidden // 5) * 5) + self.n_hidden = n_hidden + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.final_dropout = final_dropout + + aggregators = ['mean', 'min', 'max', 'std'] + scalers = ['identity', 'amplification', 'attenuation'] + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + for _ in range(self.num_gnn_layers): + conv = PNAConv(in_channels=n_hidden, out_channels=n_hidden, + aggregators=aggregators, scalers=scalers, deg=deg, + edge_dim=n_hidden, towers=5, pre_layers=1, post_layers=1, + divide_input=False) + if self.edge_updates: self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + 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.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 - - 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() - logging.debug(f"x.shape = {x.shape}") - x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) - logging.debug(f"x.shape = {x.shape}") - out = x - return self.mlp(out) - + 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() + logging.debug(f"x.shape = {x.shape}") + x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) + logging.debug(f"x.shape = {x.shape}") + out = x + return self.mlp(out) + 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 - self.num_gnn_layers = num_gnn_layers - self.n_hidden = n_hidden - self.residual = residual - self.dropout = dropout - self.final_dropout = final_dropout + + self.num_features = num_features + self.num_gnn_layers = num_gnn_layers + self.n_hidden = n_hidden + self.residual = residual + self.dropout = dropout + self.final_dropout = final_dropout self.n_classes = n_classes 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) - - self.convs = nn.ModuleList() - self.bns = nn.ModuleList() - self.mlp = nn.ModuleList() - - if self.edge_update: - self.emlps = nn.ModuleList() - self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - - for _ in range(self.num_gnn_layers): - conv = RGCNConv(self.n_hidden, self.n_hidden, num_relations, num_bases=self.n_bases) - self.convs.append(conv) - self.bns.append(nn.BatchNorm1d(self.n_hidden)) - - if self.edge_update: - self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout), Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def reset_parameters(self): - for m in self.modules(): - if isinstance(m, nn.Linear): - m.reset_parameters() - elif isinstance(m, RGCNConv): - m.reset_parameters() - elif isinstance(m, nn.BatchNorm1d): - m.reset_parameters() - + + self.convs = nn.ModuleList() + self.bns = nn.ModuleList() + self.mlp = nn.ModuleList() + + if self.edge_update: + self.emlps = nn.ModuleList() + self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + + for _ in range(self.num_gnn_layers): + conv = RGCNConv(self.n_hidden, self.n_hidden, num_relations, num_bases=self.n_bases) + self.convs.append(conv) + self.bns.append(nn.BatchNorm1d(self.n_hidden)) + + if self.edge_update: + self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout), Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def reset_parameters(self): + for m in self.modules(): + if isinstance(m, nn.Linear): + m.reset_parameters() + elif isinstance(m, RGCNConv): + m.reset_parameters() + elif isinstance(m, nn.BatchNorm1d): + 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 - - 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 \ No newline at end of file + 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 From 716c8aca50a2893272d34df8574733af075c77d6 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:52:05 +0800 Subject: [PATCH 07/10] Update training.py for edge feature fix Signed-off-by: benben951 --- training.py | 462 ++++++++++++++++++++++++++-------------------------- 1 file changed, 232 insertions(+), 230 deletions(-) diff --git a/training.py b/training.py index 797357c..45388b3 100644 --- a/training.py +++ b/training.py @@ -1,237 +1,239 @@ -import torch -import tqdm -from sklearn.metrics import f1_score -from train_util import AddEgoIds, extract_param, add_arange_ids, get_loaders, evaluate_homo, evaluate_hetero, save_model, load_model -from models import GINe, PNA, GATe, RGCN -from torch_geometric.data import Data, HeteroData -from torch_geometric.nn import to_hetero, summary -from torch_geometric.utils import degree +import torch +import tqdm +from sklearn.metrics import f1_score +from train_util import AddEgoIds, extract_param, add_arange_ids, get_loaders, evaluate_homo, evaluate_hetero, save_model, load_model +from models import GINe, PNA, GATe, RGCN +from torch_geometric.data import Data, HeteroData +from torch_geometric.nn import to_hetero, summary +from torch_geometric.utils import degree import wandb import logging - -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 - best_val_f1 = 0 - for epoch in range(config.epochs): - total_loss = total_examples = 0 - preds = [] - ground_truths = [] - for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): - optimizer.zero_grad() - #select the seed edges from which the batch was created - inds = tr_inds.detach().cpu() - batch_edge_inds = inds[batch.input_id.detach().cpu()] - batch_edge_ids = tr_loader.data.edge_attr.detach().cpu()[batch_edge_inds, 0] - mask = torch.isin(batch.edge_attr[:, 0].detach().cpu(), batch_edge_ids) - - #remove the unique edge id from the edge features, as it's no longer needed - batch.edge_attr = batch.edge_attr[:, 1:] - - batch.to(device) - out = model(batch.x, batch.edge_index, batch.edge_attr) - pred = out[mask] - ground_truth = batch.y[mask] - preds.append(pred.argmax(dim=-1)) - ground_truths.append(ground_truth) - loss = loss_fn(pred, ground_truth) - - loss.backward() - optimizer.step() - - total_loss += float(loss) * pred.numel() - total_examples += pred.numel() - - pred = torch.cat(preds, dim=0).detach().cpu().numpy() - ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() - f1 = f1_score(ground_truth, pred) - wandb.log({"f1/train": f1}, step=epoch) - logging.info(f'Train F1: {f1:.4f}') - - #evaluate - val_f1 = evaluate_homo(val_loader, val_inds, model, val_data, device, args) - te_f1 = evaluate_homo(te_loader, te_inds, model, te_data, device, args) - - wandb.log({"f1/validation": val_f1}, step=epoch) - wandb.log({"f1/test": te_f1}, step=epoch) - logging.info(f'Validation F1: {val_f1:.4f}') - logging.info(f'Test F1: {te_f1:.4f}') - - if epoch == 0: - wandb.log({"best_test_f1": te_f1}, step=epoch) - elif val_f1 > best_val_f1: - best_val_f1 = val_f1 - wandb.log({"best_test_f1": te_f1}, step=epoch) - if args.save_model: - save_model(model, optimizer, epoch, args, data_config) - - return model - -def train_hetero(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 - best_val_f1 = 0 - for epoch in range(config.epochs): - total_loss = total_examples = 0 - preds = [] - ground_truths = [] - for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): - optimizer.zero_grad() - #select the seed edges from which the batch was created - inds = tr_inds.detach().cpu() - batch_edge_inds = inds[batch['node', 'to', 'node'].input_id.detach().cpu()] - batch_edge_ids = tr_loader.data['node', 'to', 'node'].edge_attr.detach().cpu()[batch_edge_inds, 0] - mask = torch.isin(batch['node', 'to', 'node'].edge_attr[:, 0].detach().cpu(), batch_edge_ids) - - #remove the unique edge id from the edge features, as it's no longer needed - batch['node', 'to', 'node'].edge_attr = batch['node', 'to', 'node'].edge_attr[:, 1:] - batch['node', 'rev_to', 'node'].edge_attr = batch['node', 'rev_to', 'node'].edge_attr[:, 1:] - - batch.to(device) - out = model(batch.x_dict, batch.edge_index_dict, batch.edge_attr_dict) - out = out[('node', 'to', 'node')] - pred = out[mask] - ground_truth = batch['node', 'to', 'node'].y[mask] - preds.append(pred.argmax(dim=-1)) - ground_truths.append(batch['node', 'to', 'node'].y[mask]) - loss = loss_fn(pred, ground_truth) - - loss.backward() - optimizer.step() - - total_loss += float(loss) * pred.numel() - total_examples += pred.numel() - - pred = torch.cat(preds, dim=0).detach().cpu().numpy() - ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() - f1 = f1_score(ground_truth, pred) - wandb.log({"f1/train": f1}, step=epoch) - logging.info(f'Train F1: {f1:.4f}') - - #evaluate - val_f1 = evaluate_hetero(val_loader, val_inds, model, val_data, device, args) - te_f1 = evaluate_hetero(te_loader, te_inds, model, te_data, device, args) - - wandb.log({"f1/validation": val_f1}, step=epoch) - wandb.log({"f1/test": te_f1}, step=epoch) - logging.info(f'Validation F1: {val_f1:.4f}') - logging.info(f'Test F1: {te_f1:.4f}') - - if epoch == 0: - wandb.log({"best_test_f1": te_f1}, step=epoch) - elif val_f1 > best_val_f1: - best_val_f1 = val_f1 - wandb.log({"best_test_f1": te_f1}, step=epoch) - if args.save_model: - save_model(model, optimizer, epoch, args, data_config) - - return model - -def get_model(sample_batch, config, args): - n_feats = sample_batch.x.shape[1] if not isinstance(sample_batch, HeteroData) else sample_batch['node'].x.shape[1] - e_dim = (sample_batch.edge_attr.shape[1] - 1) if not isinstance(sample_batch, HeteroData) else (sample_batch['node', 'to', 'node'].edge_attr.shape[1] - 1) - - if args.model == "gin": - model = GINe( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), residual=False, edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, final_dropout=config.final_dropout - ) - elif args.model == "gat": - model = GATe( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), n_heads=round(config.n_heads), - edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, final_dropout=config.final_dropout - ) - elif args.model == "pna": - if not isinstance(sample_batch, HeteroData): - d = degree(sample_batch.edge_index[1], dtype=torch.long) - else: - index = torch.cat((sample_batch['node', 'to', 'node'].edge_index[1], sample_batch['node', 'rev_to', 'node'].edge_index[1]), 0) - d = degree(index, dtype=torch.long) - deg = torch.bincount(d, minlength=1) - model = PNA( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, deg=deg, final_dropout=config.final_dropout - ) +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 + best_val_f1 = 0 + for epoch in range(config.epochs): + total_loss = total_examples = 0 + preds = [] + ground_truths = [] + for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): + optimizer.zero_grad() + #select the seed edges from which the batch was created + inds = tr_inds.detach().cpu() + batch_edge_inds = inds[batch.input_id.detach().cpu()] + batch_edge_ids = tr_loader.data.edge_attr.detach().cpu()[batch_edge_inds, 0] + mask = torch.isin(batch.edge_attr[:, 0].detach().cpu(), batch_edge_ids) + + #remove the unique edge id from the edge features, as it's no longer needed + batch.edge_attr = batch.edge_attr[:, 1:] + + batch.to(device) + out = model(batch.x, batch.edge_index, batch.edge_attr) + pred = out[mask] + ground_truth = batch.y[mask] + preds.append(pred.argmax(dim=-1)) + ground_truths.append(ground_truth) + loss = loss_fn(pred, ground_truth) + + loss.backward() + optimizer.step() + + total_loss += float(loss) * pred.numel() + total_examples += pred.numel() + + pred = torch.cat(preds, dim=0).detach().cpu().numpy() + ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() + f1 = f1_score(ground_truth, pred) + wandb.log({"f1/train": f1}, step=epoch) + logging.info(f'Train F1: {f1:.4f}') + + #evaluate + val_f1 = evaluate_homo(val_loader, val_inds, model, val_data, device, args) + te_f1 = evaluate_homo(te_loader, te_inds, model, te_data, device, args) + + wandb.log({"f1/validation": val_f1}, step=epoch) + wandb.log({"f1/test": te_f1}, step=epoch) + logging.info(f'Validation F1: {val_f1:.4f}') + logging.info(f'Test F1: {te_f1:.4f}') + + if epoch == 0: + wandb.log({"best_test_f1": te_f1}, step=epoch) + elif val_f1 > best_val_f1: + best_val_f1 = val_f1 + wandb.log({"best_test_f1": te_f1}, step=epoch) + if args.save_model: + save_model(model, optimizer, epoch, args, data_config) + + return model + +def train_hetero(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 + best_val_f1 = 0 + for epoch in range(config.epochs): + total_loss = total_examples = 0 + preds = [] + ground_truths = [] + for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): + optimizer.zero_grad() + #select the seed edges from which the batch was created + inds = tr_inds.detach().cpu() + batch_edge_inds = inds[batch['node', 'to', 'node'].input_id.detach().cpu()] + batch_edge_ids = tr_loader.data['node', 'to', 'node'].edge_attr.detach().cpu()[batch_edge_inds, 0] + mask = torch.isin(batch['node', 'to', 'node'].edge_attr[:, 0].detach().cpu(), batch_edge_ids) + + #remove the unique edge id from the edge features, as it's no longer needed + batch['node', 'to', 'node'].edge_attr = batch['node', 'to', 'node'].edge_attr[:, 1:] + batch['node', 'rev_to', 'node'].edge_attr = batch['node', 'rev_to', 'node'].edge_attr[:, 1:] + + batch.to(device) + out = model(batch.x_dict, batch.edge_index_dict, batch.edge_attr_dict) + out = out[('node', 'to', 'node')] + pred = out[mask] + ground_truth = batch['node', 'to', 'node'].y[mask] + preds.append(pred.argmax(dim=-1)) + ground_truths.append(batch['node', 'to', 'node'].y[mask]) + loss = loss_fn(pred, ground_truth) + + loss.backward() + optimizer.step() + + total_loss += float(loss) * pred.numel() + total_examples += pred.numel() + + pred = torch.cat(preds, dim=0).detach().cpu().numpy() + ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() + f1 = f1_score(ground_truth, pred) + wandb.log({"f1/train": f1}, step=epoch) + logging.info(f'Train F1: {f1:.4f}') + + #evaluate + val_f1 = evaluate_hetero(val_loader, val_inds, model, val_data, device, args) + te_f1 = evaluate_hetero(te_loader, te_inds, model, te_data, device, args) + + wandb.log({"f1/validation": val_f1}, step=epoch) + wandb.log({"f1/test": te_f1}, step=epoch) + logging.info(f'Validation F1: {val_f1:.4f}') + logging.info(f'Test F1: {te_f1:.4f}') + + if epoch == 0: + wandb.log({"best_test_f1": te_f1}, step=epoch) + elif val_f1 > best_val_f1: + best_val_f1 = val_f1 + wandb.log({"best_test_f1": te_f1}, step=epoch) + if args.save_model: + save_model(model, optimizer, epoch, args, data_config) + + return model + +def get_model(sample_batch, config, args): + n_feats = sample_batch.x.shape[1] if not isinstance(sample_batch, HeteroData) else sample_batch['node'].x.shape[1] + e_dim = (sample_batch.edge_attr.shape[1] - 1) if not isinstance(sample_batch, HeteroData) else (sample_batch['node', 'to', 'node'].edge_attr.shape[1] - 1) + + if args.model == "gin": + model = GINe( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), residual=False, edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, final_dropout=config.final_dropout + ) + elif args.model == "gat": + model = GATe( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), n_heads=round(config.n_heads), + edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, final_dropout=config.final_dropout + ) + elif args.model == "pna": + if not isinstance(sample_batch, HeteroData): + d = degree(sample_batch.edge_index[1], dtype=torch.long) + else: + index = torch.cat((sample_batch['node', 'to', 'node'].edge_index[1], sample_batch['node', 'rev_to', 'node'].edge_index[1]), 0) + d = degree(index, dtype=torch.long) + deg = torch.bincount(d, minlength=1) + model = PNA( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, deg=deg, final_dropout=config.final_dropout + ) elif config.model == "rgcn": 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 - -def train_gnn(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, args, data_config): - #set device - device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") - - #define a model config dictionary and wandb logging at the same time - wandb.init( - mode="disabled" if args.testing else "online", - project="your_proj_name", #replace this with your wandb project name if you want to use wandb logging - - config={ - "epochs": args.n_epochs, - "batch_size": args.batch_size, - "model": args.model, - "data": args.data, - "num_neighbors": args.num_neighs, - "lr": extract_param("lr", args), - "n_hidden": extract_param("n_hidden", args), - "n_gnn_layers": extract_param("n_gnn_layers", args), - "loss": "ce", - "w_ce1": extract_param("w_ce1", args), - "w_ce2": extract_param("w_ce2", args), - "dropout": extract_param("dropout", args), - "final_dropout": extract_param("final_dropout", args), - "n_heads": extract_param("n_heads", args) if args.model == 'gat' else None - } - ) - - config = wandb.config - - #set the transform if ego ids should be used - if args.ego: - transform = AddEgoIds() - else: - transform = None - - #add the unique ids to later find the seed edges - add_arange_ids([tr_data, val_data, te_data]) - - tr_loader, val_loader, te_loader = get_loaders(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, transform, args) - - #get the model - sample_batch = next(iter(tr_loader)) - model = get_model(sample_batch, config, args) - - if args.reverse_mp: - model = to_hetero(model, te_data.metadata(), aggr='mean') - - if args.finetune: - model, optimizer = load_model(model, device, args, config, data_config) - else: - model.to(device) - optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) - - sample_batch.to(device) - sample_x = sample_batch.x if not isinstance(sample_batch, HeteroData) else sample_batch.x_dict - sample_edge_index = sample_batch.edge_index if not isinstance(sample_batch, HeteroData) else sample_batch.edge_index_dict - if isinstance(sample_batch, HeteroData): - sample_batch['node', 'to', 'node'].edge_attr = sample_batch['node', 'to', 'node'].edge_attr[:, 1:] - sample_batch['node', 'rev_to', 'node'].edge_attr = sample_batch['node', 'rev_to', 'node'].edge_attr[:, 1:] - else: - sample_batch.edge_attr = sample_batch.edge_attr[:, 1:] - sample_edge_attr = sample_batch.edge_attr if not isinstance(sample_batch, HeteroData) else sample_batch.edge_attr_dict - logging.info(summary(model, sample_x, sample_edge_index, sample_edge_attr)) - - loss_fn = torch.nn.CrossEntropyLoss(weight=torch.FloatTensor([config.w_ce1, config.w_ce2]).to(device)) - - if args.reverse_mp: - model = train_hetero(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) - 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() \ No newline at end of file + + return model + +def train_gnn(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, args, data_config): + #set device + device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + + #define a model config dictionary and wandb logging at the same time + wandb.init( + mode="disabled" if args.testing else "online", + project="your_proj_name", #replace this with your wandb project name if you want to use wandb logging + + config={ + "epochs": args.n_epochs, + "batch_size": args.batch_size, + "model": args.model, + "data": args.data, + "num_neighbors": args.num_neighs, + "lr": extract_param("lr", args), + "n_hidden": extract_param("n_hidden", args), + "n_gnn_layers": extract_param("n_gnn_layers", args), + "loss": "ce", + "w_ce1": extract_param("w_ce1", args), + "w_ce2": extract_param("w_ce2", args), + "dropout": extract_param("dropout", args), + "final_dropout": extract_param("final_dropout", args), + "n_heads": extract_param("n_heads", args) if args.model == 'gat' else None + } + ) + + config = wandb.config + + #set the transform if ego ids should be used + if args.ego: + transform = AddEgoIds() + else: + transform = None + + #add the unique ids to later find the seed edges + add_arange_ids([tr_data, val_data, te_data]) + + tr_loader, val_loader, te_loader = get_loaders(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, transform, args) + + #get the model + sample_batch = next(iter(tr_loader)) + model = get_model(sample_batch, config, args) + + if args.reverse_mp: + model = to_hetero(model, te_data.metadata(), aggr='mean') + + if args.finetune: + model, optimizer = load_model(model, device, args, config, data_config) + else: + model.to(device) + optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) + + sample_batch.to(device) + sample_x = sample_batch.x if not isinstance(sample_batch, HeteroData) else sample_batch.x_dict + sample_edge_index = sample_batch.edge_index if not isinstance(sample_batch, HeteroData) else sample_batch.edge_index_dict + if isinstance(sample_batch, HeteroData): + sample_batch['node', 'to', 'node'].edge_attr = sample_batch['node', 'to', 'node'].edge_attr[:, 1:] + sample_batch['node', 'rev_to', 'node'].edge_attr = sample_batch['node', 'rev_to', 'node'].edge_attr[:, 1:] + else: + sample_batch.edge_attr = sample_batch.edge_attr[:, 1:] + sample_edge_attr = sample_batch.edge_attr if not isinstance(sample_batch, HeteroData) else sample_batch.edge_attr_dict + logging.info(summary(model, sample_x, sample_edge_index, sample_edge_attr)) + + loss_fn = torch.nn.CrossEntropyLoss(weight=torch.FloatTensor([config.w_ce1, config.w_ce2]).to(device)) + + if args.reverse_mp: + model = train_hetero(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) + 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() From 34370e9e052fc2fcd60b47000c63a134d585116d Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 17:52:06 +0800 Subject: [PATCH 08/10] Update tests/test_edge_feature_utils.py for edge feature fix Signed-off-by: benben951 From e56d4e2af9ba52d58c74f1a30ebf2208681ee2f9 Mon Sep 17 00:00:00 2001 From: benben951 Date: Sat, 11 Apr 2026 18:50:51 +0800 Subject: [PATCH 09/10] Add regression tests for reverse MP edge features Signed-off-by: benben951 --- tests/test_edge_feature_utils.py | 82 ++++++++++++++++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/tests/test_edge_feature_utils.py b/tests/test_edge_feature_utils.py index 8bda82f..e01093b 100644 --- a/tests/test_edge_feature_utils.py +++ b/tests/test_edge_feature_utils.py @@ -1,4 +1,8 @@ +import importlib.util +import sys +import types import unittest +from types import SimpleNamespace import torch @@ -12,6 +16,34 @@ ) +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]]) @@ -40,6 +72,56 @@ 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() From 6b6a66dbadaabad9a5ab1940d8615b4dd2fe8b30 Mon Sep 17 00:00:00 2001 From: benben951 Date: Mon, 25 May 2026 23:47:55 +0800 Subject: [PATCH 10/10] Normalize PR diff line endings --- data_loading.py | 259 ++++++++++++++------------- data_util.py | 264 ++++++++++++++-------------- models.py | 372 +++++++++++++++++++-------------------- training.py | 456 ++++++++++++++++++++++++------------------------ 4 files changed, 675 insertions(+), 676 deletions(-) diff --git a/data_loading.py b/data_loading.py index 3531181..05f8ebb 100644 --- a/data_loading.py +++ b/data_loading.py @@ -1,142 +1,141 @@ -import pandas as pd -import numpy as np +import pandas as pd +import numpy as np import torch import logging import itertools 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. - - 1. The data is loaded from the csv and the necessary features are chosen. - 2. The data is split into training, validation and test data. - 3. PyG Data objects are created with the respective data splits. - ''' - - transaction_file = f"{data_config['paths']['aml_data']}/{args.data}/formatted_transactions.csv" #replace this with your path to the respective AML data objects - df_edges = pd.read_csv(transaction_file) - - logging.info(f'Available Edge Features: {df_edges.columns.tolist()}') - - df_edges['Timestamp'] = df_edges['Timestamp'] - df_edges['Timestamp'].min() - - max_n_id = df_edges.loc[:, ['from_id', 'to_id']].to_numpy().max() + 1 - df_nodes = pd.DataFrame({'NodeID': np.arange(max_n_id), 'Feature': np.ones(max_n_id)}) - timestamps = torch.Tensor(df_edges['Timestamp'].to_numpy()) - y = torch.LongTensor(df_edges['Is Laundering'].to_numpy()) - - logging.info(f"Illicit ratio = {sum(y)} / {len(y)} = {sum(y) / len(y) * 100:.2f}%") - logging.info(f"Number of nodes (holdings doing transcations) = {df_nodes.shape[0]}") - logging.info(f"Number of transactions = {df_edges.shape[0]}") - - edge_features = ['Timestamp', 'Amount Received', 'Received Currency', 'Payment Format'] - node_features = ['Feature'] - - logging.info(f'Edge features being used: {edge_features}') - logging.info(f'Node features being used: {node_features} ("Feature" is a placeholder feature of all 1s)') - - x = torch.tensor(df_nodes.loc[:, node_features].to_numpy()).float() - edge_index = torch.LongTensor(df_edges.loc[:, ['from_id', 'to_id']].to_numpy().T) - edge_attr = torch.tensor(df_edges.loc[:, edge_features].to_numpy()).float() - - n_days = int(timestamps.max() / (3600 * 24) + 1) - n_samples = y.shape[0] - logging.info(f'number of days and transactions in the data: {n_days} days, {n_samples} transactions') - - #data splitting - daily_irs, weighted_daily_irs, daily_inds, daily_trans = [], [], [], [] #irs = illicit ratios, inds = indices, trans = transactions - for day in range(n_days): - l = day * 24 * 3600 - r = (day + 1) * 24 * 3600 - day_inds = torch.where((timestamps >= l) & (timestamps < r))[0] - daily_irs.append(y[day_inds].float().mean()) - weighted_daily_irs.append(y[day_inds].float().mean() * day_inds.shape[0] / n_samples) - daily_inds.append(day_inds) - daily_trans.append(day_inds.shape[0]) - - split_per = [0.6, 0.2, 0.2] - daily_totals = np.array(daily_trans) - d_ts = daily_totals - I = list(range(len(d_ts))) - split_scores = dict() - for i,j in itertools.combinations(I, 2): - if j >= i: - split_totals = [d_ts[:i].sum(), d_ts[i:j].sum(), d_ts[j:].sum()] - split_totals_sum = np.sum(split_totals) - split_props = [v/split_totals_sum for v in split_totals] - split_error = [abs(v-t)/t for v,t in zip(split_props, split_per)] - score = max(split_error) #- (split_totals_sum/total) + 1 - split_scores[(i,j)] = score - else: - continue - - i,j = min(split_scores, key=split_scores.get) - #split contains a list for each split (train, validation and test) and each list contains the days that are part of the respective split - split = [list(range(i)), list(range(i, j)), list(range(j, len(daily_totals)))] - logging.info(f'Calculate split: {split}') - - #Now, we seperate the transactions based on their indices in the timestamp array - split_inds = {k: [] for k in range(3)} - for i in range(3): - for day in split[i]: - split_inds[i].append(daily_inds[day]) #split_inds contains a list for each split (tr,val,te) which contains the indices of each day seperately - - tr_inds = torch.cat(split_inds[0]) - val_inds = torch.cat(split_inds[1]) - te_inds = torch.cat(split_inds[2]) - - logging.info(f"Total train samples: {tr_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[tr_inds].float().mean() * 100 :.2f}% || Train days: {split[0][:5]}") - logging.info(f"Total val samples: {val_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[val_inds].float().mean() * 100:.2f}% || Val days: {split[1][:5]}") - logging.info(f"Total test samples: {te_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " - f"{y[te_inds].float().mean() * 100:.2f}% || Test days: {split[2][:5]}") - - #Creating the final data objects - tr_x, val_x, te_x = x, x, x - e_tr = tr_inds.numpy() - e_val = np.concatenate([tr_inds, val_inds]) - - tr_edge_index, tr_edge_attr, tr_y, tr_edge_times = edge_index[:,e_tr], edge_attr[e_tr], y[e_tr], timestamps[e_tr] - val_edge_index, val_edge_attr, val_y, val_edge_times = edge_index[:,e_val], edge_attr[e_val], y[e_val], timestamps[e_val] - te_edge_index, te_edge_attr, te_y, te_edge_times = edge_index, edge_attr, y, timestamps - - tr_data = GraphData (x=tr_x, y=tr_y, edge_index=tr_edge_index, edge_attr=tr_edge_attr, timestamps=tr_edge_times ) - val_data = GraphData(x=val_x, y=val_y, edge_index=val_edge_index, edge_attr=val_edge_attr, timestamps=val_edge_times) - te_data = GraphData (x=te_x, y=te_y, edge_index=te_edge_index, edge_attr=te_edge_attr, timestamps=te_edge_times ) - - #Adding ports and time-deltas if applicable - if args.ports: - logging.info(f"Start: adding ports") - tr_data.add_ports() - val_data.add_ports() - te_data.add_ports() - logging.info(f"Done: adding ports") - if args.tds: - logging.info(f"Start: adding time-deltas") - tr_data.add_time_deltas() - val_data.add_time_deltas() - te_data.add_time_deltas() - logging.info(f"Done: adding time-deltas") - + +def get_data(args, data_config): + '''Loads the AML transaction data. + + 1. The data is loaded from the csv and the necessary features are chosen. + 2. The data is split into training, validation and test data. + 3. PyG Data objects are created with the respective data splits. + ''' + + transaction_file = f"{data_config['paths']['aml_data']}/{args.data}/formatted_transactions.csv" #replace this with your path to the respective AML data objects + df_edges = pd.read_csv(transaction_file) + + logging.info(f'Available Edge Features: {df_edges.columns.tolist()}') + + df_edges['Timestamp'] = df_edges['Timestamp'] - df_edges['Timestamp'].min() + + max_n_id = df_edges.loc[:, ['from_id', 'to_id']].to_numpy().max() + 1 + df_nodes = pd.DataFrame({'NodeID': np.arange(max_n_id), 'Feature': np.ones(max_n_id)}) + timestamps = torch.Tensor(df_edges['Timestamp'].to_numpy()) + y = torch.LongTensor(df_edges['Is Laundering'].to_numpy()) + + logging.info(f"Illicit ratio = {sum(y)} / {len(y)} = {sum(y) / len(y) * 100:.2f}%") + logging.info(f"Number of nodes (holdings doing transcations) = {df_nodes.shape[0]}") + logging.info(f"Number of transactions = {df_edges.shape[0]}") + + edge_features = ['Timestamp', 'Amount Received', 'Received Currency', 'Payment Format'] + node_features = ['Feature'] + + logging.info(f'Edge features being used: {edge_features}') + logging.info(f'Node features being used: {node_features} ("Feature" is a placeholder feature of all 1s)') + + x = torch.tensor(df_nodes.loc[:, node_features].to_numpy()).float() + edge_index = torch.LongTensor(df_edges.loc[:, ['from_id', 'to_id']].to_numpy().T) + edge_attr = torch.tensor(df_edges.loc[:, edge_features].to_numpy()).float() + + n_days = int(timestamps.max() / (3600 * 24) + 1) + n_samples = y.shape[0] + logging.info(f'number of days and transactions in the data: {n_days} days, {n_samples} transactions') + + #data splitting + daily_irs, weighted_daily_irs, daily_inds, daily_trans = [], [], [], [] #irs = illicit ratios, inds = indices, trans = transactions + for day in range(n_days): + l = day * 24 * 3600 + r = (day + 1) * 24 * 3600 + day_inds = torch.where((timestamps >= l) & (timestamps < r))[0] + daily_irs.append(y[day_inds].float().mean()) + weighted_daily_irs.append(y[day_inds].float().mean() * day_inds.shape[0] / n_samples) + daily_inds.append(day_inds) + daily_trans.append(day_inds.shape[0]) + + split_per = [0.6, 0.2, 0.2] + daily_totals = np.array(daily_trans) + d_ts = daily_totals + I = list(range(len(d_ts))) + split_scores = dict() + for i,j in itertools.combinations(I, 2): + if j >= i: + split_totals = [d_ts[:i].sum(), d_ts[i:j].sum(), d_ts[j:].sum()] + split_totals_sum = np.sum(split_totals) + split_props = [v/split_totals_sum for v in split_totals] + split_error = [abs(v-t)/t for v,t in zip(split_props, split_per)] + score = max(split_error) #- (split_totals_sum/total) + 1 + split_scores[(i,j)] = score + else: + continue + + i,j = min(split_scores, key=split_scores.get) + #split contains a list for each split (train, validation and test) and each list contains the days that are part of the respective split + split = [list(range(i)), list(range(i, j)), list(range(j, len(daily_totals)))] + logging.info(f'Calculate split: {split}') + + #Now, we seperate the transactions based on their indices in the timestamp array + split_inds = {k: [] for k in range(3)} + for i in range(3): + for day in split[i]: + split_inds[i].append(daily_inds[day]) #split_inds contains a list for each split (tr,val,te) which contains the indices of each day seperately + + tr_inds = torch.cat(split_inds[0]) + val_inds = torch.cat(split_inds[1]) + te_inds = torch.cat(split_inds[2]) + + logging.info(f"Total train samples: {tr_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[tr_inds].float().mean() * 100 :.2f}% || Train days: {split[0][:5]}") + logging.info(f"Total val samples: {val_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[val_inds].float().mean() * 100:.2f}% || Val days: {split[1][:5]}") + logging.info(f"Total test samples: {te_inds.shape[0] / y.shape[0] * 100 :.2f}% || IR: " + f"{y[te_inds].float().mean() * 100:.2f}% || Test days: {split[2][:5]}") + + #Creating the final data objects + tr_x, val_x, te_x = x, x, x + e_tr = tr_inds.numpy() + e_val = np.concatenate([tr_inds, val_inds]) + + tr_edge_index, tr_edge_attr, tr_y, tr_edge_times = edge_index[:,e_tr], edge_attr[e_tr], y[e_tr], timestamps[e_tr] + val_edge_index, val_edge_attr, val_y, val_edge_times = edge_index[:,e_val], edge_attr[e_val], y[e_val], timestamps[e_val] + te_edge_index, te_edge_attr, te_y, te_edge_times = edge_index, edge_attr, y, timestamps + + tr_data = GraphData (x=tr_x, y=tr_y, edge_index=tr_edge_index, edge_attr=tr_edge_attr, timestamps=tr_edge_times ) + val_data = GraphData(x=val_x, y=val_y, edge_index=val_edge_index, edge_attr=val_edge_attr, timestamps=val_edge_times) + te_data = GraphData (x=te_x, y=te_y, edge_index=te_edge_index, edge_attr=te_edge_attr, timestamps=te_edge_times ) + + #Adding ports and time-deltas if applicable + if args.ports: + logging.info(f"Start: adding ports") + tr_data.add_ports() + val_data.add_ports() + te_data.add_ports() + logging.info(f"Done: adding ports") + if args.tds: + logging.info(f"Start: adding time-deltas") + tr_data.add_time_deltas() + val_data.add_time_deltas() + te_data.add_time_deltas() + logging.info(f"Done: adding time-deltas") + #Normalize data tr_data.x = val_data.x = te_data.x = z_norm(tr_data.x) 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 - if args.reverse_mp: - tr_data = create_hetero_obj(tr_data.x, tr_data.y, tr_data.edge_index, tr_data.edge_attr, tr_data.timestamps, args) - val_data = create_hetero_obj(val_data.x, val_data.y, val_data.edge_index, val_data.edge_attr, val_data.timestamps, args) - te_data = create_hetero_obj(te_data.x, te_data.y, te_data.edge_index, te_data.edge_attr, te_data.timestamps, args) - - logging.info(f'train data object: {tr_data}') - logging.info(f'validation data object: {val_data}') - logging.info(f'test data object: {te_data}') - - return tr_data, val_data, te_data, tr_inds, val_inds, te_inds + + #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 + if args.reverse_mp: + tr_data = create_hetero_obj(tr_data.x, tr_data.y, tr_data.edge_index, tr_data.edge_attr, tr_data.timestamps, args) + val_data = create_hetero_obj(val_data.x, val_data.y, val_data.edge_index, val_data.edge_attr, val_data.timestamps, args) + te_data = create_hetero_obj(te_data.x, te_data.y, te_data.edge_index, te_data.edge_attr, te_data.timestamps, args) + logging.info(f'train data object: {tr_data}') + logging.info(f'validation data object: {val_data}') + logging.info(f'test data object: {te_data}') + + return tr_data, val_data, te_data, tr_inds, val_inds, te_inds diff --git a/data_util.py b/data_util.py index 954f6e7..56e7ee4 100644 --- a/data_util.py +++ b/data_util.py @@ -3,141 +3,141 @@ 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 - timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) - edges = torch.cat((data.edge_index.T, timestamps), dim=1) if not isinstance(data, HeteroData) else torch.cat((data['node', 'to', 'node'].edge_index.T, timestamps), dim=1) - adj_list_out = dict([(i, []) for i in range(num_nodes)]) - adj_list_in = dict([(i, []) for i in range(num_nodes)]) - for u,v,t in edges: - u,v,t = int(u), int(v), int(t) - adj_list_out[u] += [(v, t)] - adj_list_in[v] += [(u, t)] - return adj_list_in, adj_list_out - -def to_adj_edges_with_times(data): - num_nodes = data.num_nodes - timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) - edges = torch.cat((data.edge_index.T, timestamps), dim=1) - # calculate adjacent edges with times per node - adj_edges_out = dict([(i, []) for i in range(num_nodes)]) - adj_edges_in = dict([(i, []) for i in range(num_nodes)]) - for i, (u,v,t) in enumerate(edges): - u,v,t = int(u), int(v), int(t) - adj_edges_out[u] += [(i, v, t)] - adj_edges_in[v] += [(i, u, t)] - return adj_edges_in, adj_edges_out - -def ports(edge_index, adj_list): - ports = torch.zeros(edge_index.shape[1], 1) - ports_dict = {} - for v, nbs in adj_list.items(): - if len(nbs) < 1: continue - a = np.array(nbs) - a = a[a[:, -1].argsort()] - _, idx = np.unique(a[:,[0]],return_index=True,axis=0) - nbs_unique = a[np.sort(idx)][:,0] - for i, u in enumerate(nbs_unique): - ports_dict[(u,v)] = i - for i, e in enumerate(edge_index.T): - ports[i] = ports_dict[tuple(e.numpy())] - return ports - -def time_deltas(data, adj_edges_list): - time_deltas = torch.zeros(data.edge_index.shape[1], 1) - if data.timestamps is None: - return time_deltas - for v, edges in adj_edges_list.items(): - if len(edges) < 1: continue - a = np.array(edges) - a = a[a[:, -1].argsort()] - a_tds = [0] + [a[i+1,-1] - a[i,-1] for i in range(a.shape[0]-1)] - tds = np.hstack((a[:,0].reshape(-1,1), np.array(a_tds).reshape(-1,1))) - for i,td in tds: - time_deltas[i] = td - return time_deltas - -class GraphData(Data): - '''This is the homogenous graph object we use for GNN training if reverse MP is not enabled''' - def __init__( - self, x: OptTensor = None, edge_index: OptTensor = None, edge_attr: OptTensor = None, y: OptTensor = None, pos: OptTensor = None, - readout: str = 'edge', - num_nodes: int = None, - timestamps: OptTensor = None, - node_timestamps: OptTensor = None, - **kwargs - ): - super().__init__(x, edge_index, edge_attr, y, pos, **kwargs) - self.readout = readout - self.loss_fn = 'ce' - self.num_nodes = int(self.x.shape[0]) - self.node_timestamps = node_timestamps - if timestamps is not None: - self.timestamps = timestamps - elif edge_attr is not None: - self.timestamps = edge_attr[:,0].clone() - else: - self.timestamps = None - - def add_ports(self): - '''Adds port numberings to the edge features''' - reverse_ports = True - adj_list_in, adj_list_out = to_adj_nodes_with_times(self) - in_ports = ports(self.edge_index, adj_list_in) - out_ports = [ports(self.edge_index.flipud(), adj_list_out)] if reverse_ports else [] - self.edge_attr = torch.cat([self.edge_attr, in_ports] + out_ports, dim=1) - return self - - def add_time_deltas(self): - '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' - reverse_tds = True - adj_list_in, adj_list_out = to_adj_edges_with_times(self) - in_tds = time_deltas(self, adj_list_in) - out_tds = [time_deltas(self, adj_list_out)] if reverse_tds else [] - self.edge_attr = torch.cat([self.edge_attr, in_tds] + out_tds, dim=1) - return self - -class HeteroGraphData(HeteroData): - '''This is the heterogenous graph object we use for GNN training if reverse MP is enabled''' - def __init__( - self, - readout: str = 'edge', - **kwargs - ): - super().__init__(**kwargs) - self.readout = readout - - @property - def num_nodes(self): - return self['node'].x.shape[0] - - @property - def timestamps(self): - return self['node', 'to', 'node'].timestamps - - def add_ports(self): - '''Adds port numberings to the edge features''' - adj_list_in, adj_list_out = to_adj_nodes_with_times(self) - in_ports = ports(self['node', 'to', 'node'].edge_index, adj_list_in) - out_ports = ports(self['node', 'rev_to', 'node'].edge_index, adj_list_out) - self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_ports], dim=1) - self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_ports], dim=1) - return self - - def add_time_deltas(self): - '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' - adj_list_in, adj_list_out = to_adj_edges_with_times(self) - in_tds = time_deltas(self, adj_list_in) - out_tds = time_deltas(self, adj_list_out) - self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_tds], dim=1) - self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_tds], dim=1) - return self - + +def to_adj_nodes_with_times(data): + num_nodes = data.num_nodes + timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) + edges = torch.cat((data.edge_index.T, timestamps), dim=1) if not isinstance(data, HeteroData) else torch.cat((data['node', 'to', 'node'].edge_index.T, timestamps), dim=1) + adj_list_out = dict([(i, []) for i in range(num_nodes)]) + adj_list_in = dict([(i, []) for i in range(num_nodes)]) + for u,v,t in edges: + u,v,t = int(u), int(v), int(t) + adj_list_out[u] += [(v, t)] + adj_list_in[v] += [(u, t)] + return adj_list_in, adj_list_out + +def to_adj_edges_with_times(data): + num_nodes = data.num_nodes + timestamps = torch.zeros((data.edge_index.shape[1], 1)) if data.timestamps is None else data.timestamps.reshape((-1,1)) + edges = torch.cat((data.edge_index.T, timestamps), dim=1) + # calculate adjacent edges with times per node + adj_edges_out = dict([(i, []) for i in range(num_nodes)]) + adj_edges_in = dict([(i, []) for i in range(num_nodes)]) + for i, (u,v,t) in enumerate(edges): + u,v,t = int(u), int(v), int(t) + adj_edges_out[u] += [(i, v, t)] + adj_edges_in[v] += [(i, u, t)] + return adj_edges_in, adj_edges_out + +def ports(edge_index, adj_list): + ports = torch.zeros(edge_index.shape[1], 1) + ports_dict = {} + for v, nbs in adj_list.items(): + if len(nbs) < 1: continue + a = np.array(nbs) + a = a[a[:, -1].argsort()] + _, idx = np.unique(a[:,[0]],return_index=True,axis=0) + nbs_unique = a[np.sort(idx)][:,0] + for i, u in enumerate(nbs_unique): + ports_dict[(u,v)] = i + for i, e in enumerate(edge_index.T): + ports[i] = ports_dict[tuple(e.numpy())] + return ports + +def time_deltas(data, adj_edges_list): + time_deltas = torch.zeros(data.edge_index.shape[1], 1) + if data.timestamps is None: + return time_deltas + for v, edges in adj_edges_list.items(): + if len(edges) < 1: continue + a = np.array(edges) + a = a[a[:, -1].argsort()] + a_tds = [0] + [a[i+1,-1] - a[i,-1] for i in range(a.shape[0]-1)] + tds = np.hstack((a[:,0].reshape(-1,1), np.array(a_tds).reshape(-1,1))) + for i,td in tds: + time_deltas[i] = td + return time_deltas + +class GraphData(Data): + '''This is the homogenous graph object we use for GNN training if reverse MP is not enabled''' + def __init__( + self, x: OptTensor = None, edge_index: OptTensor = None, edge_attr: OptTensor = None, y: OptTensor = None, pos: OptTensor = None, + readout: str = 'edge', + num_nodes: int = None, + timestamps: OptTensor = None, + node_timestamps: OptTensor = None, + **kwargs + ): + super().__init__(x, edge_index, edge_attr, y, pos, **kwargs) + self.readout = readout + self.loss_fn = 'ce' + self.num_nodes = int(self.x.shape[0]) + self.node_timestamps = node_timestamps + if timestamps is not None: + self.timestamps = timestamps + elif edge_attr is not None: + self.timestamps = edge_attr[:,0].clone() + else: + self.timestamps = None + + def add_ports(self): + '''Adds port numberings to the edge features''' + reverse_ports = True + adj_list_in, adj_list_out = to_adj_nodes_with_times(self) + in_ports = ports(self.edge_index, adj_list_in) + out_ports = [ports(self.edge_index.flipud(), adj_list_out)] if reverse_ports else [] + self.edge_attr = torch.cat([self.edge_attr, in_ports] + out_ports, dim=1) + return self + + def add_time_deltas(self): + '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' + reverse_tds = True + adj_list_in, adj_list_out = to_adj_edges_with_times(self) + in_tds = time_deltas(self, adj_list_in) + out_tds = [time_deltas(self, adj_list_out)] if reverse_tds else [] + self.edge_attr = torch.cat([self.edge_attr, in_tds] + out_tds, dim=1) + return self + +class HeteroGraphData(HeteroData): + '''This is the heterogenous graph object we use for GNN training if reverse MP is enabled''' + def __init__( + self, + readout: str = 'edge', + **kwargs + ): + super().__init__(**kwargs) + self.readout = readout + + @property + def num_nodes(self): + return self['node'].x.shape[0] + + @property + def timestamps(self): + return self['node', 'to', 'node'].timestamps + + def add_ports(self): + '''Adds port numberings to the edge features''' + adj_list_in, adj_list_out = to_adj_nodes_with_times(self) + in_ports = ports(self['node', 'to', 'node'].edge_index, adj_list_in) + out_ports = ports(self['node', 'rev_to', 'node'].edge_index, adj_list_out) + self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_ports], dim=1) + self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_ports], dim=1) + return self + + def add_time_deltas(self): + '''Adds time deltas (i.e. the time between subsequent transactions) to the edge features''' + adj_list_in, adj_list_out = to_adj_edges_with_times(self) + in_tds = time_deltas(self, adj_list_in) + out_tds = time_deltas(self, adj_list_out) + self['node', 'to', 'node'].edge_attr = torch.cat([self['node', 'to', 'node'].edge_attr, in_tds], dim=1) + self['node', 'rev_to', 'node'].edge_attr = torch.cat([self['node', 'rev_to', 'node'].edge_attr, out_tds], dim=1) + return self + 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() diff --git a/models.py b/models.py index cee87d7..b851e47 100644 --- a/models.py +++ b/models.py @@ -1,46 +1,46 @@ -import torch.nn as nn -from torch_geometric.nn import GINEConv, BatchNorm, Linear, GATConv, PNAConv, RGCNConv +import torch.nn as nn +from torch_geometric.nn import GINEConv, BatchNorm, Linear, GATConv, PNAConv, RGCNConv 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, - n_hidden=100, edge_updates=False, residual=True, - edge_dim=None, dropout=0.0, final_dropout=0.5): - super().__init__() - self.n_hidden = n_hidden - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.final_dropout = final_dropout - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - for _ in range(self.num_gnn_layers): - conv = GINEConv(nn.Sequential( - nn.Linear(self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden) - ), edge_dim=self.n_hidden) - if self.edge_updates: self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - + +class GINe(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, + n_hidden=100, edge_updates=False, residual=True, + edge_dim=None, dropout=0.0, final_dropout=0.5): + super().__init__() + self.n_hidden = n_hidden + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.final_dropout = final_dropout + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + for _ in range(self.num_gnn_layers): + conv = GINEConv(nn.Sequential( + nn.Linear(self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden) + ), edge_dim=self.n_hidden) + if self.edge_updates: self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + src, dst = edge_index + x = self.node_emb(x) edge_attr = self.edge_emb(edge_attr) @@ -48,45 +48,45 @@ def forward(self, x, edge_index, edge_attr): 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 = 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) - out = x - - return self.mlp(out) - -class GATe(torch.nn.Module): - def __init__(self, num_features, num_gnn_layers, n_classes=2, n_hidden=100, n_heads=4, edge_updates=False, edge_dim=None, dropout=0.0, final_dropout=0.5): - super().__init__() - # GAT specific code - tmp_out = n_hidden // n_heads - n_hidden = tmp_out * n_heads - - self.n_hidden = n_hidden - self.n_heads = n_heads - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.dropout = dropout - self.final_dropout = final_dropout - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - - for _ in range(self.num_gnn_layers): - conv = GATConv(self.n_hidden, tmp_out, self.n_heads, concat = True, dropout = self.dropout, add_self_loops = True, edge_dim=self.n_hidden) - if self.edge_updates: self.emlps.append(nn.Sequential(nn.Linear(3 * self.n_hidden, self.n_hidden),nn.ReLU(),nn.Linear(self.n_hidden, self.n_hidden),)) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - + + 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) + out = x + + return self.mlp(out) + +class GATe(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, n_hidden=100, n_heads=4, edge_updates=False, edge_dim=None, dropout=0.0, final_dropout=0.5): + super().__init__() + # GAT specific code + tmp_out = n_hidden // n_heads + n_hidden = tmp_out * n_heads + + self.n_hidden = n_hidden + self.n_heads = n_heads + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.dropout = dropout + self.final_dropout = final_dropout + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + + for _ in range(self.num_gnn_layers): + conv = GATConv(self.n_hidden, tmp_out, self.n_heads, concat = True, dropout = self.dropout, add_self_loops = True, edge_dim=self.n_hidden) + if self.edge_updates: self.emlps.append(nn.Sequential(nn.Linear(3 * self.n_hidden, self.n_hidden),nn.ReLU(),nn.Linear(self.n_hidden, self.n_hidden),)) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + src, dst = edge_index + x = self.node_emb(x) edge_attr = self.edge_emb(edge_attr) @@ -94,55 +94,55 @@ def forward(self, x, edge_index, edge_attr): 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 = 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() - logging.debug(f"x.shape = {x.shape}") - x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) - logging.debug(f"x.shape = {x.shape}") - out = x - - return self.mlp(out) - -class PNA(torch.nn.Module): - def __init__(self, num_features, num_gnn_layers, n_classes=2, - n_hidden=100, edge_updates=True, - edge_dim=None, dropout=0.0, final_dropout=0.5, deg=None): - super().__init__() - n_hidden = int((n_hidden // 5) * 5) - self.n_hidden = n_hidden - self.num_gnn_layers = num_gnn_layers - self.edge_updates = edge_updates - self.final_dropout = final_dropout - - aggregators = ['mean', 'min', 'max', 'std'] - scalers = ['identity', 'amplification', 'attenuation'] - - self.node_emb = nn.Linear(num_features, n_hidden) - self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.emlps = nn.ModuleList() - self.batch_norms = nn.ModuleList() - for _ in range(self.num_gnn_layers): - conv = PNAConv(in_channels=n_hidden, out_channels=n_hidden, - aggregators=aggregators, scalers=scalers, deg=deg, - edge_dim=n_hidden, towers=5, pre_layers=1, post_layers=1, - divide_input=False) - if self.edge_updates: self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - self.convs.append(conv) - self.batch_norms.append(BatchNorm(n_hidden)) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def forward(self, x, edge_index, edge_attr): - src, dst = edge_index - + + 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() + logging.debug(f"x.shape = {x.shape}") + x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) + logging.debug(f"x.shape = {x.shape}") + out = x + + return self.mlp(out) + +class PNA(torch.nn.Module): + def __init__(self, num_features, num_gnn_layers, n_classes=2, + n_hidden=100, edge_updates=True, + edge_dim=None, dropout=0.0, final_dropout=0.5, deg=None): + super().__init__() + n_hidden = int((n_hidden // 5) * 5) + self.n_hidden = n_hidden + self.num_gnn_layers = num_gnn_layers + self.edge_updates = edge_updates + self.final_dropout = final_dropout + + aggregators = ['mean', 'min', 'max', 'std'] + scalers = ['identity', 'amplification', 'attenuation'] + + self.node_emb = nn.Linear(num_features, n_hidden) + self.edge_emb = nn.Linear(edge_dim, n_hidden) + + self.convs = nn.ModuleList() + self.emlps = nn.ModuleList() + self.batch_norms = nn.ModuleList() + for _ in range(self.num_gnn_layers): + conv = PNAConv(in_channels=n_hidden, out_channels=n_hidden, + aggregators=aggregators, scalers=scalers, deg=deg, + edge_dim=n_hidden, towers=5, pre_layers=1, post_layers=1, + divide_input=False) + if self.edge_updates: self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + self.convs.append(conv) + self.batch_norms.append(BatchNorm(n_hidden)) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout),Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def forward(self, x, edge_index, edge_attr): + src, dst = edge_index + x = self.node_emb(x) edge_attr = self.edge_emb(edge_attr) @@ -150,28 +150,28 @@ def forward(self, x, edge_index, edge_attr): 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 = 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() - logging.debug(f"x.shape = {x.shape}") - x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) - logging.debug(f"x.shape = {x.shape}") - out = x - return self.mlp(out) - + + 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() + logging.debug(f"x.shape = {x.shape}") + x = torch.cat((x, edge_attr.view(-1, edge_attr.shape[1])), 1) + logging.debug(f"x.shape = {x.shape}") + out = x + return self.mlp(out) + 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, edge_type_index=3): super(RGCN, self).__init__() - - self.num_features = num_features - self.num_gnn_layers = num_gnn_layers - self.n_hidden = n_hidden - self.residual = residual - self.dropout = dropout - self.final_dropout = final_dropout + + self.num_features = num_features + self.num_gnn_layers = num_gnn_layers + self.n_hidden = n_hidden + self.residual = residual + self.dropout = dropout + self.final_dropout = final_dropout self.n_classes = n_classes self.edge_update = edge_update self.num_relations = num_relations @@ -180,43 +180,43 @@ def __init__(self, num_features, edge_dim, num_relations, num_gnn_layers, n_clas self.node_emb = nn.Linear(num_features, n_hidden) self.edge_emb = nn.Linear(edge_dim, n_hidden) - - self.convs = nn.ModuleList() - self.bns = nn.ModuleList() - self.mlp = nn.ModuleList() - - if self.edge_update: - self.emlps = nn.ModuleList() - self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - - for _ in range(self.num_gnn_layers): - conv = RGCNConv(self.n_hidden, self.n_hidden, num_relations, num_bases=self.n_bases) - self.convs.append(conv) - self.bns.append(nn.BatchNorm1d(self.n_hidden)) - - if self.edge_update: - self.emlps.append(nn.Sequential( - nn.Linear(3 * self.n_hidden, self.n_hidden), - nn.ReLU(), - nn.Linear(self.n_hidden, self.n_hidden), - )) - - self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout), Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), - Linear(25, n_classes)) - - def reset_parameters(self): - for m in self.modules(): - if isinstance(m, nn.Linear): - m.reset_parameters() - elif isinstance(m, RGCNConv): - m.reset_parameters() - elif isinstance(m, nn.BatchNorm1d): - m.reset_parameters() - + + self.convs = nn.ModuleList() + self.bns = nn.ModuleList() + self.mlp = nn.ModuleList() + + if self.edge_update: + self.emlps = nn.ModuleList() + self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + + for _ in range(self.num_gnn_layers): + conv = RGCNConv(self.n_hidden, self.n_hidden, num_relations, num_bases=self.n_bases) + self.convs.append(conv) + self.bns.append(nn.BatchNorm1d(self.n_hidden)) + + if self.edge_update: + self.emlps.append(nn.Sequential( + nn.Linear(3 * self.n_hidden, self.n_hidden), + nn.ReLU(), + nn.Linear(self.n_hidden, self.n_hidden), + )) + + self.mlp = nn.Sequential(Linear(n_hidden*3, 50), nn.ReLU(), nn.Dropout(self.final_dropout), Linear(50, 25), nn.ReLU(), nn.Dropout(self.final_dropout), + Linear(25, n_classes)) + + def reset_parameters(self): + for m in self.modules(): + if isinstance(m, nn.Linear): + m.reset_parameters() + elif isinstance(m, RGCNConv): + m.reset_parameters() + elif isinstance(m, nn.BatchNorm1d): + m.reset_parameters() + def forward(self, x, edge_index, edge_attr): 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) @@ -229,10 +229,10 @@ def forward(self, x, edge_index, edge_attr): x = average_residual_update(x, F.relu(self.bns[i](self.convs[i](x, edge_index, edge_type)))) if self.edge_update: 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 - + + 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 diff --git a/training.py b/training.py index 45388b3..6bf78a4 100644 --- a/training.py +++ b/training.py @@ -1,160 +1,160 @@ -import torch -import tqdm -from sklearn.metrics import f1_score -from train_util import AddEgoIds, extract_param, add_arange_ids, get_loaders, evaluate_homo, evaluate_hetero, save_model, load_model -from models import GINe, PNA, GATe, RGCN -from torch_geometric.data import Data, HeteroData -from torch_geometric.nn import to_hetero, summary -from torch_geometric.utils import degree +import torch +import tqdm +from sklearn.metrics import f1_score +from train_util import AddEgoIds, extract_param, add_arange_ids, get_loaders, evaluate_homo, evaluate_hetero, save_model, load_model +from models import GINe, PNA, GATe, RGCN +from torch_geometric.data import Data, HeteroData +from torch_geometric.nn import to_hetero, summary +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 - best_val_f1 = 0 - for epoch in range(config.epochs): - total_loss = total_examples = 0 - preds = [] - ground_truths = [] - for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): - optimizer.zero_grad() - #select the seed edges from which the batch was created - inds = tr_inds.detach().cpu() - batch_edge_inds = inds[batch.input_id.detach().cpu()] - batch_edge_ids = tr_loader.data.edge_attr.detach().cpu()[batch_edge_inds, 0] - mask = torch.isin(batch.edge_attr[:, 0].detach().cpu(), batch_edge_ids) - - #remove the unique edge id from the edge features, as it's no longer needed - batch.edge_attr = batch.edge_attr[:, 1:] - - batch.to(device) - out = model(batch.x, batch.edge_index, batch.edge_attr) - pred = out[mask] - ground_truth = batch.y[mask] - preds.append(pred.argmax(dim=-1)) - ground_truths.append(ground_truth) - loss = loss_fn(pred, ground_truth) - - loss.backward() - optimizer.step() - - total_loss += float(loss) * pred.numel() - total_examples += pred.numel() - - pred = torch.cat(preds, dim=0).detach().cpu().numpy() - ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() - f1 = f1_score(ground_truth, pred) - wandb.log({"f1/train": f1}, step=epoch) - logging.info(f'Train F1: {f1:.4f}') - - #evaluate - val_f1 = evaluate_homo(val_loader, val_inds, model, val_data, device, args) - te_f1 = evaluate_homo(te_loader, te_inds, model, te_data, device, args) - - wandb.log({"f1/validation": val_f1}, step=epoch) - wandb.log({"f1/test": te_f1}, step=epoch) - logging.info(f'Validation F1: {val_f1:.4f}') - logging.info(f'Test F1: {te_f1:.4f}') - - if epoch == 0: - wandb.log({"best_test_f1": te_f1}, step=epoch) - elif val_f1 > best_val_f1: - best_val_f1 = val_f1 - wandb.log({"best_test_f1": te_f1}, step=epoch) - if args.save_model: - save_model(model, optimizer, epoch, args, data_config) - - return model - -def train_hetero(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 - best_val_f1 = 0 - for epoch in range(config.epochs): - total_loss = total_examples = 0 - preds = [] - ground_truths = [] - for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): - optimizer.zero_grad() - #select the seed edges from which the batch was created - inds = tr_inds.detach().cpu() - batch_edge_inds = inds[batch['node', 'to', 'node'].input_id.detach().cpu()] - batch_edge_ids = tr_loader.data['node', 'to', 'node'].edge_attr.detach().cpu()[batch_edge_inds, 0] - mask = torch.isin(batch['node', 'to', 'node'].edge_attr[:, 0].detach().cpu(), batch_edge_ids) - - #remove the unique edge id from the edge features, as it's no longer needed - batch['node', 'to', 'node'].edge_attr = batch['node', 'to', 'node'].edge_attr[:, 1:] - batch['node', 'rev_to', 'node'].edge_attr = batch['node', 'rev_to', 'node'].edge_attr[:, 1:] - - batch.to(device) - out = model(batch.x_dict, batch.edge_index_dict, batch.edge_attr_dict) - out = out[('node', 'to', 'node')] - pred = out[mask] - ground_truth = batch['node', 'to', 'node'].y[mask] - preds.append(pred.argmax(dim=-1)) - ground_truths.append(batch['node', 'to', 'node'].y[mask]) - loss = loss_fn(pred, ground_truth) - - loss.backward() - optimizer.step() - - total_loss += float(loss) * pred.numel() - total_examples += pred.numel() - - pred = torch.cat(preds, dim=0).detach().cpu().numpy() - ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() - f1 = f1_score(ground_truth, pred) - wandb.log({"f1/train": f1}, step=epoch) - logging.info(f'Train F1: {f1:.4f}') - - #evaluate - val_f1 = evaluate_hetero(val_loader, val_inds, model, val_data, device, args) - te_f1 = evaluate_hetero(te_loader, te_inds, model, te_data, device, args) - - wandb.log({"f1/validation": val_f1}, step=epoch) - wandb.log({"f1/test": te_f1}, step=epoch) - logging.info(f'Validation F1: {val_f1:.4f}') - logging.info(f'Test F1: {te_f1:.4f}') - - if epoch == 0: - wandb.log({"best_test_f1": te_f1}, step=epoch) - elif val_f1 > best_val_f1: - best_val_f1 = val_f1 - wandb.log({"best_test_f1": te_f1}, step=epoch) - if args.save_model: - save_model(model, optimizer, epoch, args, data_config) - - return model - -def get_model(sample_batch, config, args): - n_feats = sample_batch.x.shape[1] if not isinstance(sample_batch, HeteroData) else sample_batch['node'].x.shape[1] - e_dim = (sample_batch.edge_attr.shape[1] - 1) if not isinstance(sample_batch, HeteroData) else (sample_batch['node', 'to', 'node'].edge_attr.shape[1] - 1) - - if args.model == "gin": - model = GINe( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), residual=False, edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, final_dropout=config.final_dropout - ) - elif args.model == "gat": - model = GATe( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), n_heads=round(config.n_heads), - edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, final_dropout=config.final_dropout - ) - elif args.model == "pna": - if not isinstance(sample_batch, HeteroData): - d = degree(sample_batch.edge_index[1], dtype=torch.long) - else: - index = torch.cat((sample_batch['node', 'to', 'node'].edge_index[1], sample_batch['node', 'rev_to', 'node'].edge_index[1]), 0) - d = degree(index, dtype=torch.long) - deg = torch.bincount(d, minlength=1) - model = PNA( - num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, - n_hidden=round(config.n_hidden), edge_updates=args.emlps, edge_dim=e_dim, - dropout=config.dropout, deg=deg, final_dropout=config.final_dropout - ) + +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 + best_val_f1 = 0 + for epoch in range(config.epochs): + total_loss = total_examples = 0 + preds = [] + ground_truths = [] + for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): + optimizer.zero_grad() + #select the seed edges from which the batch was created + inds = tr_inds.detach().cpu() + batch_edge_inds = inds[batch.input_id.detach().cpu()] + batch_edge_ids = tr_loader.data.edge_attr.detach().cpu()[batch_edge_inds, 0] + mask = torch.isin(batch.edge_attr[:, 0].detach().cpu(), batch_edge_ids) + + #remove the unique edge id from the edge features, as it's no longer needed + batch.edge_attr = batch.edge_attr[:, 1:] + + batch.to(device) + out = model(batch.x, batch.edge_index, batch.edge_attr) + pred = out[mask] + ground_truth = batch.y[mask] + preds.append(pred.argmax(dim=-1)) + ground_truths.append(ground_truth) + loss = loss_fn(pred, ground_truth) + + loss.backward() + optimizer.step() + + total_loss += float(loss) * pred.numel() + total_examples += pred.numel() + + pred = torch.cat(preds, dim=0).detach().cpu().numpy() + ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() + f1 = f1_score(ground_truth, pred) + wandb.log({"f1/train": f1}, step=epoch) + logging.info(f'Train F1: {f1:.4f}') + + #evaluate + val_f1 = evaluate_homo(val_loader, val_inds, model, val_data, device, args) + te_f1 = evaluate_homo(te_loader, te_inds, model, te_data, device, args) + + wandb.log({"f1/validation": val_f1}, step=epoch) + wandb.log({"f1/test": te_f1}, step=epoch) + logging.info(f'Validation F1: {val_f1:.4f}') + logging.info(f'Test F1: {te_f1:.4f}') + + if epoch == 0: + wandb.log({"best_test_f1": te_f1}, step=epoch) + elif val_f1 > best_val_f1: + best_val_f1 = val_f1 + wandb.log({"best_test_f1": te_f1}, step=epoch) + if args.save_model: + save_model(model, optimizer, epoch, args, data_config) + + return model + +def train_hetero(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 + best_val_f1 = 0 + for epoch in range(config.epochs): + total_loss = total_examples = 0 + preds = [] + ground_truths = [] + for batch in tqdm.tqdm(tr_loader, disable=not args.tqdm): + optimizer.zero_grad() + #select the seed edges from which the batch was created + inds = tr_inds.detach().cpu() + batch_edge_inds = inds[batch['node', 'to', 'node'].input_id.detach().cpu()] + batch_edge_ids = tr_loader.data['node', 'to', 'node'].edge_attr.detach().cpu()[batch_edge_inds, 0] + mask = torch.isin(batch['node', 'to', 'node'].edge_attr[:, 0].detach().cpu(), batch_edge_ids) + + #remove the unique edge id from the edge features, as it's no longer needed + batch['node', 'to', 'node'].edge_attr = batch['node', 'to', 'node'].edge_attr[:, 1:] + batch['node', 'rev_to', 'node'].edge_attr = batch['node', 'rev_to', 'node'].edge_attr[:, 1:] + + batch.to(device) + out = model(batch.x_dict, batch.edge_index_dict, batch.edge_attr_dict) + out = out[('node', 'to', 'node')] + pred = out[mask] + ground_truth = batch['node', 'to', 'node'].y[mask] + preds.append(pred.argmax(dim=-1)) + ground_truths.append(batch['node', 'to', 'node'].y[mask]) + loss = loss_fn(pred, ground_truth) + + loss.backward() + optimizer.step() + + total_loss += float(loss) * pred.numel() + total_examples += pred.numel() + + pred = torch.cat(preds, dim=0).detach().cpu().numpy() + ground_truth = torch.cat(ground_truths, dim=0).detach().cpu().numpy() + f1 = f1_score(ground_truth, pred) + wandb.log({"f1/train": f1}, step=epoch) + logging.info(f'Train F1: {f1:.4f}') + + #evaluate + val_f1 = evaluate_hetero(val_loader, val_inds, model, val_data, device, args) + te_f1 = evaluate_hetero(te_loader, te_inds, model, te_data, device, args) + + wandb.log({"f1/validation": val_f1}, step=epoch) + wandb.log({"f1/test": te_f1}, step=epoch) + logging.info(f'Validation F1: {val_f1:.4f}') + logging.info(f'Test F1: {te_f1:.4f}') + + if epoch == 0: + wandb.log({"best_test_f1": te_f1}, step=epoch) + elif val_f1 > best_val_f1: + best_val_f1 = val_f1 + wandb.log({"best_test_f1": te_f1}, step=epoch) + if args.save_model: + save_model(model, optimizer, epoch, args, data_config) + + return model + +def get_model(sample_batch, config, args): + n_feats = sample_batch.x.shape[1] if not isinstance(sample_batch, HeteroData) else sample_batch['node'].x.shape[1] + e_dim = (sample_batch.edge_attr.shape[1] - 1) if not isinstance(sample_batch, HeteroData) else (sample_batch['node', 'to', 'node'].edge_attr.shape[1] - 1) + + if args.model == "gin": + model = GINe( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), residual=False, edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, final_dropout=config.final_dropout + ) + elif args.model == "gat": + model = GATe( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), n_heads=round(config.n_heads), + edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, final_dropout=config.final_dropout + ) + elif args.model == "pna": + if not isinstance(sample_batch, HeteroData): + d = degree(sample_batch.edge_index[1], dtype=torch.long) + else: + index = torch.cat((sample_batch['node', 'to', 'node'].edge_index[1], sample_batch['node', 'rev_to', 'node'].edge_index[1]), 0) + d = degree(index, dtype=torch.long) + deg = torch.bincount(d, minlength=1) + model = PNA( + num_features=n_feats, num_gnn_layers=config.n_gnn_layers, n_classes=2, + n_hidden=round(config.n_hidden), edge_updates=args.emlps, edge_dim=e_dim, + dropout=config.dropout, deg=deg, final_dropout=config.final_dropout + ) elif config.model == "rgcn": model = RGCN( num_features=n_feats, edge_dim=e_dim, num_relations=8, num_gnn_layers=round(config.n_gnn_layers), @@ -162,78 +162,78 @@ def get_model(sample_batch, config, args): 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 - -def train_gnn(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, args, data_config): - #set device - device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") - - #define a model config dictionary and wandb logging at the same time - wandb.init( - mode="disabled" if args.testing else "online", - project="your_proj_name", #replace this with your wandb project name if you want to use wandb logging - - config={ - "epochs": args.n_epochs, - "batch_size": args.batch_size, - "model": args.model, - "data": args.data, - "num_neighbors": args.num_neighs, - "lr": extract_param("lr", args), - "n_hidden": extract_param("n_hidden", args), - "n_gnn_layers": extract_param("n_gnn_layers", args), - "loss": "ce", - "w_ce1": extract_param("w_ce1", args), - "w_ce2": extract_param("w_ce2", args), - "dropout": extract_param("dropout", args), - "final_dropout": extract_param("final_dropout", args), - "n_heads": extract_param("n_heads", args) if args.model == 'gat' else None - } - ) - - config = wandb.config - - #set the transform if ego ids should be used - if args.ego: - transform = AddEgoIds() - else: - transform = None - - #add the unique ids to later find the seed edges - add_arange_ids([tr_data, val_data, te_data]) - - tr_loader, val_loader, te_loader = get_loaders(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, transform, args) - - #get the model - sample_batch = next(iter(tr_loader)) - model = get_model(sample_batch, config, args) - - if args.reverse_mp: - model = to_hetero(model, te_data.metadata(), aggr='mean') - - if args.finetune: - model, optimizer = load_model(model, device, args, config, data_config) - else: - model.to(device) - optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) - - sample_batch.to(device) - sample_x = sample_batch.x if not isinstance(sample_batch, HeteroData) else sample_batch.x_dict - sample_edge_index = sample_batch.edge_index if not isinstance(sample_batch, HeteroData) else sample_batch.edge_index_dict - if isinstance(sample_batch, HeteroData): - sample_batch['node', 'to', 'node'].edge_attr = sample_batch['node', 'to', 'node'].edge_attr[:, 1:] - sample_batch['node', 'rev_to', 'node'].edge_attr = sample_batch['node', 'rev_to', 'node'].edge_attr[:, 1:] - else: - sample_batch.edge_attr = sample_batch.edge_attr[:, 1:] - sample_edge_attr = sample_batch.edge_attr if not isinstance(sample_batch, HeteroData) else sample_batch.edge_attr_dict - logging.info(summary(model, sample_x, sample_edge_index, sample_edge_attr)) - - loss_fn = torch.nn.CrossEntropyLoss(weight=torch.FloatTensor([config.w_ce1, config.w_ce2]).to(device)) - - if args.reverse_mp: - model = train_hetero(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) - 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) - + + return model + +def train_gnn(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, args, data_config): + #set device + device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + + #define a model config dictionary and wandb logging at the same time + wandb.init( + mode="disabled" if args.testing else "online", + project="your_proj_name", #replace this with your wandb project name if you want to use wandb logging + + config={ + "epochs": args.n_epochs, + "batch_size": args.batch_size, + "model": args.model, + "data": args.data, + "num_neighbors": args.num_neighs, + "lr": extract_param("lr", args), + "n_hidden": extract_param("n_hidden", args), + "n_gnn_layers": extract_param("n_gnn_layers", args), + "loss": "ce", + "w_ce1": extract_param("w_ce1", args), + "w_ce2": extract_param("w_ce2", args), + "dropout": extract_param("dropout", args), + "final_dropout": extract_param("final_dropout", args), + "n_heads": extract_param("n_heads", args) if args.model == 'gat' else None + } + ) + + config = wandb.config + + #set the transform if ego ids should be used + if args.ego: + transform = AddEgoIds() + else: + transform = None + + #add the unique ids to later find the seed edges + add_arange_ids([tr_data, val_data, te_data]) + + tr_loader, val_loader, te_loader = get_loaders(tr_data, val_data, te_data, tr_inds, val_inds, te_inds, transform, args) + + #get the model + sample_batch = next(iter(tr_loader)) + model = get_model(sample_batch, config, args) + + if args.reverse_mp: + model = to_hetero(model, te_data.metadata(), aggr='mean') + + if args.finetune: + model, optimizer = load_model(model, device, args, config, data_config) + else: + model.to(device) + optimizer = torch.optim.Adam(model.parameters(), lr=config.lr) + + sample_batch.to(device) + sample_x = sample_batch.x if not isinstance(sample_batch, HeteroData) else sample_batch.x_dict + sample_edge_index = sample_batch.edge_index if not isinstance(sample_batch, HeteroData) else sample_batch.edge_index_dict + if isinstance(sample_batch, HeteroData): + sample_batch['node', 'to', 'node'].edge_attr = sample_batch['node', 'to', 'node'].edge_attr[:, 1:] + sample_batch['node', 'rev_to', 'node'].edge_attr = sample_batch['node', 'rev_to', 'node'].edge_attr[:, 1:] + else: + sample_batch.edge_attr = sample_batch.edge_attr[:, 1:] + sample_edge_attr = sample_batch.edge_attr if not isinstance(sample_batch, HeteroData) else sample_batch.edge_attr_dict + logging.info(summary(model, sample_x, sample_edge_index, sample_edge_attr)) + + loss_fn = torch.nn.CrossEntropyLoss(weight=torch.FloatTensor([config.w_ce1, config.w_ce2]).to(device)) + + if args.reverse_mp: + model = train_hetero(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) + 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()