# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# %% [code]
# This Python 3 environment comes with many helpful analytics libraries installed
# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python
# For example, here's several helpful packages to load

import numpy as np # linear algebra
import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)

import matplotlib.pyplot as plt

# Input data files are available in the read-only "../input/" directory
# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory

import os
for dirname, _, filenames in os.walk('/kaggle/input'):
    for filename in filenames:
        print(os.path.join(dirname, filename))

        
def testmap_data(x):
    return x.split('\n', 1)[1].replace('\n', '');

def testmap_group(x):
    return x.split('\n', 1)[0].split('\t',1)[1];

def testmap_id(x):
    return x.split('\n', 1)[0].split('\t',1)[0];

def trainmap_id(x):
    return x.split('\n', 1)[0].split(' ',1)[0];

def trainmap_header(x):
    return x.split('\n', 1)[0];

def trainmap_group(x):
    if not 'OX=' in x:
        return "";
    return x.split('OX=', 1)[1].split(' ',1)[0];

def gobasic_id_ext(x):
    return (x.split('\n', 2)[1]).split(': ',1)[1];


def seqData_init(m):
    zids = np.array(list(map(testmap_id, m)));
    zgroups = np.array(list(map(testmap_group, m)));
    zdats = np.array(list(map(testmap_data, m)), dtype=object);
    return zids, zgroups, zdats;


def seqData_init2(m):
    zids = np.array(list(map(trainmap_id, m)));
    zhead = list(map(trainmap_header, m));
    zgroups = np.array(list(map(trainmap_group, zhead)));
    zdats = np.array(list(map(testmap_data, m)), dtype=object);
    return zids, zhead, zgroups, zdats; #zgroups,

###----------------cluster utils

def s_dist(a, v):
    return np.sum(np.abs(a - v), axis=1);

def sel_max(n):
    return np.argmax(n);

def gen_cents(n, data, reset=True, xdists=[]):
    
    cents = [];
    xdata = np.mean(data, axis=0);
    cents.append(xdata);              ## mean seed?
    d2 = np.zeros((data.shape[0])) + 10000;    
    if not reset:
        d2 = xdists;
    for i in range(n-1):
        dists = s_dist(data, xdata); ## iter 0
        d2 = np.minimum(d2, dists);  ## min distance < dist to mean?
        x = sel_max(d2);
        xdata = data[x, :];
        cents.append(xdata);
        if(i%10 == 0):
            print(d2[x], " iter ",i," peak");
    print(np.mean(d2), "avg minimum distances Initial");
    return cents;


def update_cents(c, data):
    
    cents = c;
    d2 = np.zeros((data.shape[0], len(c)));
    n = len(c);
    d3 = np.zeros(n);
    for i in range(n):
        dists = s_dist(data, cents[i]);
        d2[:,i] = dists; #np.minimum(d2, dists);
    xdists = np.argmin(d2, axis=1);
    for i in range(n):
        cents[i] = np.mean(data[xdists == i, :], axis=0);
    return cents;


def fetch_dists(c,data):

    cents = c;
    d2 = np.zeros((data.shape[0], len(c)));
    n = len(c);
    d4 = np.zeros(data.shape[0]);
    for i in range(n):
        dists = s_dist(data, cents[i]);
        d2[:,i] = dists; #np.minimum(d2, dists);
    xdists = np.argmin(d2, axis=1);
    for i in range(n):
        d4[xdists==i] = d2[xdists == i, i];
    return d4;


def fetch_groups(c,data):

    cents = c;    
    d2 = np.zeros((data.shape[0], len(c)));
    n = len(c);
    for i in range(n):
        dists = s_dist(data, cents[i]);
        d2[:,i] = dists; #np.minimum(d2, dists);
    xdists = np.argmin(d2, axis=1);
    return xdists;

##-------------------------
## series utils?

def s0_sampler(x, amine_code):
    
    n = x.shape[0];
    m = len(amine_code);
    
    ret = [];
    
    for i in range(n):
        if i%10000 == 0:
            print("s0", i);
            
        k2 = x[i];
        #for j in range(0,len(k2), w_stride):
        xstr = k2;#j:j+w_stride];
        a1 = [];
        for k in range(m):
            zr = xstr.count(amine_code[k]);
            a1.append(zr);
        ret.append(a1);
        #ret.append
    return ret;


