# %% [markdown]
# Hello Fellow Kagglers,
# 
# This notebook demonstrates the inference process using an EfficientNet B8 based UNet model trained on a TPU.
# 
# [Preprocessing Notebook](https://www.kaggle.com/code/markwijkhuizen/hubmap-patched-tfrecord-generation-visualization)
# 
# [Training Notebook](https://www.kaggle.com/code/masterray/hubmap-training-tf-tpu-efficientnet-b8-640-640)

# %% [markdown]
# Copied from: https://www.kaggle.com/code/markwijkhuizen/hubmap-inference-tf-tpu-efficientnet-b7-640x640

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:07.797130Z","iopub.execute_input":"2022-09-22T21:27:07.797433Z","iopub.status.idle":"2022-09-22T21:27:07.802989Z","shell.execute_reply.started":"2022-09-22T21:27:07.797401Z","shell.execute_reply":"2022-09-22T21:27:07.801961Z"}}
# Import EfficientNet models with intermediate endpoints
import sys
sys.path.append('../input/efficientnetv2-head-1x1-endpoint-v2/')
sys.path.append('../input/efficientnetv2-head-1x1-endpoint-v2/efficientnetv2/')

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:07.804914Z","iopub.execute_input":"2022-09-22T21:27:07.805453Z","iopub.status.idle":"2022-09-22T21:27:07.817526Z","shell.execute_reply.started":"2022-09-22T21:27:07.805414Z","shell.execute_reply":"2022-09-22T21:27:07.816571Z"}}
import numpy as np
import pandas as pd
import tensorflow as tf
import tensorflow.keras.backend as K
import tensorflow_addons as tfa
import matplotlib.pyplot as plt

from tensorflow.keras.mixed_precision import experimental as mixed_precision
from kaggle_datasets import KaggleDatasets
from tqdm.notebook import tqdm
#from tqdm.auto import tqdm
from multiprocessing import cpu_count
from sklearn import metrics
from sklearn.model_selection import KFold

from collections import defaultdict

import effnetv2_model
import tifffile
import re
import os
import io
import time
import pickle
import math
import random
import sys
import cv2
import gc

print(f'tensorflow version: {tf.__version__}')
print(f'tensorflow keras version: {tf.keras.__version__}')
print(f'python version: P{sys.version}')

# %% [markdown]
# # Seed

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:07.819827Z","iopub.execute_input":"2022-09-22T21:27:07.820457Z","iopub.status.idle":"2022-09-22T21:27:08.071697Z","shell.execute_reply.started":"2022-09-22T21:27:07.820422Z","shell.execute_reply":"2022-09-22T21:27:08.070786Z"}}
# Seed all random number generators
def seed_everything(seed):
    os.environ['PYTHONHASHSEED'] = str(seed)
    random.seed(seed)
    np.random.seed(seed)
    tf.random.set_seed(seed)
    
SEED = 42
seed_everything(SEED)

# %% [markdown]
# # Setting

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.082807Z","iopub.execute_input":"2022-09-22T21:27:08.083545Z","iopub.status.idle":"2022-09-22T21:27:08.092031Z","shell.execute_reply.started":"2022-09-22T21:27:08.083499Z","shell.execute_reply":"2022-09-22T21:27:08.091129Z"}}

# Threshold to classify a pixel as mask
#THRESHOLD = 0.390

# base experimento, 0.390 ALL 0.68 #B7
# Primer experimento, 0.1 ALL 0.74
# Segundo experimento, Hubmap 1, HPA 0.1 -0.20
# Tercer expermiento, Hubmap 0.1, HPA 0.3, lung 0.1   0.75 LB 154
# cUARTO expermiento, Hubmap 0.1, HPA 0.4, lung 0.1   0.75 LB 152  150
# Quinto expermiento, Hubmap 0.1, HPA 0.4, lung 0.1   0.75 LB 148 
# Quinto expermiento, Hubmap 0.4, HPA X, lung 0.1   0.75 LB 148 Best Score 

ORGAN_THRESHOLD = {
    
    'Hubmap': {
        'kidney'        : 0.18, #22
        'prostate'      : 0.14, #*
        'largeintestine': 0.16, #*
        'spleen'        : 0.14, #*
        'lung'          : 0.03, #*
    },
    'HPA': {
        'kidney'        : 0.63, #0.63,    
        'prostate'      : 0.56, #0.56,
        'largeintestine': 0.52, #0.52,
        'spleen'        : 0.41, #0.41,
        'lung'          : 0.03, #0.03,
    },
}

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.094488Z","iopub.execute_input":"2022-09-22T21:27:08.095552Z","iopub.status.idle":"2022-09-22T21:27:08.102492Z","shell.execute_reply.started":"2022-09-22T21:27:08.095513Z","shell.execute_reply":"2022-09-22T21:27:08.101714Z"}}
WEIGHTS_MODELS = {
    # kidney
    "kidney": {"w1": 0.3, "w2": 0.1, "w3": 0.2, "w4": 0.4},
    
    # Prostate
    "prostate": {"w1": 0.3, "w2": 0.1, "w3": 0.2, "w4": 0.4},
    
    # largeintestine
    "largeintestine": {"w1": 0.3, "w2": 0.1, "w3": 0.3, "w4": 0.3},
    
    # Spleen
    "spleen": {"w1": 0.1, "w2": 0.1, "w3": 0.4, "w4": 0.4},
    
    # Lung
    "lung": {"w1":0.4, "w2":0.1, "w3":0.1, "w4":0.4}
}

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.103869Z","iopub.execute_input":"2022-09-22T21:27:08.104201Z","iopub.status.idle":"2022-09-22T21:27:08.116210Z","shell.execute_reply.started":"2022-09-22T21:27:08.104167Z","shell.execute_reply":"2022-09-22T21:27:08.115209Z"}}
### Show model
SHOW_MODEL = False
CHECK_SANITY = False

