mirror of
https://github.com/Priyatham-sai-chand/compara-deep-learning.git
synced 2026-10-05 08:11:34 -07:00
Test the model
This commit is contained in:
parent
294369474a
commit
48cae1021f
3 changed files with 350 additions and 2 deletions
127
prediction.py
Normal file
127
prediction.py
Normal file
|
|
@ -0,0 +1,127 @@
|
||||||
|
import json
|
||||||
|
import gc
|
||||||
|
import pandas as pd
|
||||||
|
import numpy as np
|
||||||
|
import pickle
|
||||||
|
import tensorflow as tf
|
||||||
|
import sys
|
||||||
|
from test_prepare_functions import create_data_homology_ls,read_gene_sequences,synteny_matrix,create_tree_data
|
||||||
|
from access_data_rest import update_rest
|
||||||
|
from process_data import create_map_list
|
||||||
|
|
||||||
|
df=pd.read_csv("input_real_predictions_random_20K.txt",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")
|
||||||
|
label_dict_2=dict(ortholog_one2one=1,other_paralog=0,ortholog_one2many=1,ortholog_many2many=1,within_species_paralog=0,gene_split=4)
|
||||||
|
label_2=[]
|
||||||
|
label_1=[]
|
||||||
|
for index,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
|
||||||
|
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))
|
||||||
|
with open("ng","wb")as file:
|
||||||
|
pickle.dump(lsy,file)
|
||||||
|
|
||||||
|
gene_sequences=read_gene_sequences(df,lsy,"geneseq","gene_sequences")
|
||||||
|
print("Gene Sequences Loaded")
|
||||||
|
|
||||||
|
print("Going to update not found sequences:")
|
||||||
|
gene_sequences=update_rest(gene_sequences)
|
||||||
|
with open("gs","wb")as file:
|
||||||
|
pickle.dump(gene_sequences,file)
|
||||||
|
|
||||||
|
"""with open("ng","rb") as file:
|
||||||
|
lsy=pickle.load(file)
|
||||||
|
with open("gs","rb") as file:
|
||||||
|
gene_sequences=pickle.load(file)"""
|
||||||
|
|
||||||
|
df=df[0:20]
|
||||||
|
n=3
|
||||||
|
smg,sml,indexes=synteny_matrix(gene_sequences,df,lsy,n,0,list(),list(),list())
|
||||||
|
print(len(indexes))
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
try:
|
||||||
|
model=tf.train.import_meta_graph('saved_models/model.ckpt.meta')
|
||||||
|
except:
|
||||||
|
print("Something wrong with the model.")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
with tf.Session() as sess:
|
||||||
|
try:
|
||||||
|
model.restore(sess,"saved_models/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=sess.run([predictions],feed_dict=fd)
|
||||||
|
preds=np.array(preds)[0]
|
||||||
|
preds=np.argmax(preds,axis=1)
|
||||||
|
print(preds.shape)
|
||||||
|
|
||||||
|
with open("prediction.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[2])
|
||||||
|
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,8 +3,11 @@ import pandas as pd
|
||||||
import gzip
|
import gzip
|
||||||
import sys
|
import sys
|
||||||
import progressbar
|
import progressbar
|
||||||
|
import traceback
|
||||||
|
|
||||||
def clear_data(x):
|
def clear_data(x):
|
||||||
|
if x==None:
|
||||||
|
return x
|
||||||
x=x.split()
|
x=x.split()
|
||||||
try:
|
try:
|
||||||
x=x[1]
|
x=x[1]
|
||||||
|
|
@ -33,11 +36,13 @@ def read_data_genome(dir_name,a,dict_ind_genome):
|
||||||
try:
|
try:
|
||||||
for y in ["gene_version","gene_name","gene_source","gene_biotype","gene_id"]:
|
for y in ["gene_version","gene_name","gene_source","gene_biotype","gene_id"]:
|
||||||
data_gene[y]=data_gene[y].apply(clear_data)
|
data_gene[y]=data_gene[y].apply(clear_data)
|
||||||
except:
|
except Exception as e:
|
||||||
|
traceback.print_exc()
|
||||||
|
print(e)
|
||||||
continue
|
continue
|
||||||
#print(data_gene[0:10])
|
#print(data_gene[0:10])
|
||||||
data_gene=data_gene[(data_gene['gene_biotype']=='protein_coding') | (data_gene['gene_source']=='protein_coding')]
|
data_gene=data_gene[(data_gene['gene_biotype']=='protein_coding') | (data_gene['gene_source']=='protein_coding')]
|
||||||
#print(data_gene[data_gene["gene_id"]=="ENSPMGG00000022088"])
|
#print(data_gene[data_gene["gene_id"]=="ENSNGAG00000000407"])
|
||||||
a.append(data_gene)
|
a.append(data_gene)
|
||||||
n=lf[x].split(".")[0]
|
n=lf[x].split(".")[0]
|
||||||
dict_ind_genome[n]=len(a)-1
|
dict_ind_genome[n]=len(a)-1
|
||||||
|
|
|
||||||
216
test_prepare_functions.py
Normal file
216
test_prepare_functions.py
Normal file
|
|
@ -0,0 +1,216 @@
|
||||||
|
import pandas as pd
|
||||||
|
import progressbar
|
||||||
|
import json
|
||||||
|
import numpy as np
|
||||||
|
import multiprocessing
|
||||||
|
from threading import Thread
|
||||||
|
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
|
||||||
|
|
||||||
|
def create_data_homology_ls(df,n,a,d,ld,ldg,cmap,cimap):
|
||||||
|
buf=open("not_found.txt","w")
|
||||||
|
lsy={} #dictionary which stores +/- n genes of the given gene by id. Each key is a gene id which corresponds to the one in center.
|
||||||
|
for _,row in progressbar.progressbar(df.iterrows()):
|
||||||
|
x=row[1]
|
||||||
|
y=row[3]
|
||||||
|
xs=row[2]
|
||||||
|
ys=row[4]
|
||||||
|
try:
|
||||||
|
z=lsy[x]
|
||||||
|
except:
|
||||||
|
try:
|
||||||
|
t2=d[xs.capitalize()]#see if the species exist in genomic maps
|
||||||
|
xl,xr=get_nearest_neighbors(x,xs,n,a,d,ld,ldg,cmap,cimap)
|
||||||
|
if len(xl)!=0:#check if neighboring genes were successfully found
|
||||||
|
lsy[x]=dict(b=xl,f=xr)
|
||||||
|
except Exception as e:
|
||||||
|
buf.write(x+"\t"+xs)
|
||||||
|
buf.write("\n")
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
z=lsy[y]
|
||||||
|
except:
|
||||||
|
try:
|
||||||
|
t2=d[ys.capitalize()]
|
||||||
|
yl,yr=get_nearest_neighbors(y,ys,n,a,d,ld,ldg,cmap,cimap)
|
||||||
|
if len(yl)!=0:
|
||||||
|
lsy[y]=dict(b=yl,f=yr)
|
||||||
|
except Exception as e:
|
||||||
|
buf.write(y+"\t"+ys)
|
||||||
|
buf.write("\n")
|
||||||
|
continue
|
||||||
|
buf.close()
|
||||||
|
return lsy
|
||||||
|
|
||||||
|
def group_seq_by_species(df,g_to_sp):
|
||||||
|
sph=list(df[4])
|
||||||
|
ghsp=list(df[3])
|
||||||
|
sp=list(df[2])
|
||||||
|
gsp=list(df[1])
|
||||||
|
create_dict(gsp,sp,g_to_sp)
|
||||||
|
create_dict(ghsp,sph,g_to_sp)
|
||||||
|
return g_to_sp
|
||||||
|
|
||||||
|
def read_gene_sequences(df,lsy,data_dir,fname):
|
||||||
|
|
||||||
|
"""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
|
||||||
|
and the CDNA sequences for each gene in the species record are read and stored.
|
||||||
|
Thus we don't have to read the same file multiple times."""
|
||||||
|
grouped_genes={}
|
||||||
|
gene_by_species_dict={}
|
||||||
|
grouped_genes=group_seq_by_species(df,grouped_genes)
|
||||||
|
for i in df[2].unique():
|
||||||
|
gene_by_species_dict[i]=[]
|
||||||
|
for i in df[4].unique():
|
||||||
|
gene_by_species_dict[i]=[]
|
||||||
|
|
||||||
|
for x in progressbar.progressbar(lsy):
|
||||||
|
try:
|
||||||
|
species=grouped_genes[x]#get the species
|
||||||
|
except:
|
||||||
|
continue
|
||||||
|
if x not in gene_by_species_dict[species]:#check if the gene already exists in the species dict or not.
|
||||||
|
gene_by_species_dict[species].append(x)
|
||||||
|
xl=lsy[x]['b']
|
||||||
|
xr=lsy[x]['f']
|
||||||
|
for gxl in xl:
|
||||||
|
if gxl=="NULL_GENE":
|
||||||
|
break
|
||||||
|
if gxl not in gene_by_species_dict[species]:
|
||||||
|
gene_by_species_dict[species].append(gxl)
|
||||||
|
for gxr in xr:
|
||||||
|
if gxr=="NULL_GENE":
|
||||||
|
break
|
||||||
|
if gxr not in gene_by_species_dict[species]:
|
||||||
|
gene_by_species_dict[species].append(gxr)
|
||||||
|
|
||||||
|
|
||||||
|
s=[x for x in gene_by_species_dict if len(gene_by_species_dict[x])!=0]#select those species only whose gene sequences we have to read.
|
||||||
|
s=[x.capitalize() for x in s]
|
||||||
|
|
||||||
|
data=read_gene_seq(data_dir,s,gene_by_species_dict)
|
||||||
|
not_found={}
|
||||||
|
for species in gene_by_species_dict:
|
||||||
|
for gene in gene_by_species_dict[species]:
|
||||||
|
try:
|
||||||
|
_=data[gene]
|
||||||
|
except:
|
||||||
|
not_found[gene]=1
|
||||||
|
|
||||||
|
with open("processed/not_found.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):
|
||||||
|
smgtemp,smltemp=create_synteny_matrix_mul(gene_seq,x,y,n)
|
||||||
|
if np.all(smgtemp==0):
|
||||||
|
return
|
||||||
|
sg.append(smgtemp)
|
||||||
|
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():
|
||||||
|
th.join()
|
||||||
|
print(row)
|
||||||
|
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
|
||||||
|
for x in bl:
|
||||||
|
for i in range(len(x),maxlen):
|
||||||
|
x.append(0)
|
||||||
|
|
||||||
|
def create_tree_data(treename,df):
|
||||||
|
t=Tree(treename)
|
||||||
|
branch_lengths_s=[]
|
||||||
|
branch_lengths_hs=[]
|
||||||
|
dist=[]
|
||||||
|
ns=[]
|
||||||
|
nhs=[]
|
||||||
|
for index,row in progressbar.progressbar(df.iterrows()):
|
||||||
|
d=0
|
||||||
|
x=row[2]
|
||||||
|
y=row[4]
|
||||||
|
bl=[]
|
||||||
|
c=0
|
||||||
|
mca=t.get_common_ancestor(x,y)
|
||||||
|
node=t&x
|
||||||
|
while node.up!=mca:
|
||||||
|
d+=node.dist
|
||||||
|
bl.append(node.dist)
|
||||||
|
node=node.up
|
||||||
|
c+=1
|
||||||
|
ns.append(c)
|
||||||
|
c=0
|
||||||
|
branch_lengths_s.append(bl)
|
||||||
|
bl=[]
|
||||||
|
node=t&y
|
||||||
|
while node.up!=mca:
|
||||||
|
d+=node.dist
|
||||||
|
bl.append(node.dist)
|
||||||
|
node=node.up
|
||||||
|
c+=1
|
||||||
|
nhs.append(c)
|
||||||
|
branch_lengths_hs.append(bl)
|
||||||
|
dist.append(d)
|
||||||
|
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)
|
||||||
Loading…
Reference in a new issue