# 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 in 

# Input data files are available in the "../input/" directory.
# For example, running this (by clicking run or pressing Shift+Enter) will list the files in the input directory

import os
print(os.listdir("../input"))

# Any results you write to the current directory are saved as output.

import numpy as np
from keras import layers
from keras.layers import Input, Dense, Activation, ZeroPadding2D, BatchNormalization, Flatten, Conv2D
from keras.layers import AveragePooling2D, MaxPooling2D, Dropout, GlobalMaxPooling2D, GlobalAveragePooling2D, Concatenate
from keras.models import Model
import pandas as pd
import gc
from sklearn.preprocessing import OneHotEncoder

import tensorflow as tf
import keras.backend as K

Y_train = pd.read_csv('../input/train.csv')
files = os.listdir('../input/train/')
input_shape = (512,512,4)


Y_labels = np.zeros((int(len(files)/4.),28))

for c,i in enumerate(list(Y_train.Target.str.split(' '))):
    for j in i:
        Y_labels[c,int(j)] += 1

from sklearn.metrics import accuracy_score, f1_score     
        
def IoU(y_true, y_pred, offset = 0.001):
    thr = 0.1
    mx = 0
    
    rng = np.arange(0.1,0.9)
    
    for i in rng:
        predictions = y_pred
        predictions[predictions>=i] = 1
        predictions[predictions<i] = 0
        
        inter = np.logical_and(y_true, predictions)
        union = np.logical_or(y_true, predictions)
        
        inter = np.sum(inter)
        union = np.sum(union)
        
        out = np.round((inter + offset)/(union + offset),3)
        if out >= mx:
            mx = out
            thr = i
    
    ret = 'Max IoU achieved is ' + str(mx) + ' at threshold ' + str(thr)
    print(ret)
    


def f1_sc(y_true, y_pred):
    
    thr = 0.1
    mx = 0
    
    rng = np.arange(0.1,0.9)
    
    for i in rng:
        predictions = y_pred
        predictions[predictions>=i] = 1
        predictions[predictions<i] = 0
        
        out = f1_score(y_true = y_true, y_pred = predictions, average = 'macro')
        if out >= mx:
            mx = out
            thr = i
    ret = 'Max F1 Score achieved is ' + str(mx) + ' at threshold ' + str(thr)

    print(ret)

def acc_sc(y_true, y_pred):

    thr = 0.1
    mx = 0
    
    rng = np.arange(0.1,0.9)
    
    for i in rng:
        predictions = y_pred
        predictions[predictions>=i] = 1
        predictions[predictions<i] = 0
        
        out = accuracy_score(y_true = y_true, y_pred = predictions)
        if out >= mx:
            mx = out
            thr = i
    ret = 'Max Accuracy Score achieved is ' + str(mx) + ' at threshold ' + str(thr)

    print(ret)

import tensorflow as tf

def EMR(y_true,y_pred):
    total = K.cast(K.cast(K.sum(K.abs(y_true-K.round(y_pred + tf.constant(0.4))),axis = 1),'bool'),'int32')
    total = K.sum(total)
    num_of_s = K.cast(K.shape(y_true)[0],'int32')
    
    return (num_of_s-total)/num_of_s
    
def F1(y_true, y_pred):
    y_pred = K.round(y_pred + tf.constant(0.4))
    tp = K.sum(K.cast(y_true*y_pred, 'float'), axis=0)
    tn = K.sum(K.cast((1-y_true)*(1-y_pred), 'float'), axis=0)
    fp = K.sum(K.cast((1-y_true)*y_pred, 'float'), axis=0)
    fn = K.sum(K.cast(y_true*(1-y_pred), 'float'), axis=0)
    
    p = tp / (tp + fp + K.epsilon())
    r = tp / (tp + fn + K.epsilon())
    
    f1 = 2*p*r / (p+r+K.epsilon())
    f1 = tf.where(tf.is_nan(f1), tf.zeros_like(f1), f1)
    return K.mean(f1)
    
    
def IoUN(y_true, y_pred, smooth=0.001):

    intersection = K.sum(K.abs(y_true * K.round(y_pred + tf.constant(0.4))), axis=-1)
    sum_ = K.sum(K.abs(y_true) + K.abs(y_pred), axis=-1)
    jac = (intersection + smooth) / (sum_ - intersection + smooth)
    return jac