def s1_sampler(x, amine_code, w_stride):
    
    n = x.shape[0];
    m = len(amine_code);
    
    ret = [];
    r_id = [];
    
    for i in range(n):

        if i%10000 == 0:
            print("s1", i);
        
        k2 = x[i];
        for j in range(0,len(k2), w_stride):
            xstr = k2[j:j+w_stride];
            a1 = [];
            if len(xstr)<w_stride:
                continue;
            for k in range(m):
                zr = xstr.count(amine_code[k]);
                a1.append(zr);
            ret.append(a1);
            r_id.append(i);
        #ret.append

    return ret, r_id;



def s2_sampler(xs, xm, r=4, sr2=50):
    n = len(xs);
    #r = 4;
    s3 = [];
    print("s2sampler run::");
    for i in range(n):
        if i%2500 == 0:
            print(i);
        m = len(xs[i]);
        #s2 = np.array([i]);
        for j in range(0,m,r):
            if j>m-r:
                continue;
            s2 = np.array([i]);
            for k in range(j, j+r):
                s1 = np.array(xs[i][k]) - (xm[i, :] * sr2);
                s2= np.append(s2,s1, axis=0);
            s3.append(s2);
    return s3;


def ht_sampler(x, amine_code, wid=160):
    
    n = x.shape[0];
    m = len(amine_code);
    
    ret = [];
    
    for i in range(n):
        if i%10000 == 0:
            print("ht samp", i);
            
        k2 = x[i];
        
        
        #for j in range(0,len(k2), w_stride):
        xstr = k2[:wid];#j:j+w_stride];
        a1 = [];
        for k in range(m):
            zr = xstr.count(amine_code[k]);
            a1.append(zr);
            
        xstr = k2[-wid:];#j:j+w_stride];
        #a1 = [];
        for k in range(m):
            zr = xstr.count(amine_code[k]);
            a1.append(zr);
        ret.append(a1);
        #ret.append
    return ret;


##------------------

skipGo = False;
        
f = open("/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta");
w = f.read();
f.close();

f5 = open("/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta");
zstr = f5.read();
f5.close();

f = open("/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo");
xstr = f.read();
f.close();

#w2 = w.split('\n>')

f2 = np.array(pd.read_csv("/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv", sep='\t'));

f4 = np.array(pd.read_csv("/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset-taxon-list.tsv", sep='\t', encoding='iso-8859-1'));

f6 = np.array(pd.read_csv("/kaggle/input/cafa-5-protein-function-prediction/IA.txt", sep='\t'));

w2 = w.split('\n>');

w2[0] = w2[0][1:];

zstr2 = zstr.split('\n>');

zstr2[0] = zstr2[0][1:];

xstr2 = xstr.split("\n[Term]");

xstr2 = xstr2[1:];


test_id, test_g, test_seq = seqData_init(zstr2);

train_id, train_h, train_g, train_seq = seqData_init2(w2); #


x3 = list(map(gobasic_id_ext, xstr2));

gb_index =  dict(zip(x3, range(len(x3))));

gb_data = xstr2;


tid_index = dict(zip(train_id, range(train_id.shape[0])));


xmatch = 0;

xm_list = np.zeros((test_id.shape[0]));

for i in range(test_id.shape[0]):
    if test_id[i] in tid_index:
        xmatch += 1;
        xm_list[i] = 1;

print(xmatch, "id'd test set proteins");

print("train set:", train_seq.shape, ", test set:", test_seq.shape, ", go terms:", len(gb_data));

print(f2.shape, "protein:go table");

pg_ids = pd.unique(f2[:,0]);
pg_gos = pd.unique(f2[:,1]);

pg_class = pd.unique(f2[:,2]);

print(np.sum(f2[:,2]==pg_class[0]), np.sum(f2[:,2]==pg_class[1]), np.sum(f2[:,2]==pg_class[2]), pg_class);

