compara-deep-learning/train.py

228 lines
7.9 KiB
Python
Raw Normal View History

2019-08-11 22:56:13 -07:00
import pickle
import numpy as np
import sys
import tensorflow as tf
from model import create_model
2021-04-08 08:41:16 -07:00
2019-08-11 22:56:13 -07:00
def create_branch_length_padding(bl):
2021-04-08 08:41:16 -07:00
maxlen = 0
2019-08-11 22:56:13 -07:00
for x in bl:
2021-04-08 08:41:16 -07:00
if len(x) > maxlen:
maxlen = len(x)
2019-08-11 22:56:13 -07:00
for x in range(len(bl)):
2021-04-08 08:41:16 -07:00
temp = bl[x]
for i in range(len(temp), maxlen):
temp = np.append(temp, [0])
bl[x] = temp
2019-08-11 22:56:13 -07:00
return bl
2021-04-08 08:41:16 -07:00
def train(train_synteny_matrices_global, train_synteny_matrices_local,
train_pfam_matrices, train_branch_length_species,
train_branch_length_homology_species, train_dist_p_s,
train_dist_p_hs, train_distance, train_labels, v,
num_epochs, learning_rate, decay, size_train,
batch_size, model_name):
print("Going to train model {} for:\n Batch Size:{} \
\n Learning Rate:{} \n Decay:{}\n On {} Samples".format(
v, batch_size, learning_rate, decay,
len(train_synteny_matrices_global)
))
graph, saver = create_model()
synmg, synml, pfam, bls, blhs, dps, \
dphs, dis, lr, y = graph.get_collection("input_nodes")
loss, t_op, accuracy, init, summary = graph.get_collection("output_nodes")
2019-08-11 22:56:13 -07:00
with tf.Session(graph=graph) as sess:
writer = tf.summary.FileWriter('./'+model_name+'_v'+str(v), sess.graph)
sess.run(init)
2021-04-08 08:41:16 -07:00
learn = learning_rate
2019-08-11 22:56:13 -07:00
for j in range(num_epochs):
for i in range(size_train//batch_size):
2021-04-08 08:41:16 -07:00
feed_dict = {
synmg: train_synteny_matrices_global
[i*batch_size:(i+1)*batch_size],
synml: train_synteny_matrices_local
[i*batch_size:(i+1)*batch_size],
pfam: train_pfam_matrices
[i*batch_size:(i+1)*batch_size]
.reshape((batch_size, 7, 7, 1)),
bls: train_branch_length_species
[i*batch_size:(i+1)*batch_size],
blhs: train_branch_length_homology_species
[i*batch_size:(i+1)*batch_size],
dps: train_dist_p_s
[i*batch_size:(i+1)*batch_size]
.reshape((batch_size, 1)),
dphs: train_dist_p_hs
[i*batch_size:(i+1)*batch_size]
.reshape((batch_size, 1)),
dis: train_distance
[i*batch_size:(i+1)*batch_size]
.reshape((batch_size, 1)),
lr: learn,
y: train_labels[i*batch_size:(i+1)*batch_size]
2019-08-11 22:56:13 -07:00
}
2021-04-08 08:41:16 -07:00
sess.run(t_op, feed_dict=feed_dict)
test_dict = {
synmg: train_synteny_matrices_global[size_train:],
synml: train_synteny_matrices_local[size_train:],
pfam: train_pfam_matrices[size_train:]
.reshape((
len(train_synteny_matrices_global)-size_train, 7, 7, 1)),
bls: train_branch_length_species[size_train:],
blhs: train_branch_length_homology_species[size_train:],
dps: train_dist_p_s[size_train:]
.reshape((len(train_synteny_matrices_global)-size_train, 1)),
dphs: train_dist_p_hs[size_train:]
.reshape((len(train_synteny_matrices_global)-size_train, 1)),
dis: train_distance[size_train:]
.reshape((len(train_synteny_matrices_global)-size_train, 1)),
y: train_labels[size_train:],
lr: learn
}
accuracy_test, loss_test, summary_write = sess.run(
[accuracy, loss, summary], feed_dict=test_dict)
writer.add_summary(summary_write, i+1)
print("Epoch:{} Test Accuracy:{} Test Loss:{}".format(
j+1, accuracy_test*100, loss_test))
learn *= decay
saver.save(sess, model_name+"_v"+str(v)+"/model.ckpt")
2019-08-11 22:56:13 -07:00
writer.close()
2021-04-08 08:41:16 -07:00
def read_positive(len_p, bls, blhs, dis, dps, dphs, sml, smg, pfam, label):
rowsh = []
with open("dataset", "rb") as file:
rowsh = pickle.load(file)
shi = np.random.permutation(len(rowsh))
rows_shuffled = []
2019-08-11 22:56:13 -07:00
for i in range(len(shi)):
2021-04-08 08:41:16 -07:00
rows_shuffled.append(rowsh[shi[i]])
rowsh = rows_shuffled
spco = {}
spcp = {}
2019-08-11 22:56:13 -07:00
for row in rowsh:
if row["species"] not in spco:
2021-04-08 08:41:16 -07:00
spco[row["species"]] = 0
spcp[row["species"]] = 0
maxcount_o = int((len_p)*0.3/14)
maxcount_p = int((len_p)*0.7/14)
2019-08-11 22:56:13 -07:00
for row in rowsh:
2021-04-08 08:41:16 -07:00
if row["label"] == 2:
2019-08-11 22:56:13 -07:00
continue
2021-04-08 08:41:16 -07:00
if row["label"] == 1 and spco[row["species"]] > maxcount_o:
2019-08-11 22:56:13 -07:00
continue
2021-04-08 08:41:16 -07:00
if row["label"] == 0 and spcp[row["species"]] > maxcount_p:
2019-08-11 22:56:13 -07:00
continue
bls.append(np.array(row["bls"]))
blhs.append(np.array(row["blhs"]))
dis.append(row["dis"])
dps.append(row["dps"])
dphs.append(row["dphs"])
sml.append(row["local_alignment_matrix"])
smg.append(row["global_alignment_matrix"])
pfam.append(row["pfam_matrix"])
label.append(row["label"])
2021-04-08 08:41:16 -07:00
if row["label"] == 1:
spco[row["species"]] += 1
if row["label"] == 0:
spcp[row["species"]] += 1
def read_negative(len_n, bls, blhs, dis, dps, dphs, sml, smg, pfam, label):
rows = []
with open("dataset", "rb") as file:
rows = pickle.load(file)
rows = [row for row in rows if row["label"] == 2]
shi = np.random.permutation(len(rows))
rows_shuffled = []
2019-08-11 22:56:13 -07:00
for i in range(len(shi)):
rows_shuffled.append(rows[shi[i]])
2021-04-08 08:41:16 -07:00
rows = rows_shuffled
rows = rows[:len_n]
2019-08-11 22:56:13 -07:00
for row in rows:
bls.append(np.array(row["bls"]))
blhs.append(np.array(row["blhs"]))
dis.append(row["dis"])
dps.append(row["dps"])
dphs.append(row["dphs"])
sml.append(row["local_alignment_matrix"])
smg.append(row["global_alignment_matrix"])
label.append(row["label"])
pfam.append(row["pfam_matrix"])
2021-04-08 08:41:16 -07:00
def train_models(model_name, start, end, num_epochs, learn_rate,
decay, size_train, batch_size):
k = 1
for i in range(start//10, end//10+1):
bls = []
blhs = []
dis = []
dps = []
dphs = []
sml = []
smg = []
pfam = []
label = []
portion = float(i/10)
len_n = int(size_train*portion)
len_p = int(size_train*(1-portion))
read_positive(len_p, bls, blhs, dis, dps, dphs, sml, smg, pfam, label)
read_negative(len_n, bls, blhs, dis, dps, dphs, sml, smg, pfam, label)
bls = create_branch_length_padding(bls)
blhs = create_branch_length_padding(blhs)
bls = np.array(bls)
2019-08-11 22:56:13 -07:00
print(bls.shape)
2021-04-08 08:41:16 -07:00
blhs = np.array(blhs)
2019-08-11 22:56:13 -07:00
print(blhs.shape)
2021-04-08 08:41:16 -07:00
dis = np.array(dis)
2019-08-11 22:56:13 -07:00
print(dis.shape)
2021-04-08 08:41:16 -07:00
dps = np.array(dps)
2019-08-11 22:56:13 -07:00
print(dps.shape)
2021-04-08 08:41:16 -07:00
dphs = np.array(dphs)
2019-08-11 22:56:13 -07:00
print(dphs.shape)
2021-04-08 08:41:16 -07:00
sml = np.array(sml)
2019-08-11 22:56:13 -07:00
print(sml.shape)
2021-04-08 08:41:16 -07:00
smg = np.array(smg)
2019-08-11 22:56:13 -07:00
print(smg.shape)
2021-04-08 08:41:16 -07:00
pfam = np.array(pfam)
2019-08-11 22:56:13 -07:00
print(pfam.shape)
2021-04-08 08:41:16 -07:00
label = np.array(label)
2019-08-11 22:56:13 -07:00
print(label.shape)
2021-04-08 08:41:16 -07:00
shi = np.random.permutation(len(label))
labels = label[shi]
bls = bls[shi]
blhs = blhs[shi]
dis = dis[shi]
dps = dps[shi]
dphs = dphs[shi]
sml = sml[shi]
smg = smg[shi]
pfam = pfam[shi]
train(smg, sml, pfam, bls, blhs, dps, dphs, dis, labels, k, num_epochs,
learn_rate, decay, int(0.9*size_train), batch_size, model_name)
k += 1
2019-08-11 22:56:13 -07:00
def main():
2021-04-08 08:41:16 -07:00
arg = sys.argv
model_name = arg[-8]
start_p = int(arg[-7])
end_p = int(arg[-6])
num_epochs = int(arg[-5])
learn_rate = float(arg[-4])
decay = float(arg[-3])
size_train = float(arg[-2])
batch_size = int(arg[-1])
train_models(model_name, start_p, end_p, num_epochs,
learn_rate, decay, size_train, batch_size)
2019-08-11 22:56:13 -07:00
2021-04-08 08:41:16 -07:00
if __name__ == "__main__":
main()