#!/usr/bin/env python
# coding: utf-8



import numpy as np
import matplotlib.pyplot as plt
import pandas as pd
import random
from itertools import combinations as cb


from numpy import load
import numpy as np
dict_data = load('x_data.npz')
# extract the first array
x_data_whole=dict_data['arr_0']

dict_data = load('y_data.npz')
y_data_whole=dict_data['arr_0']


dict_data = load('snow_max.npz')
snow_max=dict_data['arr_0']
z_data_whole = y_data_whole*snow_max.tolist()
z_data_whole[z_data_whole<=2]=0
z_data_whole[z_data_whole>2]=1



x_data_whole = np.float32(x_data_whole)
y_data_whole = np.float32(y_data_whole)
z_data_whole = np.float32(z_data_whole)



def data_augmentation(X_slice, Y_slice,Z_slice):
    dice = random.randint(1,4)
    if dice == 1:
        ## updown 
        X_augmented = X_slice
        Y_augmented = Y_slice
        Z_augmented = Z_slice
    if dice == 2:
        ## left right 
        X_augmented = X_slice[:,::-1,:]
        Y_augmented = Y_slice[:,::-1]
        Z_augmented = Z_slice[:,::-1]
    if dice == 3:
        ## updown and left right 
        X_temp = X_slice[::-1,:,:]
        Y_temp = Y_slice[::-1,:]
        Z_temp = Z_slice[::-1,:]
        X_augmented = X_temp[:,::-1,:]
        Y_augmented = Y_temp[:,::-1]
        Z_augmented = Z_temp[:,::-1]       
    if dice == 4:
        X_augmented = X_slice
        Y_augmented = Y_slice
        Z_augmented = Z_slice        
    return X_augmented,Y_augmented,Z_augmented,dice



X_augmented = []
Y_augmented = []
dicexy =[]
Z_augmented = []
random.seed(10)
for j in range(y_data_whole.shape[0]):
    xx,yy,zz,dice =  data_augmentation(x_data_whole[j,:,:,:],y_data_whole[j,:,:],z_data_whole[j,:,:])
    X_augmented.append(xx)
    Y_augmented.append(yy)
    Z_augmented.append(zz)
    dicexy.append(dice)

x_data = np.stack(X_augmented,axis=0)
y_data = np.stack(Y_augmented,axis=0)
z_data = np.stack(Z_augmented,axis=0)

dicexy = np.array(dicexy)
print(x_data.shape)
print(y_data.shape)
print(dicexy.shape)
ik = 3106
part1=[1] * ik
part2=[2] * ik
part3=[3] * ik
partall = part1 + part2 + part3
partall = np.array(partall)



from tensorflow.keras.models import *
from tensorflow.keras.layers import *
def snownet(input_shape=(99, 99, 17)):
    inputs = Input(input_shape)
    x = Conv2D(32,kernel_size=1, padding='same')(inputs)
    x = BatchNormalization()(x) # size: 14*14
    x = Activation('elu')(x)
    x = Conv2D(64,kernel_size=1, padding='same')(x)
    x = BatchNormalization()(x)
    x = Activation('elu')(x)
    x = Conv2D(128,kernel_size=1, padding='same')(x)
    x = BatchNormalization()(x)
    x = Activation('elu')(x)
    x = Conv2D(128,kernel_size=1, padding='same')(x)
    x = BatchNormalization()(x)
    x = Activation('elu')(x)
    x = Conv2D(64,kernel_size=1, padding='same')(x)
    x = BatchNormalization()(x)
    x = Activation('elu')(x)
    x = Conv2D(32,kernel_size=1, padding='same')(x)
    x = BatchNormalization()(x)
    x = Activation('elu')(x)
    output1 = Conv2D(1, 1, activation='linear')(x)
    output2 = Conv2D(1, 1, activation='linear')(x)
    model = Model(inputs=inputs, outputs=[output1,output2])
    model.summary()
    return model



import datetime
import tensorflow as tf
import numpy as np
from tensorflow.keras.callbacks import ModelCheckpoint