print("listed protein ids:", pg_ids.shape, ", listed go terms:", pg_gos.shape)

pg_ids = train_id; ##?

pg_index = dict(zip(pg_ids, range(pg_ids.shape[0])));

pg_index2 = dict(zip(pg_gos, range(pg_gos.shape[0])));

pg_mat = np.zeros((pg_ids.shape[0], pg_gos.shape[0]), np.int8);

xmatch = 0;

for i in range(test_id.shape[0]):
    if test_id[i] in pg_index:
        xmatch += 1;

print(xmatch, "go'd test set proteins");

xmatch2=0;

for i in range(train_id.shape[0]):
    if train_id[i] in pg_index:
        xmatch2 += 1;

print(xmatch2, "go'd train set proteins");

test_human = test_g == '9606';

print(np.sum(test_human & (xm_list==0)), " unlabeled h test");
print(np.sum(test_human & (xm_list==1)), " labeled h test");

print(np.sum((test_human==0) & (xm_list==0)), " unlabeled not-h test");
print(np.sum((test_human==0) & (xm_list==1)), " labeled not-h test");

if not skipGo:
    
    f6check = f6[:,1] ==0;
    
    zeroVals = f6[f6check, 0];
    
    zD = dict(zip(zeroVals, range(zeroVals.shape[0])));

    
    for i in range(f2.shape[0]):

        if i%250000 == 0:
            print("iter ",i);
            
                    
        #if f2[i, 1]  in zD: #Zero IA check
            
        pg_mat[pg_index[f2[i, 0]], pg_index2[f2[i,1]]] = 1;
            #pg_mat[tid_index[f2[i, 0]], pg_index2[f2[i,1]]] = 1;
    

    print("completed", pg_mat.shape);

    pgs1 = np.sum(pg_mat, axis=0)+1;
    pgs2 = np.sum(pg_mat, axis=1);

    plt.figure();
    plt.plot(np.minimum(np.sort(pgs1)[::500], 100));
    plt.plot(np.minimum(np.sort(pgs2)[::500], 100));


    plt.figure();
    plt.plot(np.maximum(np.sort(pgs1)[::500], 100));
    plt.plot(np.maximum(np.sort(pgs2)[::500], 100));

    
    
#amines = ["R","H","K", "D","E", "S","T","N","Q", "C","G","U","P", "A","V","I","L","M","F","Y","W"];

amines = ['L', 'S', 'A', 'E', 'G', 'V', 'K', 'P', 'T', 'R', 'D', 'I', 'Q', 'N', 'F', 'Y', 'H', 'M', 'C', 'W', 'U'];


##samp0 = np.array(s0_sampler(train_seq, amines));


samp0 = np.array(ht_sampler(train_seq, amines));


print(samp0[0]);

#samp1, s1_id = s1_sampler(train_seq, amines, 160); ##40/5 :: 160/20?

#print(samp1[0]);

#print(len(samp0), len(samp1), "0samp, 1samp")


###s0 = np.array(samp0) / np.sum(samp0,axis=1).reshape(samp0.shape[0], 1) * 160;#c1_groups = fetch_groups(cen1, s0);

s0 = samp0; #np.array(samp0) / np.sum(samp0,axis=1).reshape(samp0.shape[0], 1) * 160;#c1_groups = fetch_groups(cen1, s0);


#s1 = np.array(samp1); #[::5]); ## 20%

cen1 = gen_cents(28, s0);  ## No of groups ##-----------------------------------------#

plt.figure();
#c1_dists = fetch_dists(cen1, s0);
#plt.plot(np.sort(c1_dists)[::1000] /160 * 100);

for i in range(12):       ## No of updates ## ----------------------#
    print("cent iter", i);
    
    oldcen = np.array(cen1) + 0;
    
    cen1 = update_cents(cen1, s0);
    
    print(np.sum(np.abs(oldcen - np.array(cen1))), "centr delta");
    if i%5 == 0:
        c1_dists = fetch_dists(cen1, s0);
        plt.plot(np.sort(c1_dists)[::1000][10:-10] /320 * 100); ## % maximum error? == 160*2
    
