mirror of
https://github.com/Priyatham-sai-chand/compara-deep-learning.git
synced 2026-10-05 08:11:34 -07:00
Prediction FIx
This commit is contained in:
parent
787557963a
commit
b24f3fcc0f
3 changed files with 233 additions and 64 deletions
BIN
__pycache__/select.cpython-36.pyc
Normal file
BIN
__pycache__/select.cpython-36.pyc
Normal file
Binary file not shown.
169
prediction_thread_multiple.py
Normal file
169
prediction_thread_multiple.py
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
import json
|
||||
import gc
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
import pickle
|
||||
import tensorflow as tf
|
||||
import sys
|
||||
import time
|
||||
from test_prepare_functions import create_data_homology_ls,read_gene_sequences,create_tree_data,read_data,update_rest
|
||||
from process_data import create_map_list
|
||||
from threading import Thread,Lock
|
||||
import traceback
|
||||
from threads import Procerssrunner
|
||||
|
||||
def main():
|
||||
arg=sys.argv
|
||||
fname=arg[-6]
|
||||
model_name=arg[-5]
|
||||
n_of_t=int(arg[-4])
|
||||
st=int(arg[-3])
|
||||
end=int(arg[-2])
|
||||
name=arg[-1]
|
||||
|
||||
df=pd.read_csv(fname,sep="\t",header=None)
|
||||
label_dict=dict(ortholog_one2one="Orthologs",
|
||||
other_paralog="Paralogs",
|
||||
ortholog_one2many="Orthologs",
|
||||
ortholog_many2many="Orthologs",
|
||||
within_species_paralog="Paralogs",
|
||||
gene_split="Gene Split",
|
||||
non_homolog="non_homolog")
|
||||
label_dict_2=dict(ortholog_one2one=1,other_paralog=0,non_homolog=2,ortholog_one2many=1,ortholog_many2many=1,within_species_paralog=0,gene_split=4)
|
||||
label_2=[]
|
||||
label_1=[]
|
||||
for _,row in df.iterrows():
|
||||
label_1.append(label_dict_2[row[7]])
|
||||
label_2.append(label_dict[row[7]])
|
||||
df=df.assign(label=label_1)
|
||||
df[7]=label_2
|
||||
try:
|
||||
if end<len(df):
|
||||
if st<end:
|
||||
df=df.loc[df.index.values[st:end]]
|
||||
else:
|
||||
raise ValueError()
|
||||
except:
|
||||
print("Making Predictions for the complete dataframe:)")
|
||||
print(len(df))
|
||||
|
||||
data={}
|
||||
with open("genome_maps","rb") as file:
|
||||
data=pickle.load(file)
|
||||
cmap=data["cmap"]
|
||||
cimap=data["cimap"]
|
||||
ld=data["ld"]
|
||||
ldg=data["ldg"]
|
||||
a=data["a"]
|
||||
d=data["d"]
|
||||
lsy=create_data_homology_ls(df,3,a,d,ld,ldg,cmap,cimap)
|
||||
print(len(lsy))
|
||||
data={}
|
||||
a=[]
|
||||
d=[]
|
||||
ldg=[]
|
||||
ld=[]
|
||||
cmap=[]
|
||||
cimap=[]
|
||||
gc.collect()
|
||||
|
||||
gene_sequences=read_gene_sequences(df,lsy,"geneseq","gene_sequences",name)
|
||||
print("Gene Sequences Loaded")
|
||||
|
||||
print("Going to update not found sequences:")
|
||||
gene_sequences=update_rest(gene_sequences,name)
|
||||
|
||||
gc.collect()
|
||||
part=len(df)//n_of_t
|
||||
"""
|
||||
lsy={}
|
||||
gene_sequences={}
|
||||
part=1"""
|
||||
return lsy,gene_sequences,df,model_name,fname,n_of_t,part,name
|
||||
|
||||
|
||||
if __name__=='__main__':
|
||||
n=3
|
||||
lsy,gene_sequences,df,model_name,fname,n_of_t,part,name=main()
|
||||
pr=Procerssrunner()
|
||||
pr.start_processes(n_of_t,df,gene_sequences,lsy,part,n,name)
|
||||
smg,sml,indexes=read_data(n_of_t,name)
|
||||
sml=np.array(sml)
|
||||
smg=np.array(smg)
|
||||
indexes=np.array(indexes)
|
||||
print(indexes.shape)
|
||||
df_temp=df.loc[indexes]
|
||||
bls,blhs,dis,dps,dphs=create_tree_data("species_tree.tree",df_temp)
|
||||
|
||||
assert(len(bls)==len(indexes))
|
||||
assert(len(blhs)==len(smg))
|
||||
assert(len(df_temp)==len(dps))
|
||||
|
||||
print("Data Prepared.")
|
||||
index_dict=create_map_list(indexes)
|
||||
preds=np.zeros((len(smg),3))
|
||||
w=[0.86,0.8,0.06,0.0]
|
||||
for i in range(1,4):
|
||||
try:
|
||||
model=tf.train.import_meta_graph(model_name+'_v'+str(i)+'/model.ckpt.meta')
|
||||
except:
|
||||
print("Something wrong with the model.")
|
||||
continue
|
||||
with tf.Session() as sess:
|
||||
try:
|
||||
model.restore(sess,model_name+'_v'+str(i)+"/model.ckpt")
|
||||
graph = tf.get_default_graph()
|
||||
synmgt,synmlt,blst,blhst,dpst,dphst,dist,lrt,yt=graph.get_collection("input_nodes")
|
||||
predictions=graph.get_tensor_by_name("Predictions/BiasAdd:0")
|
||||
print("Model Loaded Successfully :)")
|
||||
except:
|
||||
print(":(")
|
||||
sys.exit()
|
||||
|
||||
fd={synmgt:smg,
|
||||
synmlt:sml,
|
||||
blst:bls,
|
||||
blhst:blhs,
|
||||
dpst:dps.reshape((len(blhs),1)),
|
||||
dist:dis.reshape((len(blhs),1)),
|
||||
dphst:dphs.reshape((len(blhs),1))}
|
||||
preds_t_1=sess.run([predictions],feed_dict=fd)
|
||||
preds_t_1=np.array(preds_t_1)[0]
|
||||
fd={synmgt:smg.transpose((0,2,1,3)),
|
||||
synmlt:sml.transpose((0,2,1,3)),
|
||||
blst:blhs,
|
||||
blhst:bls,
|
||||
dpst:dphs.reshape((len(blhs),1)),
|
||||
dist:dis.reshape((len(blhs),1)),
|
||||
dphst:dps.reshape((len(blhs),1))}
|
||||
preds_t_2=sess.run([predictions],feed_dict=fd)
|
||||
preds_t_2=np.array(preds_t_2)[0]
|
||||
preds=preds+w[i-1]*(preds_t_1+preds_t_2)/2
|
||||
tf.reset_default_graph()
|
||||
|
||||
#preds=preds/6
|
||||
preds=np.argmax(preds,axis=1)
|
||||
print(preds.shape)
|
||||
|
||||
with open("prediction_"+fname+"_"+model_name+"_"+name+"_multiple.txt","w") as file:
|
||||
for index,row in df.iterrows():
|
||||
file.write(str(row[0]))
|
||||
file.write("\t")
|
||||
file.write(row[1])
|
||||
file.write("\t")
|
||||
file.write(row[3])
|
||||
file.write("\t")
|
||||
file.write(str(row["label"]))
|
||||
file.write("\t")
|
||||
if index in index_dict:
|
||||
file.write(str(preds[index_dict[index]]))
|
||||
file.write("\t")
|
||||
if preds[index_dict[index]]==row["label"]:
|
||||
file.write(str(1))
|
||||
else:
|
||||
file.write(str(0))
|
||||
else:
|
||||
file.write("Error")
|
||||
file.write("\t")
|
||||
file.write("NaN")
|
||||
file.write("\n")
|
||||
|
|
@ -3,12 +3,13 @@ import progressbar
|
|||
import json
|
||||
import numpy as np
|
||||
import time
|
||||
from multiprocessing import Process
|
||||
from threading import Thread
|
||||
import pickle
|
||||
from ete3 import Tree
|
||||
from read_get_gene_seq import read_gene_seq,create_dict
|
||||
from process_data import get_nearest_neighbors
|
||||
from create_synteny_matrix import create_synteny_matrix_mul
|
||||
from create_synteny_matrix import create_synteny_matrix_mul,update
|
||||
import requests
|
||||
import json
|
||||
|
||||
def create_data_homology_ls(df,n,a,d,ld,ldg,cmap,cimap):
|
||||
buf=open("not_found.txt","w")
|
||||
|
|
@ -54,7 +55,7 @@ def group_seq_by_species(df,g_to_sp):
|
|||
create_dict(ghsp,sph,g_to_sp)
|
||||
return g_to_sp
|
||||
|
||||
def read_gene_sequences(df,lsy,data_dir,fname):
|
||||
def read_gene_sequences(df,lsy,data_dir,fname,name):
|
||||
|
||||
"""The basic idea here is to create a list/dictionary of all the genes by their species.
|
||||
Once the mapping is done, all the respective fasta sequence files are read by Species
|
||||
|
|
@ -101,10 +102,8 @@ def read_gene_sequences(df,lsy,data_dir,fname):
|
|||
except:
|
||||
not_found[gene]=1
|
||||
|
||||
with open("processed/not_found.json","w") as file:
|
||||
with open("processed/not_found_"+name+"_.json","w") as file:
|
||||
json.dump(not_found,file)
|
||||
with open("processed/"+fname+".json","w") as file:#save the data
|
||||
json.dump(data,file)
|
||||
return data
|
||||
|
||||
def intermediate_process(gene_seq,x,y,n,index,sl,sg,ind):
|
||||
|
|
@ -115,63 +114,6 @@ def intermediate_process(gene_seq,x,y,n,index,sl,sg,ind):
|
|||
sl.append(smltemp)
|
||||
ind.append(index)
|
||||
|
||||
def synteny_matrix(gene_seq,hdf,lsy,n,enable_break,sg,sl,ind):
|
||||
#sg=[]
|
||||
#sl=[]
|
||||
t=0
|
||||
#ind=[]
|
||||
|
||||
for index,row in progressbar.progressbar(hdf.iterrows()):
|
||||
g1=str(row[1])
|
||||
g2=str(row[3])
|
||||
x=[]
|
||||
y=[]
|
||||
t+=1
|
||||
try:
|
||||
temp=lsy[g1]
|
||||
except:
|
||||
continue
|
||||
try:
|
||||
temp=lsy[g2]
|
||||
except:
|
||||
continue
|
||||
for i in range(len(lsy[g1]['b'])-1,-1,-1):
|
||||
x.append(lsy[g1]['b'][i])
|
||||
x.append(g1)
|
||||
for k in lsy[g1]['f']:
|
||||
x.append(k)
|
||||
|
||||
for i in range(len(lsy[g2]['b'])-1,-1,-1):
|
||||
y.append(lsy[g2]['b'][i])
|
||||
y.append(g2)
|
||||
for k in lsy[g2]['f']:
|
||||
y.append(k)
|
||||
|
||||
assert(len(x)==len(y))
|
||||
assert(len(x)==(2*n+1))
|
||||
smgtemp,smltemp=create_synteny_matrix_mul(gene_seq,x,y,2*n+1)
|
||||
if np.all(smgtemp==0):
|
||||
continue
|
||||
sg.append(smgtemp)
|
||||
sl.append(smltemp)
|
||||
ind.append(index)
|
||||
"""try:
|
||||
th=Thread(target=intermediate_process,name="TimeOutDetector",args=(gene_seq,x,y,2*n+1,index,sl,sg,ind,))
|
||||
th.start()
|
||||
th.join(30)
|
||||
if th.is_alive():
|
||||
raise SystemError("Long Time")
|
||||
print(row)
|
||||
th.join()
|
||||
except Exception as e:
|
||||
print(e)"""
|
||||
if t==5 and enable_break==1:
|
||||
break
|
||||
#print("Time Taken:",end-start)
|
||||
#print("Average Time:",(end-start)/len(sg))
|
||||
print(t)
|
||||
return np.array(sg),np.array(sl),np.array(ind)
|
||||
|
||||
|
||||
def create_branch_length_padding(bl):
|
||||
maxlen=29
|
||||
|
|
@ -215,3 +157,61 @@ def create_tree_data(treename,df):
|
|||
create_branch_length_padding(branch_lengths_s)
|
||||
create_branch_length_padding(branch_lengths_hs)
|
||||
return np.array(branch_lengths_s),np.array(branch_lengths_hs),np.array(dist),np.array(ns),np.array(nhs)
|
||||
|
||||
def read_data(nop,name):
|
||||
smg=[]
|
||||
sml=[]
|
||||
indexes=[]
|
||||
for i in range(nop):
|
||||
try:
|
||||
with open("temp_"+name+"/thread_"+str(i+1)+"_smg.temp","rb") as file:
|
||||
smg=smg+pickle.load(file)
|
||||
with open("temp_"+name+"/thread_"+str(i+1)+"_sml.temp","rb") as file:
|
||||
sml=sml+pickle.load(file)
|
||||
with open("temp_"+name+"/thread_"+str(i+1)+"_indexes.temp","rb") as file:
|
||||
indexes=indexes+pickle.load(file)
|
||||
except:
|
||||
continue
|
||||
print(len(indexes))
|
||||
return smg,sml,indexes
|
||||
|
||||
def update_rest(data,name):
|
||||
gids={}
|
||||
with open("processed/not_found_"+name+"_.json","r") as file:
|
||||
gids=dict(json.load(file))
|
||||
|
||||
gids=list(gids.keys())
|
||||
|
||||
geneseq={}
|
||||
|
||||
server = "https://rest.ensembl.org"
|
||||
ext = "/sequence/id?type=cds"
|
||||
headers={ "Content-Type" : "application/json", "Accept" : "application/json"}
|
||||
|
||||
for i in progressbar.progressbar(range(0,len(gids)-50,50)):
|
||||
ids=dict(ids=list(gids[i:i+50]))
|
||||
while(1):
|
||||
try:
|
||||
r = requests.post(server+ext, headers=headers, data=str(json.dumps(ids)))
|
||||
if not r.ok:
|
||||
r.raise_for_status()
|
||||
gs=r.json()
|
||||
tgs={}
|
||||
for g in gs:
|
||||
tgs[g["query"]]=g["seq"]
|
||||
geneseq.update(tgs)
|
||||
break
|
||||
except Exception as e:
|
||||
print("Error:",e)
|
||||
continue
|
||||
|
||||
data.update(geneseq)
|
||||
for genes in gids:
|
||||
try:
|
||||
_=data[genes]
|
||||
except:
|
||||
print(genes)
|
||||
update(data,genes)
|
||||
|
||||
print("Gene Sequences Updated Successfully")
|
||||
return data
|
||||
|
|
|
|||
Loading…
Reference in a new issue