MODEL_PREDICTION_MODE = 'single'
#MODEL_PREDICTION_MODE = 'ensamble'
#MODEL_PREDICTION_MODE = 'mix'


DEBUG = False
IS_TPU = True

# Image dimensions
IMG_SIZE = 640
PATCH_SIZE = 640
#IMG_SIZE = 768
#PATCH_SIZE = 768
N_CHANNELS = 3
N_PATCHES_PER_IMAGE = (IMG_SIZE // PATCH_SIZE) ** 2

INPUT_SHAPE = (PATCH_SIZE, PATCH_SIZE, N_CHANNELS)

N_FOLDS=4

# EfficientNet version, b0/b1/b2/b3/s/m/l/xl/xxl
EFN_SIZE = 'b8'
LR_MAX = 0.02
EPOCHS = 30
MOMENTUM = 0.00

# Batch size
BATCH_SIZE = 64

# DATASET PATH
if IMG_SIZE == 768:
    DATASET_DIR = '../input/ds-bojack-hubmap-hpa-tfrecords-768'
    FOLDS_MODELS_DIR = '../input/hubmaphpamtftpu680768/768/checkpoints'
else:
    DATASET_DIR = '../input/hubmap-patched-tfrecords-300x300'
    FOLDS_MODELS_DIR = '../input/hubmaphpamtftpu680768/680/checkpoints'

SINGLE_MODELS_DIR = '../input/hubmaphpamtftpu680768/b_model/model_0.h5'
    
print("[PATHS] DATASET_DIR", DATASET_DIR)
print("[PATHS] FOLD_MODELS_DIR", FOLDS_MODELS_DIR)
print("[PATHS] SINGLE_MODELS_DIR", SINGLE_MODELS_DIR)

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.118137Z","iopub.execute_input":"2022-09-22T21:27:08.118492Z","iopub.status.idle":"2022-09-22T21:27:08.146550Z","shell.execute_reply.started":"2022-09-22T21:27:08.118444Z","shell.execute_reply":"2022-09-22T21:27:08.145431Z"}}
# Dataset Mean and Standard Deviation
#MEAN = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/MEAN.npy')
#STD = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/STD.npy')
MEAN = np.load(F'{DATASET_DIR}/MEAN.npy')
STD = np.load(F'{DATASET_DIR}/STD.npy')

print(f'MEAN: {MEAN}, STD: {STD}')

# %% [markdown]
# # Hardware Configuration

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.148214Z","iopub.execute_input":"2022-09-22T21:27:08.148526Z","iopub.status.idle":"2022-09-22T21:27:08.159606Z","shell.execute_reply.started":"2022-09-22T21:27:08.148487Z","shell.execute_reply":"2022-09-22T21:27:08.158700Z"}}
# Detect hardware, return appropriate distribution strategy
try:
    TPU = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.
    print('Running on TPU ', TPU.master())
except ValueError:
    print('Running on GPU')
    TPU = None

if TPU:
    tf.config.experimental_connect_to_cluster(TPU)
    tf.tpu.experimental.initialize_tpu_system(TPU)
    strategy = tf.distribute.experimental.TPUStrategy(TPU)
else:
    strategy = tf.distribute.get_strategy() # default distribution strategy in Tensorflow. Works on CPU and single GPU.

REPLICAS = strategy.num_replicas_in_sync
print(f'REPLICAS: {REPLICAS}')

# %% [markdown]
# # FPN

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.161287Z","iopub.execute_input":"2022-09-22T21:27:08.161624Z","iopub.status.idle":"2022-09-22T21:27:08.173102Z","shell.execute_reply.started":"2022-09-22T21:27:08.161587Z","shell.execute_reply":"2022-09-22T21:27:08.172221Z"}}
def FPN(xs, output_channels, last_layer, debug=False):
    def _conv(x):
        x = tf.keras.layers.ZeroPadding2D(padding=1)(x)
        x = tf.keras.layers.Conv2D(output_channels * 2, 3, padding='SAME', kernel_initializer='he_normal', activation='relu')(x)
        x = tf.keras.layers.BatchNormalization()(x)
        x = tf.keras.layers.ZeroPadding2D(padding=1)(x)
        x = tf.keras.layers.Conv2D(output_channels, 3, padding='SAME', kernel_initializer='he_normal')(x)
        x = tf.image.resize(x, size=target_size, method=tf.image.ResizeMethod.BILINEAR)
        x = tf.nn.relu(x)
        return x

    target_size = last_layer.shape[1:3]
    xs = tf.keras.layers.Concatenate()([_conv(x) for x in xs])
    x = tf.keras.layers.Concatenate()([xs, last_layer])

    if debug:
        return x, xs
    else:
        return x

# %% [markdown]
# # ASPP

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.176434Z","iopub.execute_input":"2022-09-22T21:27:08.176671Z","iopub.status.idle":"2022-09-22T21:27:08.190995Z","shell.execute_reply.started":"2022-09-22T21:27:08.176645Z","shell.execute_reply":"2022-09-22T21:27:08.190209Z"}}
def ASPP(x, mid_c=320, dilations=[1, 2, 3, 4], out_c=640, debug=False):
    def _aspp_module(x, filters, kernel_size, padding, dilation, groups=1):
        x = tf.keras.layers.ZeroPadding2D(padding=padding)(x)
        x = tf.keras.layers.Conv2D(
                filters=filters,
                kernel_size=kernel_size,
                dilation_rate=dilation,
                groups=1 if IS_TPU else groups,
                kernel_initializer='he_uniform',
            )(x)
        x = tf.keras.layers.BatchNormalization()(x)
        x = tf.nn.relu(x)
        
        return x
    
    x0 = tf.math.reduce_max(x, axis=(1,2), keepdims=True)
    x0 = tf.keras.layers.Conv2D(filters=mid_c, kernel_size=1, strides=1, kernel_initializer='he_uniform', use_bias=False)(x0)
    x0 = tf.keras.layers.BatchNormalization(gamma_initializer=tf.constant_initializer(value=0.25))(x0)
    x0 = tf.nn.relu(x0)
                                  
                                  
    xs = (
        [_aspp_module(x, mid_c, 1, padding=0, dilation=1)] +
        [_aspp_module(x, mid_c, 3, padding=d, dilation=d, groups=4) for d in dilations]
    )
    
    x0= tf.image.resize(x0, size=xs[0].shape[1:3])
    x = tf.keras.layers.Concatenate()([x0] + xs)
    x = tf.keras.layers.Conv2D(filters=out_c, kernel_size=1, kernel_initializer='he_uniform', use_bias=False)(x)
    x = tf.keras.layers.BatchNormalization()(x)
    x = tf.nn.relu(x)
                       
    if debug:
        return x, x0, xs
    else:
        return x

# %% [markdown]
# # Upsample

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.194806Z","iopub.execute_input":"2022-09-22T21:27:08.195099Z","iopub.status.idle":"2022-09-22T21:27:08.205628Z","shell.execute_reply.started":"2022-09-22T21:27:08.195070Z","shell.execute_reply":"2022-09-22T21:27:08.204755Z"}}
def PixelShuffle(x, upscale_factor=2):
    _, w, h, c = x.shape
    n = -1

    c_out = c // upscale_factor ** 2
    w_out = w * upscale_factor
    h_out = h * upscale_factor

    x = tf.reshape(x, [-1, upscale_factor, upscale_factor, w, h, c_out])
    x = tf.transpose(x, [0, 3, 1, 4, 2, 5])
    x = tf.reshape(x, [-1, w_out, h_out, c_out])

    return x

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.209182Z","iopub.execute_input":"2022-09-22T21:27:08.209406Z","iopub.status.idle":"2022-09-22T21:27:08.220704Z","shell.execute_reply.started":"2022-09-22T21:27:08.209381Z","shell.execute_reply":"2022-09-22T21:27:08.219929Z"}}
# Inspiration: https://www.tensorflow.org/tutorials/generative/pix2pix#build_an_input_pipeline_with_tfdata
def upsample(x, concat, target_filters, name, conv2dt_kernel_init_max, relu=True, dropout=0, debug=False):
#     x = PixelShuffle(x)

    filters = concat.shape[-1]
    x_up = tf.keras.layers.Conv2DTranspose(
            filters, # Number of Convolutional Filters
            kernel_size=4, # Kernel Size
            strides=2, # Kernel Steps
            padding='SAME', # linear scaling
            name=f'Conv2DTranspose_{name}', # Name of Layer
            kernel_initializer='he_uniform',
            use_bias=False,
        )(x)
    
    concat = tf.keras.layers.BatchNormalization(
        gamma_initializer=tf.constant_initializer(value=0.25),
        name=f'BatchNormalization_{name}'
    )(concat)
    x = tf.keras.layers.Concatenate(name=f'Concatenate_{name}')([x_up, concat])
    x = tf.nn.relu(x)
    
        
    x = tf.keras.layers.Conv2D(target_filters, 3, padding='SAME', kernel_initializer='he_uniform', activation='relu', name=f'Conv2D_1_{name}')(x)
    x = tf.keras.layers.Conv2D(target_filters, 3, padding='SAME', kernel_initializer='he_uniform', name=f'Conv2D_2_{name}')(x)
    
    if relu:
        x = tf.nn.relu(x)
    
    x = tf.keras.layers.Dropout(dropout, name=f'Dropout_{name}')(x)

    if debug:
        return x, x_up, concat
    else:
        return x

# %% [markdown]
# # Model

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.223905Z","iopub.execute_input":"2022-09-22T21:27:08.224261Z","iopub.status.idle":"2022-09-22T21:27:08.241107Z","shell.execute_reply.started":"2022-09-22T21:27:08.224215Z","shell.execute_reply":"2022-09-22T21:27:08.240028Z"}}
def get_model(dropout_decoder=0, dropout_cnn=0, file_path=None, lr=1e-3, eps=1e-7, clipnorm=5.0, wd_coef=1e-2, cnn_trainable=True):
    with strategy.scope():
        # EfficientNetV2 Backbone # 
        cnn = effnetv2_model.get_model(f'efficientnet-{EFN_SIZE}', include_top=False, weights=None, model_config={ 'conv_dropout': dropout_cnn })
        cnn.trainable = cnn_trainable

        # Inputs, note the names are equal to the dictionary keys in the dataset
        image = tf.keras.layers.Input(INPUT_SHAPE, name='image', dtype=tf.float32)
        image_norm = tf.cast(image, tf.float32) / 255
        image_norm = tf.keras.layers.experimental.preprocessing.Normalization(mean=MEAN, variance=STD, dtype=tf.float32)(image_norm)

        embedding, up6, up5, up4, up3, up2, up1 = cnn(image_norm, with_endpoints=True)
        print(f'embedding shape: {embedding.shape} up1 shape: {up1.shape}, up2 shape: {up2.shape}')
        print(f'up3 shape: {up3.shape}, up4 shape: {up4.shape}, up5 shape: {up5.shape}, up6 shape: {up6.shape}')
        
        dec0 = ASPP(up2)
        dec0 = tf.keras.layers.Dropout(0.50)(dec0)

        dec1 = upsample(dec0, up3, up4.shape[-1] * 4, 'upsample1', 0.02, dropout=dropout_decoder)
        dec2 = upsample(dec1, up4, up5.shape[-1] * 2, 'upsample2', 0.02, dropout=dropout_decoder)
        dec3 = upsample(dec2, up5, up6.shape[-1] * 2, 'upsample3', 0.02)
        dec4 = upsample(dec3, up6, 64, 'upsample4', 0.02)
        
        print(f'dec0 shape: {dec0.shape}, dec1 shape: {dec1.shape}, dec2 shape: {dec2.shape}, dec3 shape: {dec3.shape}, dec4 shape: {dec4.shape}')
        
        dec_fpn = FPN([dec0, dec1, dec2, dec3], 32, dec4)
        
        print(f'dec_fpn shape: {dec_fpn.shape}')
        
        # Head
        x = tf.keras.layers.Dropout(0.10)(dec_fpn)
        x = tf.keras.layers.Conv2D(
            filters=1,
            kernel_size=1,
            padding='SAME',
            kernel_initializer=tf.random_normal_initializer(0.00, 0.05),
            activation='sigmoid',
            name='Conv2D_3_head'
        )(x)
        output = tf.image.resize(x, size=[IMG_SIZE, IMG_SIZE], method=tf.image.ResizeMethod.BILINEAR)
        
        model = tf.keras.models.Model(inputs=image, outputs=output)

        if file_path:
            print('Loading pretrained weights...')
            model.load_weights(file_path)
            
        model.trainable = False

        return model

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.242616Z","iopub.execute_input":"2022-09-22T21:27:08.243313Z","iopub.status.idle":"2022-09-22T21:27:08.255540Z","shell.execute_reply.started":"2022-09-22T21:27:08.243278Z","shell.execute_reply":"2022-09-22T21:27:08.254603Z"}}
class WeightedAverageLayer(tf.keras.layers.Layer): 
     
        def __init__(self, w1, w2, w3, w4): 
            super(WeightedAverageLayer, self).__init__() 
            self.w1 = w1 
            self.w2 = w2 
            self.w3 = w3 
            self.w4 = w4 
 
 
        def call(self, inputs): 
            return self.w1 * inputs[0] + self.w2 * inputs[1] + self.w3 * inputs[2] + self.w4 * inputs[3]

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.256952Z","iopub.execute_input":"2022-09-22T21:27:08.257530Z","iopub.status.idle":"2022-09-22T21:27:08.270264Z","shell.execute_reply.started":"2022-09-22T21:27:08.257492Z","shell.execute_reply":"2022-09-22T21:27:08.269358Z"}}
def build_ensamble(models, w1, w2, w3, w4, input_shape=INPUT_SHAPE):
    model_input = tf.keras.layers.Input(input_shape, name='image', dtype=tf.float32)
    model_outputs = [model(model_input) for model in models]
    ensemble_output = WeightedAverageLayer(w1 , w2, w3, w4)(model_outputs)
    ensemble_model = tf.keras.Model(inputs=model_input, outputs=ensemble_output)
    
    return ensemble_model

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.273937Z","iopub.execute_input":"2022-09-22T21:27:08.274197Z","iopub.status.idle":"2022-09-22T21:27:08.282358Z","shell.execute_reply.started":"2022-09-22T21:27:08.274153Z","shell.execute_reply":"2022-09-22T21:27:08.281524Z"}}
def load_model_folds(models_dir, n_folds=N_FOLDS):
    models = []
    for fold in range(n_folds):
        model_path = F'{models_dir}/fold{fold}/best_model_{fold}.h5'
        if not os.path.exists(model_path):
            print("[LOAD-MODELS] DOEN'T FIND MODEL", model_path)
        model = get_model(file_path=model_path)
        models.append(model)
    
    return models

# %% [markdown]
# # Metrics

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.283881Z","iopub.execute_input":"2022-09-22T21:27:08.284619Z","iopub.status.idle":"2022-09-22T21:27:08.293695Z","shell.execute_reply.started":"2022-09-22T21:27:08.284581Z","shell.execute_reply":"2022-09-22T21:27:08.292805Z"}}
def dice_coef(groundtruth_mask, pred_mask):
    intersect = np.sum(pred_mask*groundtruth_mask)
    total_sum = np.sum(pred_mask) + np.sum(groundtruth_mask)
    dice = np.mean(2*intersect/total_sum)
    return round(dice, 3) #round up to 3 decimal places

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.297345Z","iopub.execute_input":"2022-09-22T21:27:08.297581Z","iopub.status.idle":"2022-09-22T21:27:08.304665Z","shell.execute_reply.started":"2022-09-22T21:27:08.297554Z","shell.execute_reply":"2022-09-22T21:27:08.303635Z"}}
def iou(y_true, y_pred):
    intersection = np.count_nonzero(y_true * y_pred)
    union = np.count_nonzero(y_true + y_pred)
    return intersection / union

# %% [markdown]
# # Utility Funtions

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.306244Z","iopub.execute_input":"2022-09-22T21:27:08.307163Z","iopub.status.idle":"2022-09-22T21:27:08.316172Z","shell.execute_reply.started":"2022-09-22T21:27:08.307124Z","shell.execute_reply":"2022-09-22T21:27:08.315383Z"}}
def rle2mask(rle, width, height):
    s = rle.split()
    
    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]
    starts -= 1
    ends = starts + lengths
    
    mask = np.zeros(width * height, dtype=np.uint8)
    for start, end in zip(starts, ends):
        mask[start:end] = 1
    
    return mask.reshape(width, height).T

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.317418Z","iopub.execute_input":"2022-09-22T21:27:08.318281Z","iopub.status.idle":"2022-09-22T21:27:08.326569Z","shell.execute_reply.started":"2022-09-22T21:27:08.318241Z","shell.execute_reply":"2022-09-22T21:27:08.325898Z"}}
def mask2rle(mask):
    mask = mask.T.flatten()
    mask = np.concatenate([[0], mask, [0]])
    runs = np.where(mask[1:] != mask[:-1])[0] + 1
    runs[1::2] -= runs[::2]
    return ' '.join(str(x) for x in runs)

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.328353Z","iopub.execute_input":"2022-09-22T21:27:08.328655Z","iopub.status.idle":"2022-09-22T21:27:08.341108Z","shell.execute_reply.started":"2022-09-22T21:27:08.328617Z","shell.execute_reply":"2022-09-22T21:27:08.340113Z"}}
# Resized a tensor to the specified size
def resize_tensor(tensor, size=IMG_SIZE, dtype=np.uint8):
    return cv2.resize(tensor, [size, size], interpolation=cv2.INTER_CUBIC).astype(dtype)

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.342236Z","iopub.execute_input":"2022-09-22T21:27:08.342516Z","iopub.status.idle":"2022-09-22T21:27:08.352366Z","shell.execute_reply.started":"2022-09-22T21:27:08.342488Z","shell.execute_reply":"2022-09-22T21:27:08.351551Z"}}
# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode
def get_mask(image_id):
    row = train.loc[train['id'] == image_id].squeeze()
    h, w = row[['img_height', 'img_width']]
    mask = np.zeros(shape=[h * w], dtype=np.uint8)
    s = row['rle'].split()
    starts, lengths = [ np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2]) ]
    starts -= 1
    ends = starts + lengths
    for lo, hi in zip(starts, ends):
        mask[lo : hi] = 1
        
    mask = mask.reshape([h, w]).T
        
    mask = resize_tensor(mask)
    
    mask = np.expand_dims(mask, axis=2)
        
    return mask

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.353951Z","iopub.execute_input":"2022-09-22T21:27:08.354530Z","iopub.status.idle":"2022-09-22T21:27:08.362750Z","shell.execute_reply.started":"2022-09-22T21:27:08.354491Z","shell.execute_reply":"2022-09-22T21:27:08.362047Z"}}
# Reads an image and returns the image and original image size
def get_image(image_id, folder, negative=True):
    image = tifffile.imread(f'/kaggle/input/hubmap-organ-segmentation/{folder}_images/{image_id}.tiff')
    if len(image.shape) == 5:
        image = image.squeeze().transpose(1, 2, 0)
    
    # Image Size
    image_size, _, _ = image.shape
    
    # Reverse pixels to make tissue colored and background black
    if negative:
        image = image - image.min()
        image = image / (image.max() - image.min())
        image = image * 255
        image = 255 - image.astype(np.uint8)
        
    # Resize
    image = resize_tensor(image)
    
    return image, image_size

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.364481Z","iopub.execute_input":"2022-09-22T21:27:08.365042Z","iopub.status.idle":"2022-09-22T21:27:08.374383Z","shell.execute_reply.started":"2022-09-22T21:27:08.365007Z","shell.execute_reply":"2022-09-22T21:27:08.373839Z"}}
# extract patches from an image
def extract_patches(image):
    _, _, c = image.shape
    image = tf.expand_dims(image, 0)
    image_patches = tf.image.extract_patches(image, [1,PATCH_SIZE,PATCH_SIZE,1], [1, PATCH_SIZE, PATCH_SIZE, 1], [1, 1, 1, 1], padding='SAME')
    image_patches = tf.reshape(image_patches, [N_PATCHES_PER_IMAGE, PATCH_SIZE, PATCH_SIZE, c])
    image_patches = image_patches.numpy()

    return image_patches

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.376054Z","iopub.execute_input":"2022-09-22T21:27:08.376667Z","iopub.status.idle":"2022-09-22T21:27:08.389422Z","shell.execute_reply.started":"2022-09-22T21:27:08.376628Z","shell.execute_reply":"2022-09-22T21:27:08.388702Z"}}
#https://www.kaggle.com/bguberfain/memory-aware-rle-encoding
#with transposed mask
def rle_encode_less_memory(img):
    #the image should be transposed
    pixels = img.T.flatten()
    
    # This simplified method requires first and last pixel to be zero
    pixels[0] = 0
    pixels[-1] = 0
    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2
    runs[1::2] -= runs[::2]
    
    return ' '.join(str(x) for x in runs)

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.391153Z","iopub.execute_input":"2022-09-22T21:27:08.391746Z","iopub.status.idle":"2022-09-22T21:27:08.401109Z","shell.execute_reply.started":"2022-09-22T21:27:08.391699Z","shell.execute_reply":"2022-09-22T21:27:08.400569Z"}}
# Reconstruct the original image from patches
def merge_patches(patches):
    image = np.zeros(shape=[IMG_SIZE, IMG_SIZE, patches.shape[-1]], dtype=patches.dtype)
    s = int(N_PATCHES_PER_IMAGE ** 0.50)
    for r in range(s):
        for c in range(s):
            start_x = r * PATCH_SIZE
            end_x = (r + 1) * PATCH_SIZE
            start_y = c * PATCH_SIZE
            end_y = (c + 1) * PATCH_SIZE
            image[start_x:end_x, start_y:end_y] = patches[r * s + c]
            
    return image

