-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathmain.py
More file actions
54 lines (41 loc) · 1.65 KB
/
Copy pathmain.py
File metadata and controls
54 lines (41 loc) · 1.65 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
import warnings
from args import Args
from utils.train import train
from utils.util import *
from utils.data_loader import *
from model.acgcn_mmp import ACGCN_MMP
from model.acgcn_sub import ACGCN_SUB
from sklearn.model_selection import train_test_split
warnings.filterwarnings("ignore")
def main(args):
device = args['DEVICE']
random_seed = args['RANDOM_SEED']
### Make dataset ###
data = pd.read_csv('./data/' + args['TARGET_NAME'] + '_mmps.csv')
data['label'] = data['label'].astype(int)
label = data['label']
counter = label.value_counts()
tot = counter.sum()
class_weight = [tot / (2 * i) for i in counter]
train_data, test_data = train_test_split(data, test_size=0.2, random_state=random_seed, stratify=label)
if args['MODEL'] == 'acgcn-mmp':
model = ACGCN_MMP(args).to(device)
train_loader = ACGCN_MMP_Dataset(args, train_data, True)
test_loader = ACGCN_MMP_Dataset(args, test_data, False)
elif args['MODEL'] == 'acgcn-sub':
model = ACGCN_SUB(args).to(device)
train_loader = ACGCN_SUB_Dataset(args, train_data, True)
test_loader = ACGCN_SUB_Dataset(args, test_data, False)
y_actual = get_actual_label(test_loader)
y_proba = train(args, model, train_loader, test_loader, class_weight)
print_metrics(y_proba, y_actual)
if __name__ == '__main__':
args = Args().params
random_seed = args['RANDOM_SEED']
torch.manual_seed(random_seed)
torch.cuda.manual_seed(random_seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
np.random.seed(random_seed)
random.seed(random_seed)
main(args)