def PrModel(input_shape):

    X_input = Input(input_shape)

    # Layer 1
    X = Conv2D(32, (4, 4), name = 'conv0')(X_input)
    X = BatchNormalization(axis = -1, name = 'bn0')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    X = MaxPooling2D((2, 2), name='max_pool0')(X)
    
    # Layer 2
    X = Conv2D(32, (4, 4), strides = 1, name = 'conv1', padding = 'same')(X)
    X = BatchNormalization(axis = -1, name = 'bn1')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    X = MaxPooling2D((2, 2), name='max_pool1')(X)
    
    # Layer 3
    X = Conv2D(32, (4, 4), strides = 2, name = 'conv2', padding = 'valid')(X)
    X = BatchNormalization(axis = -1, name = 'bn2.1')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    X = MaxPooling2D((2, 2), name='max_pool2')(X)
    
    # Layer 4
    X = Conv2D(64, (4, 4), strides = 2, name = 'conv3', padding = 'same')(X)
    X = BatchNormalization(axis = -1, name = 'bn3')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    X = MaxPooling2D((2, 2), name='max_pool3')(X)
    
    # Layer 5
    X = Conv2D(64, (4, 4), strides = 1, name = 'conv4', padding = 'same')(X)
    X = BatchNormalization(axis = -1, name = 'bn4')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    X = MaxPooling2D((2, 2), name='max_pool4')(X)
    
    # Layer 6
    X = Conv2D(128, (2, 2), name = 'conv5', padding = 'valid')(X)
    X = BatchNormalization(axis = -1, name = 'bn5')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    
    X = Flatten()(X)
    X = Dense(256, name='fc1')(X)
    X = BatchNormalization(axis = -1, name = 'bn6')(X)
    X = Activation('relu')(X)
    X = Dense(128, name='fc2')(X)
    X = BatchNormalization(axis = -1, name = 'bn7')(X)
    X = Activation('relu')(X)
    X = Dropout(rate = 0.3)(X)
    
    X = Dense(28, activation='sigmoid', name='fc3')(X)

    model = Model(inputs = X_input, outputs = X, name='ProteinRecognizer')
    
    
    return model
    
full = PrModel((512,512,4))
full.summary()

full.compile(optimizer = 'adam', loss = 'binary_crossentropy', metrics = [IoUN,EMR,F1])

np.random.seed(0)

valid_idx = np.random.randint(0,len(list(Y_train.Id)), 128)
train_idx = set(np.arange(0,len(list(Y_train.Id)))) - set(valid_idx)

training_set = [list(Y_train.Id)[x] for x in train_idx]
valid_set = [list(Y_train.Id)[x] for x in valid_idx]

training_labels = Y_labels[list(train_idx)]
valid_labels = Y_labels[list(valid_idx)]

######################### Class Weights ####################################


from sklearn.utils import compute_class_weight

class_list = []

for item in Y_train.Target.values:
    cls = [i for i in item.split(' ')]
    class_list += cls

final = np.ndarray.astype(np.array(class_list),int)

weights_raw = (len(files)-128) / (28 * np.bincount(final))
#weights_raw = np.log(weights_raw)

weights = dict(zip(sorted(np.unique(final)), weights_raw))

######################### Prepare minority classes for Preprocessing ####################################

class_list = []

for item in Y_train.Target.values:
    cls = [i for i in item.split(' ')]
    class_list += cls

final = np.ndarray.astype(np.array(class_list),int)
bin_cnts = np.bincount(final)

bincounts = dict(zip(np.arange(0,28),bin_cnts))
import operator
sorted_bincounts = sorted(bincounts.items(), key=operator.itemgetter(1))

class_indexes = dict()
for key,value in sorted_bincounts:
    if value < 800:
        class_indexes[str(key)] = [ val for val in Y_train[Y_train.Target.str.contains(str(key))].index.values if val not in valid_idx]

###################################    Data Generator     ############################################

from keras.preprocessing.image import ImageDataGenerator

datagen = ImageDataGenerator(
    featurewise_center=True,
    featurewise_std_normalization=True,
    horizontal_flip=True,
    vertical_flip=True,
    data_format = "channels_last")

######################################################################################################

import cv2, time

for epoch in range(0,1):
    
    batch_size = 32
    data = np.zeros((batch_size,512,512,4))
    generated_data = np.zeros((len(class_indexes),512,512,4))
    Y_gen = np.zeros((len(class_indexes),28))

    path = '../input/train/'
    
    print('Start batching')

    for batches in range(0,int(len(training_set)/batch_size)):
        print('Batch: ' + str(batches+1) + '/' + str(int(len(training_set)/batch_size)))
        
        fl = 0
        gen_fl = 0
        for j in training_set[batches*batch_size:(batches+1)*batch_size]:
            
            ####################### Normal Batches ################################
            green = cv2.imread(path + j + '_green.png', cv2.IMREAD_GRAYSCALE)
            red = cv2.imread(path + j + '_red.png', cv2.IMREAD_GRAYSCALE)
            blue = cv2.imread(path + j + '_blue.png', cv2.IMREAD_GRAYSCALE)
            yellow = cv2.imread(path + j + '_yellow.png', cv2.IMREAD_GRAYSCALE)
            
            tmp = np.stack((green,red,blue,yellow), axis = -1)/255.
            
            data[fl,:,:,:] = tmp
            fl += 1
               
        Y = training_labels[(batches*batch_size):((batches+1)*batch_size)]
        
        full.fit(x = data, y = Y, epochs = 1, class_weight = weights) 
        
            ####################### Minority Class Generated Batches #############

