Download scripts/train.py from OneScience-Group/NABP-LSTM-Att: direct link, hf CLI and curl.
- Browser
- Download file 4.17 kB
-
https://huggingface.co/OneScience-Group/NABP-LSTM-Att/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/NABP-LSTM-Att/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/NABP-LSTM-Att/resolve/main/scripts/train.py
4.17 kB
| 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) | |