mirror of
https://github.com/Priyatham-sai-chand/compara-deep-learning.git
synced 2026-10-05 08:11:34 -07:00
227 lines
7.9 KiB
Python
227 lines
7.9 KiB
Python
import pickle
|
|
import numpy as np
|
|
import sys
|
|
import tensorflow as tf
|
|
from model import create_model
|
|
|
|
|
|
def create_branch_length_padding(bl):
|
|
maxlen = 0
|
|
for x in bl:
|
|
if len(x) > maxlen:
|
|
maxlen = len(x)
|
|
|
|
for x in range(len(bl)):
|
|
temp = bl[x]
|
|
for i in range(len(temp), maxlen):
|
|
temp = np.append(temp, [0])
|
|
bl[x] = temp
|
|
return bl
|
|
|
|
|
|
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")
|
|
with tf.Session(graph=graph) as sess:
|
|
writer = tf.summary.FileWriter('./'+model_name+'_v'+str(v), sess.graph)
|
|
sess.run(init)
|
|
learn = learning_rate
|
|
for j in range(num_epochs):
|
|
for i in range(size_train//batch_size):
|
|
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]
|
|
}
|
|
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")
|
|
writer.close()
|
|
|
|
|
|
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 = []
|
|
for i in range(len(shi)):
|
|
rows_shuffled.append(rowsh[shi[i]])
|
|
rowsh = rows_shuffled
|
|
spco = {}
|
|
spcp = {}
|
|
for row in rowsh:
|
|
if row["species"] not in spco:
|
|
spco[row["species"]] = 0
|
|
spcp[row["species"]] = 0
|
|
|
|
maxcount_o = int((len_p)*0.3/14)
|
|
maxcount_p = int((len_p)*0.7/14)
|
|
for row in rowsh:
|
|
if row["label"] == 2:
|
|
continue
|
|
if row["label"] == 1 and spco[row["species"]] > maxcount_o:
|
|
continue
|
|
if row["label"] == 0 and spcp[row["species"]] > maxcount_p:
|
|
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"])
|
|
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 = []
|
|
for i in range(len(shi)):
|
|
rows_shuffled.append(rows[shi[i]])
|
|
rows = rows_shuffled
|
|
rows = rows[:len_n]
|
|
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"])
|
|
|
|
|
|
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)
|
|
print(bls.shape)
|
|
blhs = np.array(blhs)
|
|
print(blhs.shape)
|
|
dis = np.array(dis)
|
|
print(dis.shape)
|
|
dps = np.array(dps)
|
|
print(dps.shape)
|
|
dphs = np.array(dphs)
|
|
print(dphs.shape)
|
|
sml = np.array(sml)
|
|
print(sml.shape)
|
|
smg = np.array(smg)
|
|
print(smg.shape)
|
|
pfam = np.array(pfam)
|
|
print(pfam.shape)
|
|
label = np.array(label)
|
|
print(label.shape)
|
|
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
|
|
|
|
|
|
def main():
|
|
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)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|