def nan_kge(y_actual, y_predicted):
    y_predicted=tf.squeeze(y_predicted)
    y_actual=tf.squeeze(y_actual)    
    is_finites=tf.math.logical_and(tf.math.is_finite(y_actual), tf.math.is_finite(y_predicted))
    sample_number = tf.math.reduce_sum(tf.cast(is_finites,tf.float32))
    is_nans=tf.math.logical_or(tf.math.is_nan(y_actual), tf.math.is_nan(y_predicted))
    y_actual_cleaned = tf.where(is_nans,tf.zeros_like(y_actual),y_actual)
    y_predicted_cleaned = tf.where(is_nans,tf.zeros_like(y_actual),y_predicted)
    actual_mean = tf.math.divide_no_nan(tf.math.reduce_sum(y_actual_cleaned),sample_number)
    predicted_mean = tf.math.divide_no_nan(tf.math.reduce_sum(y_predicted_cleaned),sample_number)
    predicted_temp1_std = tf.where(is_nans,tf.zeros_like(y_actual),tf.math.square(tf.math.subtract(y_predicted,predicted_mean)))
    predicted_std = tf.math.sqrt(tf.math.divide_no_nan(tf.math.reduce_sum(predicted_temp1_std),sample_number))
    actual_temp1_std = tf.where(is_nans,tf.zeros_like(y_actual),tf.math.square(tf.math.subtract(y_actual,actual_mean)))
    actual_std = tf.math.sqrt(tf.math.divide_no_nan(tf.math.reduce_sum(actual_temp1_std),sample_number))
    diff = tf.where(is_nans,tf.zeros_like(y_actual),tf.math.abs(tf.math.subtract(y_actual,y_predicted)))
    errorx = tf.where(is_nans,tf.zeros_like(y_actual),tf.math.subtract(y_predicted,y_actual))
    mae = tf.math.divide_no_nan(tf.math.reduce_sum(diff),sample_number)
    meandiff = tf.math.abs(tf.math.subtract(actual_mean,predicted_mean))
    stddiff = tf.math.abs(tf.math.subtract(actual_std,predicted_std))
    kgeloss = tf.cast(tf.math.add_n([mae,meandiff,stddiff]),tf.float32)
    logcosh=tf.math.reduce_sum(errorx + tf.math.softplus(-2. * errorx) - tf.cast(tf.math.log(2.), errorx.dtype))
    logcoshloss = tf.math.divide_no_nan(logcosh,sample_number)
    return tf.cast(logcoshloss,tf.float32)

def nan_entropy(y_actual, y_predicted):
    y_predicted=tf.cast(tf.squeeze(y_predicted),tf.float64)
    y_actual=tf.cast(tf.squeeze(y_actual),tf.float64)  
    preds = tf.reshape(y_predicted, [-1])
    orig = tf.reshape(y_actual,[-1])
    is_nans=tf.math.logical_or(tf.math.is_nan(orig), tf.math.is_nan(preds))
    y_actual_cleaned = tf.where(is_nans,tf.zeros_like(orig),orig)
    y_predicted_cleaned = tf.where(is_nans,tf.zeros_like(preds),preds)
    const = tf.constant([1],dtype=tf.double)
    constneg = tf.constant([-1],dtype=tf.double)
    t_loss1 = tf.math.subtract(tf.math.maximum(y_predicted_cleaned,0),tf.math.multiply(y_actual_cleaned, y_predicted_cleaned))
    t_loss2 = tf.math.log(tf.math.add(const,tf.math.exp(tf.math.multiply(constneg, tf.math.abs(y_predicted_cleaned)))))
    t_loss = tf.math.add(t_loss1,t_loss2)
    t_final = tf.where(is_nans,tf.zeros_like(t_loss),t_loss)
    t_error = tf.cast(tf.math.reduce_mean(t_final),tf.float32)
    return t_error






import random
random.seed(10)
idxlist = random.sample(range(0, x_data.shape[0]), x_data.shape[0])
five_split = np.array_split(idxlist, 5)
five_split = np.array(five_split,dtype=object)
 
# Get all combinations of [1, 2, 3]
# and length 2
comb = cb([0,1,2,3,4], 4)
combines = list(comb)
testgplist = []
# Print the obtained permutations
for i in range(len(combines)):
    train_list = np.concatenate(list(five_split[np.array(combines[i])]))
#    x_train = x_data[trainidx,:,:,:]
#    y_train = y_data[trainidx,:,:]
    nontraingp = np.setxor1d(combines[i], [0,1,2,3,4])
    test_list = five_split[nontraingp[0]][0:len(five_split[nontraingp[0]])//2]
    valid_list = five_split[nontraingp[0]][len(five_split[nontraingp[0]])//2 : len(five_split[nontraingp[0]])]
    x_train = x_data[train_list,:,:,:]
    y_train = y_data[train_list,:,:]
    z_train = z_data[train_list,:,:]
    x_valid = x_data[valid_list,:,:,:]
    y_valid = y_data[valid_list,:,:]
    z_valid = z_data[valid_list,:,:]
    x_test = x_data[test_list,:,:,:]
    y_test = y_data[test_list,:,:]
    z_test = z_data[test_list,:,:]    
    filepath = "/adapt/nobackup/people/gkonapal/snowmodel_best_multi_output_ghcnpoints"+ str(i) + "fold" +'.hdf5'

    checkpoint = ModelCheckpoint(filepath=filepath, 
                                 monitor='val_loss',
                                 verbose=1, 
                                 save_best_only=True,
                                 mode='min')
    callbacks = [checkpoint]

    model = snownet()
    model.compile(optimizer=tf.keras.optimizers.Adam(),loss=[nan_kge,nan_entropy])
    model.fit(x_train, [y_train,z_train], batch_size=64, epochs=200, validation_data=(x_valid, [y_valid,z_valid]),
              callbacks=callbacks,verbose=2)        
    
    


    