# %% [markdown]
# # Functions Sanity Check

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.402843Z","iopub.execute_input":"2022-09-22T21:27:08.403390Z","iopub.status.idle":"2022-09-22T21:27:08.413379Z","shell.execute_reply.started":"2022-09-22T21:27:08.403349Z","shell.execute_reply":"2022-09-22T21:27:08.412635Z"}}
def visualize_test_images(image, mask_prob, mask_pred, title="Original"):
    size = 26
    plt.figure(figsize = (size, size* 5))
    
    # Image Original
    plt.subplot(141)
    plt.title(f'Image: {title}')
    plt.imshow(image)
    
    # Mask prob
    plt.subplot(142)
    plt.title(f"Mask Prob")
    plt.imshow(mask_prob)
    
    # Mask pred
    plt.subplot(143)
    plt.title(f"Mask Pred")
    plt.imshow(mask_pred)
    
    # Image + Mask pred
    plt.subplot(144)
    plt.title(f"Image+Mask")
    plt.imshow(image)
    plt.imshow(mask_pred, alpha = 0.5)
    
    plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.415053Z","iopub.execute_input":"2022-09-22T21:27:08.415639Z","iopub.status.idle":"2022-09-22T21:27:08.424120Z","shell.execute_reply.started":"2022-09-22T21:27:08.415533Z","shell.execute_reply":"2022-09-22T21:27:08.423604Z"}}
