-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
84 lines (74 loc) · 3.68 KB
/
Copy pathtrain.py
File metadata and controls
84 lines (74 loc) · 3.68 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
import os
import random
import torch
from dgl.dataloading import GraphDataLoader
from model.model import STGNN_AutoEncoder
from model.train_graph import batch_level_train
from model.train_entity import entity_level_train
from utils.dataloader import load_data, load_metadata
from utils.utils import set_random_seed, create_optimizer
from utils.config import build_args
# Function to shuffle and create dataloaders for training
def extract_dataloaders(entries):
random.shuffle(entries)
train_idx = torch.arange(len(entries))
train_loader = GraphDataLoader(entries, batch_size=1, sampler=train_idx)
return train_loader
# Main function for training the model
def main(main_args):
# Set device based on availability of GPU
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
dataset_name = main_args.dataset
# Graph level model
if dataset_name in ['wget', 'streamspot', 'SC2', 'Unicorn-Cadets', 'wget-long', 'clearscope-e3']:
# Set specific parameters based on dataset
if dataset_name in ['wget', 'streamspot', 'Unicorn-Cadets', 'wget-long', 'clearscope-e3']:
if (dataset_name == 'wget-long'):
main_args.max_epoch = 1
else:
main_args.max_epoch = 6
out_dim = 64
gnn_layer = 5
nsnapshot = main_args.snapshot
elif dataset_name in ['SC2']:
main_args.max_epoch = 1
out_dim = 64
gnn_layer = 3
nsnapshot = main_args.snapshot
# Load dataset and set parameters
dataset = load_data(dataset_name)
n_node_feat = dataset['n_feat']
n_edge_feat = dataset['e_feat']
train_index = dataset['train_index']
main_args.n_dim = n_node_feat
main_args.e_dim = n_edge_feat
main_args.optimizer = "adamw"
set_random_seed(0)
# Initialize model and optimizer, perform training
model = STGNN_AutoEncoder(main_args.n_dim, main_args.e_dim, out_dim, out_dim, gnn_layer, 4, device, nsnapshot, 'prelu', 0.1, args.negative_slope, True, 'BatchNorm', main_args.pooling, alpha_l=args.alpha_l, use_all_hidden=True).to(device)
optimizer = create_optimizer(main_args.optimizer, model, main_args.lr, main_args.weight_decay)
train_loader = extract_dataloaders(train_index)
model = batch_level_train(model, train_loader, optimizer, main_args.max_epoch, device, main_args.n_dim, main_args.e_dim, dataset_name, validation= False)
else:
# Set parameters for entity level model
main_args.max_epoch = 50
info = load_metadata(dataset_name)
out_dim = 128
if dataset_name == 'cadets-e3':
gnn_layer = 5
main_args.optimizer = "adamw"
else:
gnn_layer = 4
nsnapshot = main_args.snapshot
n_node_feat = info['node_feature_dim']
n_edge_feat = info['edge_feature_dim']
main_args.n_dim = n_node_feat
main_args.e_dim = n_edge_feat
set_random_seed(0)
# Initialize model and optimizer, perform training
model = STGNN_AutoEncoder(main_args.n_dim, main_args.e_dim, out_dim, out_dim, gnn_layer, 4, device, nsnapshot, 'prelu', 0.1, args.negative_slope, True, 'BatchNorm', main_args.pooling, alpha_l=args.alpha_l, use_all_hidden=False).to(device)
optimizer = create_optimizer(main_args.optimizer, model, main_args.lr, main_args.weight_decay)
model = entity_level_train(model, nsnapshot, optimizer, main_args.max_epoch, device, dataset_name)
if __name__ == '__main__':
args = build_args() # Parse command line arguments
main(args) # Call the main function with parsed arguments