# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-11-15T11:41:10.025639Z","iopub.execute_input":"2021-11-15T11:41:10.025943Z","iopub.status.idle":"2021-11-15T11:41:10.034699Z","shell.execute_reply.started":"2021-11-15T11:41:10.025910Z","shell.execute_reply":"2021-11-15T11:41:10.033553Z"}}
import matplotlib.pyplot as plt
import plotly.express as px
import pandas as pd
import numpy as np
import time
import os

import seaborn as sns
sns.set_context("talk", font_scale=1.4)
sns.set_style("whitegrid")

from pathlib import Path
from pydicom import filereader
import nibabel as nib

from numpy.random import seed


import math
import numpy as np
from skimage.transform import resize
import math


from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import make_scorer, auc


import os
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'  # 3 = INFO, WARNING, and ERROR messages are not printed
import tensorflow as tf
seed(1)
tf.random.set_seed(2)
from tensorflow.keras.layers import Concatenate, Dense, Conv3D, Input, concatenate, Reshape, Dropout, Flatten, BatchNormalization, LeakyReLU, MaxPool3D
from tensorflow.keras.models import Model
import imgaug.augmenters as iaa

# %% [markdown]
# ## Preprocessing

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-11-15T11:41:10.321716Z","iopub.execute_input":"2021-11-15T11:41:10.322010Z","iopub.status.idle":"2021-11-15T11:41:10.329394Z","shell.execute_reply.started":"2021-11-15T11:41:10.321977Z","shell.execute_reply":"2021-11-15T11:41:10.328292Z"}}
def read_meta():
    folder_struct = []
    for path in (data_path_nifti).rglob('*.nii.gz'):
        
        rmi_type = path.name.split('.')[0].split('_')[1]
        filename = path.name
        type_data = str(path).split('/')[4]
        patient_id = str(path).split('/')[5]
        
                
        folder_struct.append([type_data, patient_id, rmi_type, str(path)])
                             
    image_meta = pd.DataFrame(folder_struct, columns=["type", "id", "mri", "path"])
    
    labels = pd.read_csv(data_path/"train_labels.csv")
    labels['id'] = labels["BraTS21ID"].apply(lambda num: f"{num:05d}")
    labels['label'] = labels['MGMT_value']
    labels.drop(columns=["BraTS21ID", 'MGMT_value'], inplace=True)
    image_meta = pd.merge(image_meta, labels, on = "id", how="left")
    
    return image_meta

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:42:51.576476Z","iopub.execute_input":"2021-11-15T11:42:51.576765Z","iopub.status.idle":"2021-11-15T11:42:51.582631Z","shell.execute_reply.started":"2021-11-15T11:42:51.576729Z","shell.execute_reply":"2021-11-15T11:42:51.581712Z"}}
def get_img_array(image_path):
    img = nib.load(str(image_path))
    return img.get_fdata()
def plot_image_slice(patient_id, rmi_type, image_slice):
    arr = get_img_array(image_meta_tr[(image_meta_tr["id"] == patient_id) &  (image_meta_tr["mri"] == rmi_type)
                                     ]['path'].values[0])
    fig, ax = plt.subplots()
    ax.imshow(arr[:, :, image_slice])
    ax.set_axis_off()

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:42:52.561531Z","iopub.execute_input":"2021-11-15T11:42:52.562064Z","iopub.status.idle":"2021-11-15T11:42:52.581172Z","shell.execute_reply.started":"2021-11-15T11:42:52.562021Z","shell.execute_reply":"2021-11-15T11:42:52.580368Z"}}
def timeit(method):
    def timed(*args, **kw):
        ts = time.time()
        result = method(*args, **kw)
        te = time.time()
        print(method.__name__, "{:.1f}".format(te - ts))        
        return result    
    return timed

def cacheimg(method):
    def cache(*args, **kw):
        
        tmp_save_path = kw['path'].replace("../input/", "/tmp/") + ".npy"
        if not os.path.isfile(tmp_save_path):
            result = method(*args, **kw)
            os.makedirs(os.path.dirname(tmp_save_path))
            np.save(tmp_save_path, result)
        else:
            result = np.load(tmp_save_path)
        
        
        return result    
    return cache