def visualize_train_images(image, mask, mask_prob, mask_pred, title="Original"):
    size = 26
    plt.figure(figsize = (size, size* 5))
    
    # Image Original
    plt.subplot(151)
    plt.title(f'Image: {title}')
    plt.imshow(image)
    
    # Mask Original
    plt.subplot(152)
    plt.title(f"Mask original")
    plt.imshow(mask)
    
    # Mask prob
    plt.subplot(153)
    plt.title(f"Mask Prob")
    plt.imshow(mask_prob)
    
    # Mask pred
    plt.subplot(154)
    plt.title(f"Mask Pred")
    plt.imshow(mask_pred)
    
    # Image + Mask pred
    plt.subplot(155)
    plt.title(f"Image+Mask")
    plt.imshow(image)
    plt.imshow(mask_pred, alpha = 0.5)
    
    plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.429566Z","iopub.execute_input":"2022-09-22T21:27:08.430083Z","iopub.status.idle":"2022-09-22T21:27:08.442829Z","shell.execute_reply.started":"2022-09-22T21:27:08.430043Z","shell.execute_reply":"2022-09-22T21:27:08.442029Z"}}
def check_train_sanity(models, data, random_sample=False, random_state=42, n_samples=10, show_img=False):
    
    size = len(data)
    n_samples = n_samples if n_samples < size else size
    sample = data.sample(n=n_samples, random_state=random_state) if random_sample else data[:n_samples]
    
    metrics = []
    print("[CHECK-SANITY] SHOW SAMPLES", n_samples)
    for row_idx, row in tqdm(sample.iterrows(), total=n_samples):
        image_id = row['id']
        data_source = row['data_source']
        organ = row['organ']
        rle = row['rle'] 
        height = row['img_height'] 
        width = row['img_width'] 
        
        # Check if dictionary
        model = models[organ] if type(models) is dict else models
        
        # Mask Original
        mask = rle2mask(rle, width, height)
        
        # Preprocess Image
        image, image_size = get_image(image_id, 'train')
        image_patches = extract_patches(image)
        
        # Image Original
        image_original = resize_tensor(image, size=image_size, dtype=np.uint8)

        # Make Prediction
        mask_patches_pred = model.predict(image_patches)
        mask_pred = merge_patches(mask_patches_pred)
        mask_pred_resized = resize_tensor(mask_pred, size=image_size, dtype=np.float32)
        mask_binary = (mask_pred_resized > ORGAN_THRESHOLD[data_source][organ]).astype(np.int8)
        
        dice_metric = dice_coef(mask, mask_binary)
        iou_metric = iou(mask, mask_binary)
        title = F"{image_id} o:{organ} dice:{dice_metric:.3f} iou:{iou_metric:.3f}"
        
        metrics.append((image_id, organ, dice_metric, iou_metric))
        # Visualize
        if show_img:
            visualize_train_images(image_original, mask, mask_pred_resized, mask_binary, title)
    
    metric_df = pd.DataFrame(data=metrics, columns=["id", "organ", "dice", "ioc"])
    return metric_df

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.444795Z","iopub.execute_input":"2022-09-22T21:27:08.445412Z","iopub.status.idle":"2022-09-22T21:27:08.459074Z","shell.execute_reply.started":"2022-09-22T21:27:08.445371Z","shell.execute_reply":"2022-09-22T21:27:08.458296Z"}}
def metrics_threshold_folds(model, train, folds_dict):
    print(train.shape)
    thresholds = np.arange(0, 1.01, 0.01)
    organs = sorted(train['organ'].unique())
    # Iterate folds
    model_metrics = {}
    for fold, val_idxs in zip(folds_dict['folds'], folds_dict['val_idxs']):
        val = train.loc[val_idxs]
        print(fold, len(val_idxs), len(val))
        val_metrics = {
            "dice": { o:defaultdict(list) for o in organs},
            "iou": { o:defaultdict(list) for o in organs}
        }
        for row_idx, row in tqdm(val.iterrows(), total=len(val)):
            image_id = row['id']
            data_source = row['data_source']
            organ = row['organ']
            rle = row['rle'] 
            height = row['img_height'] 
            width = row['img_width'] 
        
            # Mask Original
            mask = rle2mask(rle, width, height)

            # Preprocess Image
            image, image_size = get_image(image_id, 'train')
            image_patches = extract_patches(image)

            # Image Original
            image_original = resize_tensor(image, size=image_size, dtype=np.uint8)

            # Make Prediction
            mask_patches_pred = model.predict(image_patches)
            mask_pred = merge_patches(mask_patches_pred)
            mask_pred_resized = resize_tensor(mask_pred, size=image_size, dtype=np.float32)
            
            # Metrics thresholds
            for threshold in thresholds:
                mask_binary = (mask_pred_resized > threshold).astype(np.int8)
                dice_metric = dice_coef(mask, mask_binary)
                iou_metric = iou(mask, mask_binary)
                val_metrics["dice"][organ][threshold].append(dice_metric)
                val_metrics["iou"][organ][threshold].append(iou_metric)
            
            model_metrics[fold] = val_metrics
    return model_metrics

