-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgraph_influence.py
More file actions
95 lines (86 loc) · 3.87 KB
/
Copy pathgraph_influence.py
File metadata and controls
95 lines (86 loc) · 3.87 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
import random
import argparse
import torch
import logging
import os
from loader import load_data, sample_subgraph
from models.model import load_model
from influence import Influence
from shadow_influence import ShadowInfluence
from utils import init_default_config, save_json
import torch.nn.functional as F
def get_args() -> list:
parser = argparse.ArgumentParser()
parser.add_argument('--dataset', type=str, default='Cora',
help='Dataset (Cora, Flickr, PubMed, CiteSeer)')
parser.add_argument('--model', type=str, default='GCN',
help='Model (GCN, GAT, GIN, ARMA)')
parser.add_argument('--sampling', type=str, default=False,
help='Sampling method (shadowkhop, graphsaint, random)')
parser.add_argument('--batch_size', type=int, default=1024,
help='Batch size for sampling')
parser.add_argument('--hidden_layers', type=int, default=256,
help='Number of hidden layers')
parser.add_argument('--num_layers', type=int, default=2,
help='Number of layers for GCN')
parser.add_argument('--heads', type=int, default=8,
help='Number of heads for GAT')
parser.add_argument('--seed', type=int, default=123,
help='Random seed')
parser.add_argument('--device', type=str, default='cuda',
help='Device to train')
parser.add_argument('--node_ids', nargs='+', type=int, default=[False],
help='Testing node ids')
parser.add_argument('--recursion_depth', type=int, default=1,
help='Recursion depth for s_test calculation')
parser.add_argument('--r_averaging', type=int, default=1,
help='R averaging')
parser.add_argument('--debug', dest="loglevel", action='store_const',
default=logging.INFO, const=logging.DEBUG,
help='Display additional debug info')
parser.add_argument('--experiment_name', type=str, default=False,
help='Experiment name to save the results')
args = parser.parse_args()
return args
if __name__ == '__main__':
args = get_args()
init_default_config(args)
logging.info(f"Dataset: {args.dataset}")
dataset = load_data(args.dataset)
if args.sampling:
train_l, test_l, _ = sample_subgraph(dataset, args.sampling, args.batch_size)
logging.info("Training Sub-graphs:")
logging.info(dataset[0])
for s in train_l:
logging.info(s)
device = torch.device(args.device if torch.cuda.is_available() and
args.device != 'cpu'else 'cpu')
logging.info(f"Using: {device}")
model = load_model(args.model,
in_channels=dataset.num_features,
hidden_channels=args.hidden_layers,
heads=args.heads,
num_layers=args.num_layers,
out_channels=dataset.num_classes,
batch=args.sampling and True)
model = model.to(device)
logging.info(model)
if args.sampling:
for idx,graph in enumerate(train_l):
logging.info(f"Subgraph index {idx}")
influence = ShadowInfluence(model, graph, device, args.recursion_depth, args.r_averaging)
result = influence.calculate()
print(result)
else:
influence = Influence(model, dataset[0], device, args.recursion_depth, args.r_averaging)
for node_id in args.node_ids:
result = influence.calculate(node_id)
if args.experiment_name:
save_json(args.experiment_name, {
'model': args.model,
'dataset': args.dataset,
'seed': args.seed,
'node_id': node_id,
'influence': result
})
print(result)