c1_dists = fetch_dists(cen1, s0);
plt.plot(np.sort(c1_dists)[::1000][10:-10] /320 * 100);

print(np.mean(c1_dists), "final avg cent dist");

plt.savefig("c1_dists.jpg");

c1_groups = fetch_groups(cen1, s0);
plt.figure();

for i in range(len(cen1)):
    xgd = c1_dists[c1_groups == i];
    plt.plot(np.sort(xgd)[::100][2:-3]);    

plt.savefig("c1_tree.jpg");

s1_m = [];

c_block = [];

##----------------------------------------
"""
for i in range(len(cen1)):

    if np.sum(c1_groups==i)<1000:
        continue;

    samp1_L, sL_id = s1_sampler(train_seq[c1_groups==i], amines, 160); ##40/5 :: 160/20?
    s1_m.append(samp1_L);

plt.figure();

for i in range(len(s1_m)):
    print("batch iter ",i);
    s1 = np.array(s1_m[i]);

    cen2 = gen_cents(8, s1);

    #plt.figure();
#c2_dists = fetch_dists(cen2, s1);
#plt.plot(np.sort(c2_dists)[::1000] /160 * 100);

    for j in range(10):
        print("cent iter", i, j);
        cen2 = update_cents(cen2, s1);
    
        #if i%5 == 0:
        #    c2_dists = fetch_dists(cen2, s1);
        #    plt.plot(np.sort(c2_dists)[::1000][10:-10] /320 * 100);

    c_block.append(cen2);
    
    c2_dists = fetch_dists(cen2, s1);

    print(np.mean(c2_dists), "final avg cent dist");

    plt.plot(np.sort(c2_dists)[::100][2:-3]);
    #plt.savefig("c2_dists_"+str(i)+".jpg")


    "" "c2_groups = fetch_groups(cen2, s1);
    plt.figure();

    for j in range(len(cen2)):
        xgd = c2_dists[c2_groups == j];
        plt.plot(np.sort(xgd)[::100]);    
        
    plt.savefig("c2_tree_"+str(i)+".jpg")"" ";
    
plt.savefig("c2_dists_All.jpg")

plt.figure();
plt.plot(np.array(cen1).transpose());
plt.savefig("cent1.jpg")

cb = np.array(c_block)
cb2 = cb.reshape(np.prod(cb.shape[:2]), 21);

cen3 = gen_cents(32, cb2);

for j in range(10):
    print("cent iter", i);
    cen3 = update_cents(cen3, cb2);

plt.figure();
plt.plot(np.array(cen3).transpose());

c3_dists = fetch_dists(cen3, cb2);

plt.figure();
plt.plot(np.sort(c3_dists))

"""

##---

#c1_groups
#pg_mat
#pgs1 pgs2

plt.figure();

go_param = [];

go_param2 = [];

print("Id align check?");
miscount=0;

for i in range(pg_mat.shape[0]):
    if not pg_ids[i] == train_id[i]: #tid_index[i] == 
        miscount += 1;

print("Alignment errors?", miscount);

##doesn't detect fix? ## else
## fixed / but check broken

#break; 

P_sel_max = 30; ##No of predictions ##---------------------------------------#

for i in range(len(cen1)):
    
    print("go sampling", i);
    
    gs0 = c1_groups == i;

    sub_p = pg_mat[gs0, :];

    sub_r = np.sum(sub_p, axis=0);

    sub_n = np.mean(sub_p, axis=0);
    
    go_param.append(sub_n);
    
    sub_p2 = (sub_r / pgs1);

    sub_v = sub_n * (sub_r / pgs1);
    
    sub_x = sub_n / (pgs1 / pg_mat.shape[0]); ## rel freq?
    
    #print(np.sum(sub_n>0), "subn +Zero counts?");
    #print(np.sum(sub_n>0.1), "subn +10% counts?");
    #print(np.sum(sub_n>0.5), "subn +50% counts?");
    #print(np.sum(sub_p2>0.5), "subp +50% counts?");
    #print(np.sum(sub_x>2), "subx *2 counts?");
    #print(np.sum(sub_x>5), "subx *5 counts?");
    
    print(np.max(sub_x), "subx peak?");
    
    go_param2.append(np.flip(np.argsort(sub_x))[:P_sel_max]);  ##?
    
    print(np.max(sub_v));
    print(np.argmax(sub_v));
    
    #plt.plot(np.minimum(np.flip(np.sort(sub_x))[:5000], 32));
    #print("?", x3[np.argmax(sub_v)], gb_data[np.argmax(sub_v)]);
    
    x_ind = np.flip(np.argsort(sub_v));
    
    for i in range(-1):
    
        print(pg_gos[x_ind[i]], gb_data[gb_index[pg_gos[x_ind[i]]]]);
        
        
