-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathrun.py
More file actions
119 lines (105 loc) · 4.56 KB
/
Copy pathrun.py
File metadata and controls
119 lines (105 loc) · 4.56 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
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
import torch
import pickle
import argparse
import json
import os
import yaml
import wandb
import random
import datetime
import torchmetrics
import csv
import numpy as np
from tqdm import tqdm
from utils import *
from torch.utils.data import Dataset
from dataset import SelectionDataset
from model import Selection_model, Naive_classifier
from trainer import trainer
if __name__ == "__main__":
parser = argparse.ArgumentParser(description='Load config')
parser.add_argument('--config_name', type=str, default=None)
parser.add_argument('--seed', type=int, default=None)
parser.add_argument('--loss', type=str, default=None)
parser.add_argument('--load', type=str, default=None)
parser.add_argument('--test_file', type=str, default=None)
parser.add_argument('--exp_name', type=str, default=None)
parser.add_argument('--gpu_id', type=int, default=None)
args = parser.parse_args()
if args.load is not None:
with open(f"train_logs/{args.load}/config.json", 'r', encoding='utf-8') as config_file:
config = yaml.load(config_file.read(), Loader=yaml.FullLoader)
config['load_path'] = args.load
else:
with open(args.config_name, 'r', encoding='utf-8') as config_file:
config = yaml.load(config_file.read(), Loader=yaml.FullLoader)
# params
name = config['name']
seed = args.seed if args.seed is not None else config['seed']
config['train_params']['loss'] = args.loss if args.loss is not None else config['train_params']['loss']
seed_everything(seed)
logger_name = config['logger']
load_path = config['load_path']
config['model_params']['problem_type'] = config['problem_type']
config['model_params']['output_dim'] = config['train_params']['num_classes']
# Initialize logger
ts = datetime.datetime.utcnow() + datetime.timedelta(hours=+8)
ts_name = f'-ts{ts.month}-{ts.day}-{ts.hour}-{ts.minute}-{ts.second}'
log_config = config.copy()
param_config = log_config['train_params'].copy()
log_config.pop('train_params')
model_params_config = log_config['model_params'].copy()
log_config.pop('model_params')
log_config.update(param_config)
log_config.update(model_params_config)
logger = {}
if(logger_name == 'wandb'):
logger['wandb'] = wandb.init(project="selection",
name=name + ts_name,
config=log_config)
else:
logger['wandb'] = None
if args.test_file is None:
if load_path is not None:
log_dir = f'train_logs/{load_path}'
else:
log_dir = f'train_logs/{args.config_name}_{args.loss}_{args.seed}'
os.mkdir(log_dir)
logger['file'] = csv_logger(log_dir)
if not os.path.exists(log_dir + '/config.json'):
with open(log_dir + '/config.json', 'w') as f:
json.dump(config, f)
else:
config['train_params']['num_epochs'] = 0 # only test
log_dir = 'results'
if not os.path.exists(log_dir):
os.mkdir(log_dir)
logger['file'] = csv_logger(log_dir, args.exp_name)
print(config)
# Prepare datasets
name = args.test_file if args.test_file is not None else config['name']
train_set, train_label, test_set, test_label = prepare_dataset(config['problem_type'], name=name)
train_dataset = SelectionDataset(train_set, train_label, manual_feature=config['train_params']['manual_feature'], data_aug=config['train_params']['data_aug'])
test_dataset = SelectionDataset(test_set, test_label, manual_feature=config['train_params']['manual_feature'])
config['model_params']['ns_feature'] = config['train_params']['ns_feature']
if config['train_params']['ns_feature']:
representative_set = representative(train_set, train_label)
model = Selection_model(**config['model_params'])
encoder_representative = model.encoder
else:
representative_set = None
encoder_representative = None
# Initialize models
if config['train_params']['manual_feature']:
model = Naive_classifier(**config['model_params'])
else:
model = Selection_model(**config['model_params'])
# Initialize trainer
cuda_device_num = config['cuda_device_num'] if args.gpu_id is None else args.gpu_id
trainer = trainer(model=model,
logger=logger,
cuda_device_num=cuda_device_num,
encoder_representative=encoder_representative,
train_params=config['train_params'])
# Training
trainer.run(train_dataset, test_dataset, representative_set, log_dir, load_path)