# %% [markdown]
# # Inference

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.460950Z","iopub.execute_input":"2022-09-22T21:27:08.461485Z","iopub.status.idle":"2022-09-22T21:27:08.598554Z","shell.execute_reply.started":"2022-09-22T21:27:08.461442Z","shell.execute_reply":"2022-09-22T21:27:08.597560Z"}}
# Training DataFrame
train = pd.read_csv('/kaggle/input/hubmap-organ-segmentation/train.csv')

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.600139Z","iopub.execute_input":"2022-09-22T21:27:08.600432Z","iopub.status.idle":"2022-09-22T21:27:08.608599Z","shell.execute_reply.started":"2022-09-22T21:27:08.600395Z","shell.execute_reply":"2022-09-22T21:27:08.607803Z"}}
# Test DataFrame
test = pd.read_csv('/kaggle/input/hubmap-organ-segmentation/test.csv')

# %% [markdown]
# # Load model

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:08.609963Z","iopub.execute_input":"2022-09-22T21:27:08.610204Z","iopub.status.idle":"2022-09-22T21:27:30.061575Z","shell.execute_reply.started":"2022-09-22T21:27:08.610162Z","shell.execute_reply":"2022-09-22T21:27:30.060644Z"}}
# Pretrained File Path: '/kaggle/input/sartorius-training-dataset/model.h5'
if MODEL_PREDICTION_MODE == "single":
    print("[LOAD-MODEL] SINGLE LOADING....")
    model = get_model(file_path=SINGLE_MODELS_DIR)
