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)
|