Test the model

This commit is contained in:
HarshitGupta11 2019-06-26 17:41:43 +05:30 committed by GitHub
parent 294369474a
commit 48cae1021f
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 350 additions and 2 deletions

127
prediction.py Normal file
View 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")

View file

@ -3,8 +3,11 @@ import pandas as pd
import gzip
import sys
import progressbar
import traceback
def clear_data(x):
if x==None:
return x
x=x.split()
try:
x=x[1]
@ -33,11 +36,13 @@ def read_data_genome(dir_name,a,dict_ind_genome):
try:
for y in ["gene_version","gene_name","gene_source","gene_biotype","gene_id"]:
data_gene[y]=data_gene[y].apply(clear_data)
except:
except Exception as e:
traceback.print_exc()
print(e)
continue
#print(data_gene[0:10])
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)
n=lf[x].split(".")[0]
dict_ind_genome[n]=len(a)-1

216
test_prepare_functions.py Normal file
View 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)