-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathtraining_5folds.py
More file actions
82 lines (65 loc) · 2.5 KB
/
Copy pathtraining_5folds.py
File metadata and controls
82 lines (65 loc) · 2.5 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
import sys, os
import torch
import torch.nn as nn
from torch_geometric.data import DataLoader
from gnn import GNNNet
from utils import *
from emetrics import *
from data_process import create_dataset_for_5folds
from gat import GATNet
from gnn import GNNNet
from gin import GINConvNet
from gat_gcn import GAT_GCN
from gcnnII import GCNIIdense_model
datasets = [['davis', 'kiba'][int(sys.argv[1])]]
cuda_name = ['cuda:0', 'cuda:1', 'cuda:2', 'cuda:3'][int(sys.argv[2])]
modeling = GNNNet
print('cuda_name:', cuda_name)
fold = [0, 1, 2, 3, 4][int(sys.argv[3])]
cross_validation_flag = True
# print(int(sys.argv[3]))
TRAIN_BATCH_SIZE = 64
TEST_BATCH_SIZE = 64
LR = 0.001
NUM_EPOCHS = 2000
print('Learning rate: ', LR)
print('Epochs: ', NUM_EPOCHS)
models_dir = 'models'
results_dir = 'results'
if not os.path.exists(models_dir):
os.makedirs(models_dir)
if not os.path.exists(results_dir):
os.makedirs(results_dir)
# Main program: iterate over different datasets
result_str = ''
USE_CUDA = torch.cuda.is_available()
device = torch.device(cuda_name if USE_CUDA else 'cpu')
model = modeling()
model.to(device)
model_st = modeling.__name__
# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
loss_fn = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=LR)
for dataset in datasets:
train_data, valid_data = create_dataset_for_5folds(dataset, fold)
train_loader = torch.utils.data.DataLoader(train_data, batch_size=TRAIN_BATCH_SIZE, shuffle=True,
collate_fn=collate)
valid_loader = torch.utils.data.DataLoader(valid_data, batch_size=TEST_BATCH_SIZE, shuffle=False,
collate_fn=collate)
best_mse = 1000
best_test_mse = 1000
best_epoch = -1
model_file_name = 'models/model_' + model_st + '_' + dataset + '_' + str(fold) + '.model'
for epoch in range(NUM_EPOCHS):
train(model, device, train_loader, optimizer, epoch + 1)
print('predicting for valid data')
G, P = predicting(model, device, valid_loader)
val = get_mse(G, P)
print('valid result:', val, best_mse)
if val < best_mse:
best_mse = val
best_epoch = epoch + 1
torch.save(model.state_dict(), model_file_name)
print('rmse improved at epoch ', best_epoch, '; best_test_mse', best_mse, model_st, dataset, fold)
else:
print('No improvement since epoch ', best_epoch, '; best_test_mse', best_mse, model_st, dataset, fold)