#        if batches % 2==0:
#            for key in class_indexes.keys():
#                idx = np.random.choice(class_indexes[str(key)])
#                
#                green = cv2.imread(path + Y_train.Id[idx] + '_green.png', cv2.IMREAD_GRAYSCALE)
#                red = cv2.imread(path + Y_train.Id[idx] + '_red.png', cv2.IMREAD_GRAYSCALE)
#                blue = cv2.imread(path + Y_train.Id[idx] + '_blue.png', cv2.IMREAD_GRAYSCALE)
#                yellow = cv2.imread(path + Y_train.Id[idx] + '_yellow.png', cv2.IMREAD_GRAYSCALE)
#                
#                tmp = np.stack((green,red,blue,yellow), axis = -1)/255.
#                
#                generated_data[gen_fl,:,:,:] = tmp
#                
#                Y_gen[gen_fl,:] = Y_labels[idx]
#                gen_fl += 1
#                
#            datagen.fit(generated_data)
#            for x_batch, y_batch in datagen.flow(generated_data, Y_gen, batch_size=len(class_indexes)):
#                full.fit(x = x_batch, y = y_batch, epochs = 1, class_weight = weights)
#                break
                
        if batches % 50 == 0:
            tst_data = np.zeros((len(valid_set),512,512,4))
            tst_fl = 0
            for tst_j in valid_set:
                tst_green = cv2.imread(path + tst_j + '_green.png', cv2.IMREAD_GRAYSCALE)
                tst_red = cv2.imread(path + tst_j + '_red.png', cv2.IMREAD_GRAYSCALE)
                tst_blue = cv2.imread(path + tst_j + '_blue.png', cv2.IMREAD_GRAYSCALE)
                tst_yellow = cv2.imread(path + tst_j + '_yellow.png', cv2.IMREAD_GRAYSCALE)
                
                tst_tmp = np.stack((tst_green,tst_red,tst_blue,tst_yellow), axis = -1)/255.
                
                tst_data[tst_fl,:,:,:] = tst_tmp
                tst_fl += 1
            preds = full.predict(x = tst_data)
            
            prnt = preds
            prnt[prnt >= 0.1] = 1
            prnt[prnt < 0.1] = 0
            
            print('-----------------------------------------------')         
            print(prnt[0])
            print(valid_labels[0])
            print('-----------------------------------------------')
            print(prnt[1])
            print(valid_labels[1])
            print('-----------------------------------------------')
            print(prnt[2])
            print(valid_labels[2])
            print('-----------------------------------------------')

            print()
            print('-----------------------------------------------')
            acc_sc(y_true = valid_labels, y_pred = preds)
            IoU(y_true = valid_labels, y_pred = preds)
            f1_sc(y_true = valid_labels, y_pred = preds)
            print('-----------------------------------------------')
            print()
            time.sleep(10)
            
full.save_weights('model.h5')

"""
test = os.listdir("../input/test")
test_path = '../input/test/'

test_files = [x.split('_')[0] for x in test if x.split('_')[1] == 'green.png']

test_data = np.zeros((1,512,512,4))
output = pd.DataFrame(columns=['Id','Predicted'])

full.save_weights('my_model_weights')

for i,file in enumerate(test_files):
    print('Predicting file: ',i)
    test_green = cv2.imread(test_path + file + '_green.png', cv2.IMREAD_GRAYSCALE)
    test_red = cv2.imread(test_path + file + '_red.png', cv2.IMREAD_GRAYSCALE)
    test_blue = cv2.imread(test_path + file + '_blue.png', cv2.IMREAD_GRAYSCALE)
    test_yellow = cv2.imread(test_path + file + '_yellow.png', cv2.IMREAD_GRAYSCALE)
    
    test_tmp = np.stack((test_green,test_red,test_blue,test_yellow), axis = -1)/255.
            
    test_data[0,:,:,:] = test_tmp
    preds = full.predict(test_data.reshape(1,512,512,4))
    preds[preds>=0.1] = 1
    preds[preds<0.1] = 0
    
    pred = np.nonzero(preds)[1]
    if pred.size == 0:
        pred = np.array([0])
    
    predicted = ''
    for pr in pred:
        predicted = predicted + str(pr) + ' '
    predicted = pd.Series(predicted[:-1])

    output = output.append({'Id':file,'Predicted':predicted[0]}, ignore_index = True).copy()
    print(predicted)
output.to_csv('protein_classification.csv', header=True, index=False)
Y_train.Target.unique
"""