class CustomDataGen(tf.keras.utils.Sequence):
    
    def __init__(self, df, batch_size, input_size=(224, 224, 3), shuffle=True,
                scaling=None, augment_params = {}):
        
        self.df = df.copy()
        self.batch_size = batch_size
        self.input_size = input_size
        self.shuffle = shuffle
        self.n = len(self.df)
        self.augmentations = None
        self.scaling = scaling
        self.augment_params = augment_params
        
        
    def on_epoch_end(self):
        
        if self.shuffle:
            self.df = self.df.sample(frac=1).reset_index(drop=True)
    
    #@timeit
    def _reduce_image_slices(self, image):
        """ reduce 3rd dim of image by simply taking slices at equal tensor dim intervals
        does not consider spatial distances! """
        image_zero_dropped = np.compress((image!=0).sum(axis=(0,1)), image, axis=2)
        self.input_size[2] # smaller equal to image_zero_dropped.shape[2] 
        slice_interval = math.ceil(image_zero_dropped.shape[2] / self.input_size[2])
        
        return image_zero_dropped[:, :, ::slice_interval]
        
    #@timeit
    def _padd_z_dim(self, image):
        """pads z dimension to self.input_size[2] length with 0s. Assumes size not longer than input size"""
        image_resize = resize(image, (self.input_size[0], self.input_size[1], image.shape[2]))
        # pad in z-direction (slices) with black background
        to_pad_z = self.input_size[2] - image.shape[2]
        npad = ((0, 0), (0, 0), (0, to_pad_z))        
        image_padded = np.pad(image_resize, pad_width=npad, mode='constant', constant_values=0)
        #print('padded image to ', image_padded.shape)
        return image_padded
    
    def _scale_image_min_max(self, image):
        """Normalize/scale 3D image by scaling between min 0 and max provided via self.scaling"""
        
        #mean, std = self.scaling
        #image -= mean
        #image /= std
        image_norm = image/self.scaling
        
        return image_norm
    
    #@timeit
    @cacheimg
    def _preprocess_image(self, path):

        image_loaded = self._load_image(path)
        image_reduced = self._reduce_image_slices(image_loaded)
        image_resized = self._padd_z_dim(image_reduced)
        image_scaled = self._scale_image_min_max(image_resized) if self.scaling else image_resized
        image_augm = self._augment_image(image_scaled)
        
        return image_scaled
    
    def __get_data(self, batches):
        # Generates data containing batch_size samples
        
        paths = batches["path"].values
        image_batch = np.asarray([self._preprocess_image(path=path) 
                                  for path in paths])
        #print('final batch dim ', image_batch.shape)
        labels = tf.keras.utils.to_categorical(batches['label'].values, num_classes=2)
        
        return image_batch, labels
    
    def __getitem__(self, idx):
        
        batches = self.df[idx * self.batch_size:(idx + 1) * self.batch_size]
        X, y = self.__get_data(batches)        
        return X, y
    
    def __len__(self):
        return math.ceil(self.n / self.batch_size)
    
    #@timeit
    def _load_image(self, image_path):
        img = nib.load(str(image_path))
        return img.get_fdata()
    
    def _augment_image(self, image):

        if len(self.augment_params) > 0:
            augmentation_steps = []

            if "gblur" in self.augment_params:
                augmentation_steps.append(iaa.GaussianBlur(sigma=self.augment_params["gblur"]))

            seq = iaa.Sequential(augmentation_steps)

            return seq(images=image)
        else:
            return image

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:43:24.197201Z","iopub.execute_input":"2021-11-15T11:43:24.197474Z","iopub.status.idle":"2021-11-15T11:43:24.203069Z","shell.execute_reply.started":"2021-11-15T11:43:24.197443Z","shell.execute_reply":"2021-11-15T11:43:24.202335Z"}}
max_img = 9335.21021664294
print(f"normalization factor (max_img): {max_img}") 