elif MODEL_PREDICTION_MODE == "ensamble":
    print("[LOAD-MODEL] ENSAMBLE LOADING....")
    models = load_model_folds(models_dir=FOLDS_MODELS_DIR, n_folds=N_FOLDS)
    
    arams = WEIGHTS_MODELS.copy()
    params["models"] = models
    params["input_shape"] = INPUT_SHAPE
    model = build_ensamble(models, w1=0.2, w2=0.2, w3=0.3, w4=0.3, input_shape=INPUT_SHAPE)
else:
    print("[LOAD-MODEL] MIX-MODELS...")
    b_model = get_model(file_path=SINGLE_MODELS_DIR)
    mix_models = load_model_folds(models_dir=FOLDS_MODELS_DIR, n_folds=N_FOLDS)
    mix_models = mix_models[1:] 
    mix_models.append(b_model)
    
    model = {}
    for organ, param in WEIGHTS_MODELS.items():
        print(organ, param)
        mparams = param.copy()
        mparams["models"] = mix_models
        mparams["input_shape"] = INPUT_SHAPE
        organ_model = build_ensamble(**mparams)
        model[organ] = organ_model

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.063352Z","iopub.execute_input":"2022-09-22T21:27:30.063670Z","iopub.status.idle":"2022-09-22T21:27:30.068437Z","shell.execute_reply.started":"2022-09-22T21:27:30.063630Z","shell.execute_reply":"2022-09-22T21:27:30.067367Z"}}
# Plot model summary
if SHOW_MODEL:
    model.summary()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.070189Z","iopub.execute_input":"2022-09-22T21:27:30.070769Z","iopub.status.idle":"2022-09-22T21:27:30.080621Z","shell.execute_reply.started":"2022-09-22T21:27:30.070727Z","shell.execute_reply":"2022-09-22T21:27:30.079657Z"}}