#plt.plot(np.array(cen2).transpose());
#plt.savefig("subX.jpg")

#cenS = [];

cenS = np.zeros((len(pg_gos), s0.shape[1]));

cenS_v = np.zeros((len(pg_gos), s0.shape[1]));

px_mat = np.zeros((len(pg_gos), len(pg_gos)), np.uint16);


for j in range(len(pg_gos)):
    
    pv1 = pgs1 > (pgs1[j]/2); ## shrink axis2 if axis1 large? (##remove smalls from large)
    
    pv1 = pv1 & (pgs1 < (pgs1[j]*4)); ## ##remove large from smalls
    

    #gs0 = c1_groups == i;
    
    #sub_pg = pg_mat[pg_mat[:, j]==1, :];
    
    px_mat[j, pv1] = np.sum(pg_mat[pg_mat[:, j]==1, :][:,pv1], axis=0);
    
    #px_mat[j, :] = np.sum(pg_mat[:, j] * pg_mat, axis=0);
    
    #px_mat[j, :] = np.sum(pg_mat * pg_mat[:, j].reshape((142246, 1)), axis=0)

    #sub_p = s0[pg_mat[:, j]==1, :];
    
    #cenS.append( np.mean(sub_p, axis=0));
    #cenS[j, :] = np.mean(sub_p, axis=0);
    
    #cenS_v[j, :] = np.mean(np.abs(sub_p - cenS[j, :]), axis=0);
    
    if j%1000 ==0:
        print("censII iter", j);
        
    if j<1000:
        if j%50 ==0:
            print("censII iter", j);


joinz = [];

j_mat = np.zeros(px_mat.shape, np.uint8);

for j in range(len(pg_gos)):

    sub_px = px_mat[j, :];
    
    v1 = sub_px > (sub_px[j]/2);
    
    #for k in range(v1.shape[0]):
        
        #if px_mat[k,j] > px_mat[k,k]/2:
            
            #joinz.append([j, k]);
            
    j_mat[j,v1] = 1;
            
    if j%1000 ==0:
        print("censIII iter", j);


for j in range(len(pg_gos)):

    #gs0 = c1_groups == i;

    sub_p = s0[pg_mat[:, j]==1, :];
    
    #cenS.append( np.mean(sub_p, axis=0));
    cenS[j, :] = np.mean(sub_p, axis=0);
    
    cenS_v[j, :] = np.mean(np.abs(sub_p - cenS[j, :]), axis=0);
    
    if j%1000 ==0:
        print("cens iter", j);
    

    
plt.figure();
plt.plot(np.sort(np.sum(cenS_v, axis=1))[::100]);
    
#xndata = np.mean(cenS, axis=0);
#cents.append(xdata);              ## mean seed?
#d2 = np.zeros((data.shape[0])) + 10000;    
#    if not reset:
#        d2 = xdists;
#    for i in range(n-1):
#xndists = s_dist(cenS, xndata); ## iter 0
#d2 = np.minimum(d2, dists);  ## min distance < dist to mean?

#xn = np.argmax(xndists); ##sel_max(xndists);
    
cenX1 = gen_cents(45, cenS);  ## No of groups ##-----------------------------------------#

#plt.figure();
#c1_dists = fetch_dists(cen1, s0);
#plt.plot(np.sort(c1_dists)[::1000] /160 * 100);

