Prediction FIx

This commit is contained in:
HarshitGupta11 2019-07-22 18:30:27 +05:30
parent 787557963a
commit b24f3fcc0f
3 changed files with 233 additions and 64 deletions

Binary file not shown.

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

View file

@ -3,12 +3,13 @@ import progressbar
import json import json
import numpy as np import numpy as np
import time import time
from multiprocessing import Process import pickle
from threading import Thread
from ete3 import Tree from ete3 import Tree
from read_get_gene_seq import read_gene_seq,create_dict from read_get_gene_seq import read_gene_seq,create_dict
from process_data import get_nearest_neighbors 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): def create_data_homology_ls(df,n,a,d,ld,ldg,cmap,cimap):
buf=open("not_found.txt","w") 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) create_dict(ghsp,sph,g_to_sp)
return 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. """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 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: except:
not_found[gene]=1 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) json.dump(not_found,file)
with open("processed/"+fname+".json","w") as file:#save the data
json.dump(data,file)
return data return data
def intermediate_process(gene_seq,x,y,n,index,sl,sg,ind): 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) sl.append(smltemp)
ind.append(index) 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): def create_branch_length_padding(bl):
maxlen=29 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_s)
create_branch_length_padding(branch_lengths_hs) 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) 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