def prepare_gens(input_img_dim, batch_size, train, val):

    data_train_gen = CustomDataGen(train, batch_size=batch_size, input_size=input_img_dim, shuffle=True,
                                  scaling=max_img, augment_params={"gblur": (0, 2.0)})
    data_val_gen = CustomDataGen(val, batch_size=batch_size, input_size=input_img_dim, shuffle=False,
                                scaling=max_img)
    
    return data_train_gen, data_val_gen

# %% [markdown]
# # Neural Networks 

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:43:52.616765Z","iopub.execute_input":"2021-11-15T11:43:52.617055Z","iopub.status.idle":"2021-11-15T11:43:52.627345Z","shell.execute_reply.started":"2021-11-15T11:43:52.617025Z","shell.execute_reply":"2021-11-15T11:43:52.626258Z"}}
def conv_block(inp, filters):
    conv_out = Conv3D(filters=filters, kernel_size=3, strides=(1, 1, 1), activation="relu")(inp)
    return conv_out

    
def pool_block(inp):
    return MaxPool3D(pool_size=(3, 3, 3))(inp)


def dense_block(inp, units):
    dense_out =  Dense(units=units, activation="relu")(inp)
    return dense_out

def model_hu20_seq3(input_shape, dropoutrate=0., lr=1e-4):
    
    inp = Input(shape=input_shape)

    conv1_out = conv_block(inp, filters=16)
    pool1_out = pool_block(conv1_out)
    
    conv2_out = conv_block(pool1_out, filters=16)
    pool2_out = pool_block(conv2_out)  
    
    conv3_out = conv_block(pool2_out, filters=6)
    pool3_out = pool_block(conv3_out)  
    
    flattened_out = Flatten()(pool3_out)
    
    dense1_out = dense_block(flattened_out, units=128)
    drop_out = Dropout(dropoutrate)(dense1_out)
    dense2_out = Dense(units=2, activation="sigmoid")(drop_out)

    model = Model(inputs=[inp], outputs=dense2_out)
    #model.summary(line_length=150)

    
    auc = tf.keras.metrics.AUC()
    
    model.compile(
        optimizer=tf.optimizers.Adam(learning_rate=lr),
        loss=tf.losses.BinaryCrossentropy(),
        metrics=[auc, 
                 tf.keras.metrics.BinaryAccuracy(),
                 tf.keras.metrics.Precision(),
                 tf.keras.metrics.Recall()
                ],
    )
    return model

# %% [markdown]
# # V. Hyperparam Optimization Methods

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:44:26.727286Z","iopub.execute_input":"2021-11-15T11:44:26.728013Z","iopub.status.idle":"2021-11-15T11:44:27.408442Z","shell.execute_reply.started":"2021-11-15T11:44:26.727977Z","shell.execute_reply":"2021-11-15T11:44:27.407628Z"}}
import optuna
#from optuna.integration import TFKerasPruningCallback
from optuna.trial import TrialState
from optuna.integration import SkoptSampler

from optuna.samplers import RandomSampler
from optuna.samplers import TPESampler # Tree Parzen Estimator (TPE)
from optuna.integration import TFKerasPruningCallback

# %% [code] {"execution":{"iopub.status.busy":"2021-11-15T11:44:38.207475Z","iopub.execute_input":"2021-11-15T11:44:38.208193Z","iopub.status.idle":"2021-11-15T11:44:38.410395Z","shell.execute_reply.started":"2021-11-15T11:44:38.208154Z","shell.execute_reply":"2021-11-15T11:44:38.409766Z"}}
import os
def save_trial_results(df):
    # if file does not exist write header 
    if not os.path.isfile('best_vals.csv'):
       df.to_csv('best_vals.csv')
    else:
       df.to_csv('best_vals.csv', mode='a', header=False)
        