for i in range(25):       ## No of updates ## ----------------------#
    print("centX iter", i);
    
    #oldcen = np.array(cen1) + 0;
    
    cenX1 = update_cents(cenX1, cenS);
    
    
c1X_groups = fetch_groups(cenX1, cenS);
c1X_dists = fetch_dists(cenX1, cenS);

plt.figure();
plt.plot(np.sort(c1X_groups)[::100]);
plt.savefig("centX_groups.jpg");


plt.figure();

for i in range(len(cenX1)):
    xgd = c1X_dists[c1X_groups == i];
    plt.plot(np.sort(xgd)[::100]);    

plt.savefig("c1X_tree.jpg");

plt.figure();
plt.plot(np.sort(c1X_dists)[::100]);    
plt.savefig("c1X_dists.jpg");


intra_dists = fetch_dists(cenX1, np.array(cen1));
intra_groups = fetch_groups(cenX1, np.array(cen1));

## type A / type B? 

#intra_dists = fetch_dists(cen1, np.array(cenX1));
#intra_groups = fetch_groups(cen1, np.array(cenX1));

plt.figure();

for i in range(len(cenX1)):
    xgd = intra_dists[intra_groups == i];
    plt.plot(np.sort(xgd));    

plt.savefig("intra_tree.jpg");

plt.figure();
plt.plot(np.sort(intra_groups));
plt.savefig("intra_groups.jpg");


plt.figure();
plt.plot(np.array(cen1).transpose());
plt.savefig("centroids_1.jpg");

plt.figure();
plt.plot(np.array(cenX1).transpose()[:,::10]);
plt.savefig("centroids_X1.jpg");

plt.figure();
c2_dists = fetch_dists(cenX1, s0);

plt.plot(np.sort(c2_dists)[::1000][10:-10] /320 * 100);
print(np.mean(c2_dists), "go_Cross avg cent dist");

plt.savefig("c2_dists.jpg");

c2_groups = fetch_groups(cenX1, s0);
plt.figure();

for i in range(len(cenX1)):
    xgd = c2_dists[c2_groups == i];
    plt.plot(np.sort(xgd)[::100][2:-3]);    

plt.savefig("c2_tree.jpg");


errX = np.zeros((s0.shape[0]));

errY = np.zeros((s0.shape[0]));

errC = np.zeros((pg_mat.shape[1]));

errC2 = np.zeros((pg_mat.shape[1]));

for i in range(s0.shape[0]):
    
    id1 = c2_groups[i];
    
    v1 = c1X_groups == id1;
    
    correct = np.sum(pg_mat[i, v1]);
    
    errC[v1] += pg_mat[i, v1];
    
    errC2 += v1;
    
    countA = np.sum(v1);
    
    countB = np.sum(pg_mat[i, :]);
    
    if i<40:
    
        print("Sample", i, correct, countA, countB, ", ", correct / countA, "prec", correct/ countB, "recall", );
        
    errX[i] = correct;
    
    errY[i] = countA;

pgt = np.sum(pg_mat, axis=0);

pgx = errC / pgt;
pgy = errC / (errC2+1);

plt.figure();
plt.plot(np.sort(pgx)[::100]);
plt.savefig("class_Prec.jpg");

plt.figure();
plt.plot(np.sort(pgy)[::100]);
plt.savefig("class_Recall.jpg");

plt.figure();
plt.plot(np.sort(errC)[::100]);
plt.plot(np.sort(errC2)[::100]);
plt.savefig("class_correctx_pred.jpg");


    
plt.figure();
plt.plot(np.sort(errX / errY)[::1000]);
plt.savefig("training_prec.jpg");

print(np.mean(errX / errY), "Mean training precision");
print(np.mean(errX / np.sum(pg_mat, axis=1)), "Mean training recall");

print(np.sum(pg_mat) / np.sum(errY), "true / predict proportions (counts)");
print(np.sum(errX) / np.sum(pg_mat), "correct / true proportions (sum recall)");
print(np.sum(errX) / np.sum(errY), "correct / predict proportions (sum precision)");

