File size: 4,167 Bytes
bf928ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
120
121
122
123
124
125
126
import os
import pickle
from _bootstrap import DATA_DIR, OUTPUT_DIR
from model import get_model
from tf_keras.callbacks import Callback
from datetime import datetime
from sklearn.metrics import roc_auc_score,average_precision_score

cdr_kmer = 3
ag_kmer = 1
epochs=100
batch_size=64
data_path = DATA_DIR / 'features' / ('cdr_kmer' + str(cdr_kmer) + '_ag_kmer' + str(ag_kmer))
directory_path = OUTPUT_DIR / 'checkpoints' / ('cdr_kmer' + str(cdr_kmer) + '_ag_kmer' + str(ag_kmer))


def create_directory_if_not_exists(directory_path):
    if not os.path.exists(directory_path):
        os.makedirs(directory_path)
        print(f"Directory '{directory_path}' created.")
    else:
        print(f"Directory '{directory_path}' already exists.")

class roc_callback(Callback):
    def __init__(self, val_data):
        self.cdr_ids = val_data[0]
        self.cdr_number_ids = val_data[1]
        self.ag_ids = val_data[2]
        self.labels = val_data[3]

    def on_train_begin(self, logs={}):
        return

    def on_train_end(self, logs={}):
        return

    def on_epoch_begin(self, epoch, logs={}):
        return

    def on_epoch_end(self, epoch, logs={}):
        labels_pred = self.model.predict([self.cdr_ids, self.cdr_number_ids, self.ag_ids])
        auc_val = roc_auc_score(self.labels, labels_pred)
        aupr_val = average_precision_score(self.labels, labels_pred)

        create_directory_if_not_exists(directory_path)
        self.model.save_weights(str(directory_path / ("Model%d.weights.h5" % epoch)))

        print('\r auc_val: %s ' %str(round(auc_val, 4)), end=100 * ' ' + '\n')
        print('\r aupr_val: %s ' % str(round(aupr_val, 4)), end=100 * ' ' + '\n')
        return

    def on_batch_begin(self, batch, logs={}):
        return

    def on_batch_end(self, batch, logs={}):
        return

t1 = datetime.now().strftime('%Y-%m-%d-%H:%M:%S')


with (data_path / 'cdr_features_tr.pickle').open('rb') as binary_reader:
    cdr_features_tr = pickle.load(binary_reader)
with (data_path / 'ag_features_tr.pickle').open('rb') as binary_reader:
    ag_features_tr = pickle.load(binary_reader)

with (data_path / 'cdr_features_val.pickle').open('rb') as binary_reader:
    cdr_features_val = pickle.load(binary_reader)
with (data_path / 'ag_features_val.pickle').open('rb') as binary_reader:
    ag_features_val = pickle.load(binary_reader)

# Training data
dtrain_cdr_ids = []
dtrain_cdr_number_ids = []
dtrain_ag_ids = []
dtrain_labels = []
dtrain_labels_pos = 0
dtrain_labels_neg = 0

for feature in cdr_features_tr:
    dtrain_cdr_ids.append(feature.input_ids)
    dtrain_cdr_number_ids.append(feature.cdr_number_ids)
    dtrain_labels.append(feature.label_id)
    dtrain_labels_pos = dtrain_labels_pos + feature.label_id
dtrain_labels_neg = len(dtrain_labels) - dtrain_labels_pos

for feature in ag_features_tr:
    dtrain_ag_ids.append(feature.input_ids)

########################################################
# validation data
dval_cdr_ids = []
dval_cdr_number_ids = []
dval_ag_ids = []
dval_labels = []
dval_labels_pos = 0
dval_labels_neg = 0

for feature in cdr_features_val:
    dval_cdr_ids.append(feature.input_ids)
    dval_cdr_number_ids.append(feature.cdr_number_ids)
    dval_labels.append(feature.label_id)
    dval_labels_pos = dval_labels_pos + feature.label_id
dval_labels_neg = len(dval_labels) - dval_labels_pos

for feature in ag_features_val:
    dval_ag_ids.append(feature.input_ids)

#################################
# get the model
model=None
model=get_model()
model.summary()

print ('Training the model')

back = roc_callback(val_data=[dval_cdr_ids, dval_cdr_number_ids, dval_ag_ids, dval_labels])

history=model.fit([dtrain_cdr_ids, dtrain_cdr_number_ids, dtrain_ag_ids], dtrain_labels,
                  validation_data=([dval_cdr_ids, dval_cdr_number_ids, dval_ag_ids], dval_labels),
                  epochs=epochs,
                  batch_size=batch_size,
                  callbacks=[back])

t2 = datetime.now().strftime('%Y-%m-%d-%H:%M:%S')
print("开始时间:"+t1+"结束时间:"+t2)