def plot_stats_lines(results, var='loss', var_val='val_loss'):
    fig = px.line(data_frame=results.groupby("trial").mean().reset_index(),
               x='trial', 
               y=var,
            error_y = results.groupby("trial").std().reset_index()[var],
                )
    
    if var_val is not None:
    
        fig.add_traces(list(px.line(data_frame=results.groupby("trial").mean().reset_index(),
                   x='trial', 
                   y=var_val,
                    error_y = results.groupby("trial").std().reset_index()[var_val],  
                    ).select_traces()))
        fig.data[1].showlegend = True
        fig.data[1].line.color = "red"
        fig.data[1].name = var_val  
        
    fig.data[0].name = var
    fig.data[0].showlegend = True
    fig.show()
        
        
def show_best_vals():

    results_best_epochs = pd.read_csv("best_vals.csv")
    
    plot_stats_lines(results_best_epochs, var='loss', var_val='val_loss')
    plot_stats_lines(results_best_epochs, var='auc', var_val='val_auc')
    plot_stats_lines(results_best_epochs, var='epochs', var_val=None)
    plot_stats_lines(results_best_epochs, var='precision', var_val='val_precision')
    plot_stats_lines(results_best_epochs, var='recall', var_val='val_recall')
    
    
def show_result(study):

    pruned_trials = study.get_trials(deepcopy=False, states=[TrialState.PRUNED])
    complete_trials = study.get_trials(deepcopy=False, states=[TrialState.COMPLETE])

    print("Study statistics: ")
    print("  Number of finished trials: ", len(study.trials))
    print("  Number of pruned trials: ", len(pruned_trials))
    print("  Number of complete trials: ", len(complete_trials))

    print("Best trial:")
    trial = study.best_trial

    print("  Value: ", trial.value)

    print("  Params: ")
    for key, value in trial.params.items():
        print("    {}: {}".format(key, value))

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2021-11-15T11:44:50.561977Z","iopub.execute_input":"2021-11-15T11:44:50.562693Z","iopub.status.idle":"2021-11-15T11:44:50.572643Z","shell.execute_reply.started":"2021-11-15T11:44:50.562646Z","shell.execute_reply":"2021-11-15T11:44:50.571834Z"}}
def objective_cv(trial):
    
    # hyperparams
    lr = trial.suggest_float("lr", 1e-4, 1e-3, log=True)
    dropout = trial.suggest_float("dropout", 0.0, 0.8)
    input_img_dim = (61, 73, 61)
    batch_size = 32
    
    kfold = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
    valid_auc_folds = []
    
    for fold_idx, (train_idx, val_idx) in enumerate(kfold.split(image_meta_tr, image_meta_tr['label'])):
        
        print("fold ", fold_idx)
        train = image_meta_tr.iloc[train_idx]
        val = image_meta_tr.iloc[val_idx]
    
        # Clear clutter from previous TensorFlow graphs.
        tf.keras.backend.clear_session()

        # load data generators
        data_train_gen, data_val_gen = prepare_gens(input_img_dim, batch_size, train, val)

        # initialize model based on model hyperparams
        model_init = model_hu20_seq3(input_img_dim + (1, ), dropoutrate=dropout, lr=lr)

        # callbacks definition
        early_stopping = tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=10)
        #pruning_trials = TFKerasPruningCallback(trial, "val_auc")

        history = model_init.fit(data_train_gen,
                  validation_data=data_val_gen,
                  epochs=70,
                  use_multiprocessing=False,
                  workers=1,
                  steps_per_epoch=train.shape[0]//batch_size,
                  validation_steps=val.shape[0]//batch_size,
                  verbose = 0,
                  callbacks=[early_stopping])


        results = pd.DataFrame.from_dict(history.history)
        results['epochs'] = results.index + 1
        best_auc = results['val_auc'].max()
        
        results['fold'] = fold_idx + 1
        results['trial'] = trial.number # get trial number
        
        best_epoch_vals = results[results['val_auc'] == best_auc]
        save_trial_results(best_epoch_vals)
        
        
        valid_auc_folds.append(best_auc)
    
    avg_best_auc = np.mean(valid_auc_folds)
    
    return avg_best_auc

# %% [code] {"jupyter":{"outputs_hidden":false}}


# %% [code] {"jupyter":{"outputs_hidden":false}}


# %% [code] {"jupyter":{"outputs_hidden":false}}


# %% [code] {"jupyter":{"outputs_hidden":false}}