plt.figure();
plt.plot(np.sort(errX / np.sum(pg_mat, axis=1))[::1000]);
plt.savefig("training_recall.jpg");

plt.figure();
plt.plot(np.sort(errX)[::1000]);
plt.plot(np.sort(errY)[::1000]);
plt.plot(np.sort(np.sum(pg_mat, axis=1))[::1000]);
plt.savefig("trainset_totals.jpg");


plt.figure();
plt.plot(np.sort(errX)[::1000]);
plt.savefig("trainset_correct.jpg");

###-------------------------------------------------------------------------------------

## test?

#"""
#samp0_Test = np.array(s0_sampler(test_seq, amines));

samp0_Test = np.array(ht_sampler(test_seq, amines));

print(samp0_Test[0]);



#s0_Test = np.array(samp0_Test) / np.sum(samp0_Test,axis=1).reshape(samp0_Test.shape[0], 1) * 160;

s0_Test = samp0_Test; #) / np.sum(samp0_Test,axis=1).reshape(samp0_Test.shape[0], 1) * 160;

c1_groups_T = fetch_groups(cen1, s0_Test);

c1_dists_T = fetch_dists(cen1, s0_Test);


c2_groups_T = fetch_groups(cenX1, s0_Test);

# c1X_groups = fetch_groups(cenX1, cenS);

n = s0_Test.shape[0];

limit_G = 250;

xrange = np.arange(pg_mat.shape[1]);

f = open("submission.tsv", "w");

for i in range(n): #(20):
    
    id1 = c2_groups_T[i];
    
    v1 = c1X_groups == id1;
    
    v2 = xrange[v1][:limit_G];
    
    ids = test_id[i];
    
    goids = pg_gos[v2];
    
    if len(goids)<1:
        print("Type II empty guess?", ids);
    
    for j in range(len(goids)):
    
        f.write(ids + "\t"+ goids[j]+"\t"+ str(0.5)[:5] + "\n");        
        
f.close();

print("res2?");



print("compressed Result?");

n = s0_Test.shape[0];
n2 = 5;

f = open("x_submission_Old.tsv", "w");

for i in range(n):
    
    tx = c1_groups[i];
    
    preds = go_param2[tx];
    
    vals = np.minimum(np.maximum(go_param[tx][preds] * 0.8, 0.001),1.0);
    
    ids = test_id[i];
    
    goids = pg_gos[preds];
    
    for j in range(P_sel_max):
        
        if i<n2:
            print(ids, goids[j], vals[j]); #pg_index2[##pg_gos[preds[j]]
            
        if go_param[tx][preds][j] ==0:
            print("zero_Guess",i,j);
            
        if not go_param[tx][preds][j] ==0:
            f.write(ids + "\t"+ goids[j]+"\t"+ str(vals[j])[:5] + "\n");

        
f.close();


"""

testMap = np.zeros((samp0_Test.shape[0], pg_mat.shape[1]));

for i in range(len(cen1)):
    
    testMap[c1_groups_T==i, :] = go_param[i];
    
    
n = samp0_Test.shape[0];

xout = [];

for i in range(n):
    
    if i%500 == 0:
        print("xStep ",i);
    
    v = np.argsort(testMap[i,:])[:10];
    
    for j in range(10):
        
        xout.append([test_id[i], gb_index[ pg_gos[v[j]] ], testMap[i, v[j]] ]);
    
                     
print("done");
print(len(xout), xout[0]); #""";

f6check = f6[:,1] ==0;
    
zeroVals = f6[f6check, 0];

predictX = pd.unique(np.array(go_param2).flatten());

skipZero = True;

if not skipZero:

    for i in range(predictX.shape[0]):
    
        if pg_gos[predictX[i]] in zeroVals:
            print(pg_gos[predictX[i]], "weight 0, iter ", i);
            

#pd.DataFrame(j_mat).to_csv("go_vec.csv");

print("save go vec?");

## np.save(j_mat, "go_vec.npy");

print("save go vec Complete");

#pd.DataFrame(pg_mat).to_csv("protein_go.csv");

# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using "Save & Run All" 
# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session