if SHOW_MODEL:
    tf.keras.util.plot_model(model, show_shapes=True, show_dtype=True, show_layer_names=True, expand_nested=False)

# %% [markdown]
# # Sanity Check

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.082068Z","iopub.execute_input":"2022-09-22T21:27:30.082526Z","iopub.status.idle":"2022-09-22T21:27:30.092295Z","shell.execute_reply.started":"2022-09-22T21:27:30.082486Z","shell.execute_reply":"2022-09-22T21:27:30.091448Z"}}
if CHECK_SANITY:
    sample = train[train["organ"] == 'lung']
    p_metric_df = check_train_sanity(model, sample, random_sample=True, random_state=SEED, n_samples=len(sample), show_img=False)

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.093660Z","iopub.execute_input":"2022-09-22T21:27:30.094041Z","iopub.status.idle":"2022-09-22T21:27:30.103245Z","shell.execute_reply.started":"2022-09-22T21:27:30.094001Z","shell.execute_reply":"2022-09-22T21:27:30.102278Z"}}
#metric_df.describe()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.104549Z","iopub.execute_input":"2022-09-22T21:27:30.105312Z","iopub.status.idle":"2022-09-22T21:27:30.114353Z","shell.execute_reply.started":"2022-09-22T21:27:30.105268Z","shell.execute_reply":"2022-09-22T21:27:30.113401Z"}}
#print(metric_df[metric_df.organ == 'kidney'].describe())
#metric_df[metric_df.organ == 'kidney'].dice.plot.hist(bins=10)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.115659Z","iopub.execute_input":"2022-09-22T21:27:30.115949Z","iopub.status.idle":"2022-09-22T21:27:30.125969Z","shell.execute_reply.started":"2022-09-22T21:27:30.115911Z","shell.execute_reply":"2022-09-22T21:27:30.124787Z"}}
#print(metric_df[metric_df.organ == 'prostate'].describe())
#metric_df[metric_df.organ == 'prostate'].dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.127928Z","iopub.execute_input":"2022-09-22T21:27:30.128316Z","iopub.status.idle":"2022-09-22T21:27:30.138077Z","shell.execute_reply.started":"2022-09-22T21:27:30.128252Z","shell.execute_reply":"2022-09-22T21:27:30.136965Z"}}
#print(metric_df[metric_df.organ == 'largeintestine'].describe())
#metric_df[metric_df.organ == 'largeintestine'].dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.139233Z","iopub.execute_input":"2022-09-22T21:27:30.139487Z","iopub.status.idle":"2022-09-22T21:27:30.152260Z","shell.execute_reply.started":"2022-09-22T21:27:30.139450Z","shell.execute_reply":"2022-09-22T21:27:30.151241Z"}}
#print(metric_df[metric_df.organ == 'spleen'].describe())
#metric_df[metric_df.organ == 'spleen'].dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.154137Z","iopub.execute_input":"2022-09-22T21:27:30.154558Z","iopub.status.idle":"2022-09-22T21:27:30.163013Z","shell.execute_reply.started":"2022-09-22T21:27:30.154516Z","shell.execute_reply":"2022-09-22T21:27:30.161980Z"}}
#print(p_metric_df.describe())
#p_metric_df.dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.164347Z","iopub.execute_input":"2022-09-22T21:27:30.164591Z","iopub.status.idle":"2022-09-22T21:27:30.174570Z","shell.execute_reply.started":"2022-09-22T21:27:30.164563Z","shell.execute_reply":"2022-09-22T21:27:30.173243Z"}}
#print(p_metric_df.describe())
#p_metric_df.dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.176239Z","iopub.execute_input":"2022-09-22T21:27:30.176597Z","iopub.status.idle":"2022-09-22T21:27:30.187581Z","shell.execute_reply.started":"2022-09-22T21:27:30.176513Z","shell.execute_reply":"2022-09-22T21:27:30.186721Z"}}

