-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
79 lines (55 loc) · 2.26 KB
/
Copy pathtrain.py
File metadata and controls
79 lines (55 loc) · 2.26 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
import os, joblib
import os
os.environ['KMP_DUPLICATE_LIB_OK']='True'
import pandas as pd
from datetime import datetime, timezone, timedelta
# from lightgbm.callback import record_evaluation
import xgboost as xgb
from models.ML_model import get_ml_model
from modules.utils import load_yaml, load_pkl, make_directory, save_pkl, save_yaml
# CONFIG
PROJECT_DIR = os.path.dirname(os.path.abspath(__file__))
TRAIN_CONFIG_PATH = os.path.join(PROJECT_DIR, 'config/train_config.yaml')
config = load_yaml(TRAIN_CONFIG_PATH)
# DATA
DATA_DIR = config['DIRECTORY']['data']
# SEED
RANDOM_SEED = config['SEED']['random_seed']
# MODEL
MODEL_STR = config['MODEL']
PARAMETER = config['PARAMETER']
#LABEL_ENCODE
LABEL_ENCODING = config['LABEL_ENCODING']
# TRAIN
EARLY_STOPPING_ROUND = config['TRAIN']['early_stopping_round']
# time offset set
KST = timezone(timedelta(hours=9))
TRAIN_TIMESTAMP = datetime.now(tz=KST).strftime("%Y%m%d_%H%M%S")
TRAIN_SERIAL = MODEL_STR + '_' + TRAIN_TIMESTAMP
# PERFORMANCE RECORD
PERFORMANCE_RECORD_DIR = os.path.join(PROJECT_DIR, 'results', 'train', TRAIN_SERIAL)
if __name__ == '__main__':
make_directory(PERFORMANCE_RECORD_DIR)
save_yaml(os.path.join(PERFORMANCE_RECORD_DIR,'train_config.yaml'),config)
train_df = pd.read_csv(os.path.join(DATA_DIR, 'train.csv'))
valid_df = pd.read_csv(os.path.join(DATA_DIR, 'valid.csv'))
train_X, train_y = train_df.loc[:,train_df.columns!='leaktype'], train_df['leaktype']
valid_X, valid_y = valid_df.loc[:,train_df.columns!='leaktype'], valid_df['leaktype']
train_y = train_y.replace(LABEL_ENCODING)
valid_y = valid_y.replace(LABEL_ENCODING)
early_stop = xgb.callback.EarlyStopping(rounds=EARLY_STOPPING_ROUND,
metric_name='CustomErr',
data_name='Train')
#MODEL
model = get_ml_model(MODEL_STR, PARAMETER)
history = dict()
model.fit(
train_X,
train_y,
early_stopping_rounds=EARLY_STOPPING_ROUND,
eval_set = [(train_X, train_y),(valid_X, valid_y)],
verbose=1
)
# SAVE MODEL
joblib.dump(model, os.path.join(PERFORMANCE_RECORD_DIR,'model.pkl'))
save_pkl(os.path.join(PERFORMANCE_RECORD_DIR,'loss_history.pkl'),history)