#print(p_metric_df.describe())
#p_metric_df.dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.189960Z","iopub.execute_input":"2022-09-22T21:27:30.190531Z","iopub.status.idle":"2022-09-22T21:27:30.201394Z","shell.execute_reply.started":"2022-09-22T21:27:30.190492Z","shell.execute_reply":"2022-09-22T21:27:30.200429Z"}}
#print(metric_df[metric_df.organ == 'lung'].describe())
#metric_df[metric_df.organ == 'lung'].dice.plot.hist(bins=12)
#plt.show()

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.202935Z","iopub.execute_input":"2022-09-22T21:27:30.203502Z","iopub.status.idle":"2022-09-22T21:27:30.211825Z","shell.execute_reply.started":"2022-09-22T21:27:30.203452Z","shell.execute_reply":"2022-09-22T21:27:30.210812Z"}}
#metric_df[(metric_df.organ == 'prostate')].sort_values("dice", ascending=False)

# %% [markdown]
# # Inference Loop

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.213141Z","iopub.execute_input":"2022-09-22T21:27:30.213372Z","iopub.status.idle":"2022-09-22T21:27:30.225662Z","shell.execute_reply.started":"2022-09-22T21:27:30.213347Z","shell.execute_reply":"2022-09-22T21:27:30.224531Z"}}
def inference_submit(models, test):
    # Predictions are stored as a list of dictionaries
    test_rows = []

    # Iterate over all test images
    for row_idx, row in tqdm(test.iterrows(), total=len(test)):
        organ = row['organ']
        data_source = row['data_source']

        # Preprocess Image
        image, image_size = get_image(row['id'], 'test')
        image_patches = extract_patches(image)
        
        # Check if dictionary
        model = models[organ] if type(models) is dict else models
        
        # Make Prediction
        mask_patches_pred = model.predict(image_patches)
        # Merge patches
        mask_pred = merge_patches(mask_patches_pred)
        # Resize mask to original size
        mask_pred_resized = resize_tensor(mask_pred, size=image_size, dtype=np.float32)

        # Resize and Binarize Mask
        mask_binary = (mask_pred_resized > ORGAN_THRESHOLD[data_source][organ]).astype(np.int8)

        if row_idx == 0:
            image_original = resize_tensor(image, size=image_size, dtype=np.uint8)
            visualize_test_images(image_original, mask_pred_resized, mask_binary, title=F"{organ}")

        # Append to Result
        test_rows.append({
            'id': row['id'],
            'rle': rle_encode_less_memory(mask_binary)
        })
    
    # Create dataset
    test_df = pd.DataFrame(test_rows)

    return test_df

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:30.227312Z","iopub.execute_input":"2022-09-22T21:27:30.229037Z","iopub.status.idle":"2022-09-22T21:27:43.565858Z","shell.execute_reply.started":"2022-09-22T21:27:30.228985Z","shell.execute_reply":"2022-09-22T21:27:43.564815Z"}}
test_df = inference_submit(model, test)

# %% [markdown]
# # Make Submission CSV

# %% [code] {"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-09-22T21:27:43.567465Z","iopub.execute_input":"2022-09-22T21:27:43.568274Z","iopub.status.idle":"2022-09-22T21:27:43.575552Z","shell.execute_reply.started":"2022-09-22T21:27:43.568233Z","shell.execute_reply":"2022-09-22T21:27:43.574492Z"}}
# Write Submission CSV
test_df.to_csv('submission.csv', index=False)

# %% [code] {"jupyter":{"outputs_hidden":false}}
print("END INFERENCE")