{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/efficientnetv2-head-1x1-endpoint/')\nsys.path.append('/kaggle/input/efficientnetv2-head-1x1-endpoint/efficientnetv2/')","metadata":{"papermill":{"duration":0.074686,"end_time":"2022-08-29T16:30:39.779455","exception":false,"start_time":"2022-08-29T16:30:39.704769","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:31.762541Z","iopub.execute_input":"2022-08-30T20:27:31.763064Z","iopub.status.idle":"2022-08-30T20:27:31.786551Z","shell.execute_reply.started":"2022-08-30T20:27:31.762969Z","shell.execute_reply":"2022-08-30T20:27:31.785847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/tfswin/keras_cv_attention_models')\nsys.path.append('/kaggle/input/keras_cv_attention_models/keras_cv_attention_models')","metadata":{"papermill":{"duration":0.064215,"end_time":"2022-08-29T16:30:39.899932","exception":false,"start_time":"2022-08-29T16:30:39.835717","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:31.810238Z","iopub.execute_input":"2022-08-30T20:27:31.810696Z","iopub.status.idle":"2022-08-30T20:27:31.815241Z","shell.execute_reply.started":"2022-08-30T20:27:31.810661Z","shell.execute_reply":"2022-08-30T20:27:31.814000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow.keras.backend as K\nimport tensorflow_addons as tfa\nimport matplotlib.pyplot as plt\n\nfrom kaggle_datasets import KaggleDatasets\nfrom tqdm.notebook import tqdm\nfrom multiprocessing import cpu_count\nfrom sklearn import metrics\nfrom sklearn.model_selection import StratifiedKFold\nfrom albumentations import *\nfrom keras_cv_attention_models import swin_transformer_v2\n\nimport effnetv2_model\nimport re\nimport os\nimport io\nimport time\nimport pickle\nimport math\nimport random\nimport sys\nimport cv2\nimport gc\n\nimport matplotlib as mpl\nmpl.rcParams['figure.dpi'] = 150\n\nprint(f'tensorflow version: {tf.__version__}')\nprint(f'tensorflow keras version: {tf.keras.__version__}')\nprint(f'tensorflow addons version: {tfa.__version__}')\nprint(f'python version: P{sys.version}')","metadata":{"papermill":{"duration":10.930594,"end_time":"2022-08-29T16:30:50.885001","exception":false,"start_time":"2022-08-29T16:30:39.954407","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:31.857591Z","iopub.execute_input":"2022-08-30T20:27:31.857863Z","iopub.status.idle":"2022-08-30T20:27:42.306075Z","shell.execute_reply.started":"2022-08-30T20:27:31.857837Z","shell.execute_reply":"2022-08-30T20:27:42.304977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect hardware, return appropriate distribution strategy\ntry:\n    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.\n    print('Running on TPU ', TPU.master())\nexcept ValueError:\n    print('Running on GPU')\n    TPU = None\n\nif TPU:\n    IS_TPU = True\n    tf.config.experimental_connect_to_cluster(TPU)\n    tf.tpu.experimental.initialize_tpu_system(TPU)\n    strategy = tf.distribute.experimental.TPUStrategy(TPU)\nelse:\n    IS_TPU = False\n    strategy = tf.distribute.get_strategy() # default distribution strategy in Tensorflow. Works on CPU and single GPU.\n\nREPLICAS = strategy.num_replicas_in_sync\nprint(f'REPLICAS: {REPLICAS}, IS_TPU: {IS_TPU}')","metadata":{"papermill":{"duration":6.400114,"end_time":"2022-08-29T16:30:57.339277","exception":false,"start_time":"2022-08-29T16:30:50.939163","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:42.308008Z","iopub.execute_input":"2022-08-30T20:27:42.308277Z","iopub.status.idle":"2022-08-30T20:27:48.506278Z","shell.execute_reply.started":"2022-08-30T20:27:42.308247Z","shell.execute_reply":"2022-08-30T20:27:48.505406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For TPU's the dataset needs to be stored in Google Cloud\n# Retrieve the Google Cloud location of the dataset\nGCS_DS_PATH = KaggleDatasets().get_gcs_path('hubmap-hpa-hacking-the-human-body-1024x1024')","metadata":{"papermill":{"duration":0.53954,"end_time":"2022-08-29T16:30:57.934893","exception":false,"start_time":"2022-08-29T16:30:57.395353","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:48.507566Z","iopub.execute_input":"2022-08-30T20:27:48.507877Z","iopub.status.idle":"2022-08-30T20:27:48.887440Z","shell.execute_reply.started":"2022-08-30T20:27:48.507848Z","shell.execute_reply":"2022-08-30T20:27:48.886489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 43\nDEBUG = False\n\n# Original Image Size\nIMG0_SIZE_FULL = 1024\nIMG0_SIZE = 1024\nN_PATCHES_PER_IMAGE0 = (IMG0_SIZE_FULL // IMG0_SIZE) ** 2\nN_CHANNELS = 3\n\nINPUT_SHAPE0 = (IMG0_SIZE, IMG0_SIZE, N_CHANNELS)\n\n# Image dimensions\nIMG_SIZE_FULL = 640\nIMG_SIZE = 640\n\nMASK0_SIZE = 1024\nMASK_SIZE = 640\n\n\nN_PATCHES_PER_IMAGE = (IMG_SIZE_FULL // IMG_SIZE) ** 2\nN_CHANNELS = 3\n\nINPUT_SHAPE = (IMG_SIZE, IMG_SIZE, N_CHANNELS)\n\n# EfficientNet version, b0/b1/b2/b3/b4/b5/b6/b7\nEFN_SIZE = 'b0' if DEBUG else 'b7'\nLR_HEAD = 1e-3\nLR_MAX_WHOLE = 8e-4\nN_FOLDS = 5\nTRAIN_FOLDS = [0]\n\nCOMBINE_EPOCHS = 1 if DEBUG else 10\nEPOCHS_WHOLE = (5 if DEBUG else 200) // COMBINE_EPOCHS\nMOMENTUM = 0.00\n\n# Batch size\nBATCH_SIZE = 8 * REPLICAS\n\n# Dataset Mean and Standard Deviation\nMEAN = np.load(f'/kaggle/input/hubmap-hpa-hacking-the-human-body-1024x1024/MEAN.npy')\nSTD = np.load(f'/kaggle/input/hubmap-hpa-hacking-the-human-body-1024x1024/STD.npy')\n\n# Tensorflow AUTO flag\nAUTO = tf.data.experimental.AUTOTUNE\nif TPU:\n    NUM_PARALLEL_CALLS = AUTO\nelse:\n    NUM_PARALLEL_CALLS = cpu_count()\n\nprint(f'BATCH_SIZE: {BATCH_SIZE}, NUM_PARALLEL_CALLS: {NUM_PARALLEL_CALLS}')","metadata":{"papermill":{"duration":0.087464,"end_time":"2022-08-29T16:30:58.079118","exception":false,"start_time":"2022-08-29T16:30:57.991654","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:48.889921Z","iopub.execute_input":"2022-08-30T20:27:48.890722Z","iopub.status.idle":"2022-08-30T20:27:48.918076Z","shell.execute_reply.started":"2022-08-30T20:27:48.890673Z","shell.execute_reply":"2022-08-30T20:27:48.917443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'MEAN: {MEAN}, STD: {STD}')","metadata":{"papermill":{"duration":0.071336,"end_time":"2022-08-29T16:30:58.206326","exception":false,"start_time":"2022-08-29T16:30:58.134990","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:48.919345Z","iopub.execute_input":"2022-08-30T20:27:48.919579Z","iopub.status.idle":"2022-08-30T20:27:48.925144Z","shell.execute_reply.started":"2022-08-30T20:27:48.919553Z","shell.execute_reply":"2022-08-30T20:27:48.924476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Seed","metadata":{"papermill":{"duration":0.054447,"end_time":"2022-08-29T16:30:58.318199","exception":false,"start_time":"2022-08-29T16:30:58.263752","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Seed all random number generators\ndef seed_everything(seed=SEED):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    \n\nseed_everything()","metadata":{"papermill":{"duration":0.063167,"end_time":"2022-08-29T16:30:58.436859","exception":false,"start_time":"2022-08-29T16:30:58.373692","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:48.926550Z","iopub.execute_input":"2022-08-30T20:27:48.926771Z","iopub.status.idle":"2022-08-30T20:27:48.937356Z","shell.execute_reply.started":"2022-08-30T20:27:48.926745Z","shell.execute_reply":"2022-08-30T20:27:48.936284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"papermill":{"duration":0.05446,"end_time":"2022-08-29T16:30:58.547538","exception":false,"start_time":"2022-08-29T16:30:58.493078","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/hubmap-organ-segmentation/train.csv')\nN_SAMPLES = len(train)\nprint(f'N_SAMPLES: {N_SAMPLES}')\n\n# Add Ordinal Encoded Organ\ntrain['organ_ordinal'] = train['organ'].astype('category').cat.codes\nN_ORGANS = train['organ'].nunique()\nORGANS = sorted(train['organ'].unique())\norg_ord2org = dict(enumerate(train['organ'].astype('category').cat.categories))\nprint(f'N_ORGANS: {N_ORGANS}, ORGANS: {ORGANS}')\n\ndisplay(train.head())\ndisplay(train.info())","metadata":{"papermill":{"duration":0.476157,"end_time":"2022-08-29T16:30:59.078873","exception":false,"start_time":"2022-08-29T16:30:58.602716","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:48.938853Z","iopub.execute_input":"2022-08-30T20:27:48.939096Z","iopub.status.idle":"2022-08-30T20:27:49.324899Z","shell.execute_reply.started":"2022-08-30T20:27:48.939063Z","shell.execute_reply":"2022-08-30T20:27:49.324068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility Functions","metadata":{"papermill":{"duration":0.057836,"end_time":"2022-08-29T16:30:59.194379","exception":false,"start_time":"2022-08-29T16:30:59.136543","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def tf_rand_int(minval, maxval, dtype=tf.int64):\n    minval = tf.cast(minval, dtype)\n    maxval = tf.cast(maxval, dtype)\n    return tf.random.uniform(shape=(), minval=minval, maxval=maxval, dtype=dtype)\n\ndef one_in(k):\n    return 0 == tf_rand_int(0, k)","metadata":{"papermill":{"duration":0.066974,"end_time":"2022-08-29T16:30:59.318080","exception":false,"start_time":"2022-08-29T16:30:59.251106","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.326215Z","iopub.execute_input":"2022-08-30T20:27:49.326442Z","iopub.status.idle":"2022-08-30T20:27:49.332043Z","shell.execute_reply.started":"2022-08-30T20:27:49.326416Z","shell.execute_reply":"2022-08-30T20:27:49.331079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{"papermill":{"duration":0.056565,"end_time":"2022-08-29T16:30:59.433679","exception":false,"start_time":"2022-08-29T16:30:59.377114","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Source: https://github.com/shruti-jadon/Semantic-Segmentation-Loss-Functions/blob/master/loss_functions.py\nimport tensorflow as tf\nimport tensorflow.keras.backend as K\nfrom tensorflow.keras.losses import binary_crossentropy\n\nbeta = 0.25\nalpha = 0.25\ngamma = 2\nepsilon = 1e-5\nsmooth = 1\nthreshold=0.50\n\nclass Semantic_loss_functions(object):\n    def __init__(self):\n        print ('semantic loss functions initialized')\n        \n    def iou(self, y_true, y_pred):\n        y_true = tf.cast(y_true, tf.float32)\n        y_pred = tf.where(y_pred > threshold, x=1.0, y=0.0)\n        y_true_f = K.flatten(y_true)\n        y_pred_f = K.flatten(y_pred)\n        intersection = K.sum(y_true_f * y_pred_f)\n        union = K.sum(y_true_f + y_pred_f) - intersection\n        return intersection / (union + epsilon)\n\n    def dice_coef(self, y_true, y_pred):\n        y_true = tf.cast(y_true, tf.float32)\n        y_true_f = K.flatten(y_true)\n        y_pred_f = K.flatten(y_pred)\n        intersection = K.sum(y_true_f * y_pred_f, axis=1)\n        return (2. * intersection + K.epsilon()) / (\n                    K.sum(y_true_f) + K.sum(y_pred_f) + K.epsilon())\n\n    def sensitivity(self, y_true, y_pred):\n        true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))\n        possible_positives = K.sum(K.round(K.clip(y_true, 0, 1)))\n        return true_positives / (possible_positives + K.epsilon())\n\n    def specificity(self, y_true, y_pred):\n        true_negatives = K.sum(\n            K.round(K.clip((1 - y_true) * (1 - y_pred), 0, 1)))\n        possible_negatives = K.sum(K.round(K.clip(1 - y_true, 0, 1)))\n        return true_negatives / (possible_negatives + K.epsilon())\n\n    def convert_to_logits(self, y_pred):\n        y_pred = tf.clip_by_value(y_pred, tf.keras.backend.epsilon(),\n                                  1 - tf.keras.backend.epsilon())\n        return tf.math.log(y_pred / (1 - y_pred))\n\n    def weighted_cross_entropyloss(self, y_true, y_pred):\n        y_pred = self.convert_to_logits(y_pred)\n        pos_weight = beta / (1 - beta)\n        loss = tf.nn.weighted_cross_entropy_with_logits(logits=y_pred,\n                                                        targets=y_true,\n                                                        pos_weight=pos_weight)\n        return tf.reduce_mean(loss)\n\n    def focal_loss_with_logits(self, logits, targets, alpha, gamma, y_pred):\n        weight_a = alpha * (1 - y_pred) ** gamma * targets\n        weight_b = (1 - alpha) * y_pred ** gamma * (1 - targets)\n\n        return (tf.math.log1p(tf.exp(-tf.abs(logits))) + tf.nn.relu(\n            -logits)) * (weight_a + weight_b) + logits * weight_b\n\n    def focal_loss(self, y_true, y_pred):\n        y_pred = tf.clip_by_value(y_pred, tf.keras.backend.epsilon(),\n                                  1 - tf.keras.backend.epsilon())\n        logits = tf.math.log(y_pred / (1 - y_pred))\n\n        loss = self.focal_loss_with_logits(logits=logits, targets=y_true,\n                                      alpha=alpha, gamma=gamma, y_pred=y_pred)\n\n        return tf.reduce_mean(loss)\n\n    def depth_softmax(self, matrix):\n        sigmoid = lambda x: 1 / (1 + K.exp(-x))\n        sigmoided_matrix = sigmoid(matrix)\n        softmax_matrix = sigmoided_matrix / K.sum(sigmoided_matrix, axis=0)\n        return softmax_matrix\n\n    def generalized_dice_coefficient(self, y_true, y_pred):\n        smooth = 1e0\n        y_true = tf.cast(y_true, tf.float32)\n        y_true_f = K.flatten(y_true)\n        y_pred_f = K.flatten(y_pred)\n        intersection = K.sum(y_true_f * y_pred_f)\n        return (2. * intersection + smooth) / (\n                    K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n        \n    def dice_loss(self, y_true, y_pred):\n        loss = 1 - self.generalized_dice_coefficient(y_true, y_pred)\n        return loss\n    \n    def dice_loss_symmetric(self, y_true, y_pred):\n        loss = 1 - self.generalized_dice_coefficient(y_true, y_pred)\n        loss_neg = 1 - self.generalized_dice_coefficient(1 - y_true, 1 - y_pred)\n        return 0.50 * (loss + loss_neg)\n\n    def bce_dice_loss(self, y_true, y_pred):\n        loss = binary_crossentropy(y_true, y_pred) + \\\n               self.dice_loss(y_true, y_pred)\n        return loss / 2.0\n\n    def confusion(self, y_true, y_pred):\n        smooth = 1\n        y_pred_pos = K.clip(y_pred, 0, 1)\n        y_pred_neg = 1 - y_pred_pos\n        y_pos = K.clip(y_true, 0, 1)\n        y_neg = 1 - y_pos\n        tp = K.sum(y_pos * y_pred_pos)\n        fp = K.sum(y_neg * y_pred_pos)\n        fn = K.sum(y_pos * y_pred_neg)\n        prec = (tp + smooth) / (tp + fp + smooth)\n        recall = (tp + smooth) / (tp + fn + smooth)\n        return prec, recall\n\n    def true_positive(self, y_true, y_pred):\n        smooth = 1\n        y_pred_pos = K.round(K.clip(y_pred, 0, 1))\n        y_pos = K.round(K.clip(y_true, 0, 1))\n        tp = (K.sum(y_pos * y_pred_pos) + smooth) / (K.sum(y_pos) + smooth)\n        return tp\n\n    def true_negative(self, y_true, y_pred):\n        smooth = 1\n        y_pred_pos = K.round(K.clip(y_pred, 0, 1))\n        y_pred_neg = 1 - y_pred_pos\n        y_pos = K.round(K.clip(y_true, 0, 1))\n        y_neg = 1 - y_pos\n        tn = (K.sum(y_neg * y_pred_neg) + smooth) / (K.sum(y_neg) + smooth)\n        return tn\n\n    def tversky_index(self, y_true, y_pred):\n        y_true = tf.cast(y_true, tf.float32)\n        y_true_pos = K.flatten(y_true)\n        y_pred_pos = K.flatten(y_pred)\n        true_pos = K.sum(y_true_pos * y_pred_pos)\n        false_neg = K.sum(y_true_pos * (1 - y_pred_pos))\n        false_pos = K.sum((1 - y_true_pos) * y_pred_pos)\n        alpha = 0.75\n        return (true_pos + smooth) / (true_pos + alpha * false_neg + (\n                    1 - alpha) * false_pos + smooth)\n\n    def tversky_loss(self, y_true, y_pred):\n        return 1 - self.tversky_index(y_true, y_pred)\n\n    def focal_tversky(self, y_true, y_pred, gamme=2.0):\n        pt_1 = self.tversky_index(y_true, y_pred)\n        return K.pow((1 - pt_1), gamma)\n\n    def log_cosh_dice_loss(self, y_true, y_pred):\n        y_true = tf.cast(y_true, tf.float32)\n        x = self.dice_loss(y_true, y_pred)\n        return tf.math.log((tf.exp(x) + tf.exp(-x)) / 2.0)","metadata":{"papermill":{"duration":0.096834,"end_time":"2022-08-29T16:30:59.587310","exception":false,"start_time":"2022-08-29T16:30:59.490476","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.333271Z","iopub.execute_input":"2022-08-30T20:27:49.333482Z","iopub.status.idle":"2022-08-30T20:27:49.368130Z","shell.execute_reply.started":"2022-08-30T20:27:49.333457Z","shell.execute_reply":"2022-08-30T20:27:49.367292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEMANTIC_LOSS_FUNCTIONS = Semantic_loss_functions()","metadata":{"papermill":{"duration":0.071533,"end_time":"2022-08-29T16:30:59.718080","exception":false,"start_time":"2022-08-29T16:30:59.646547","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.370530Z","iopub.execute_input":"2022-08-30T20:27:49.370892Z","iopub.status.idle":"2022-08-30T20:27:49.380088Z","shell.execute_reply.started":"2022-08-30T20:27:49.370864Z","shell.execute_reply":"2022-08-30T20:27:49.379169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# FPN ","metadata":{"papermill":{"duration":0.057285,"end_time":"2022-08-29T16:30:59.834630","exception":false,"start_time":"2022-08-29T16:30:59.777345","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def FPN(xs, output_channels, last_layer, debug=False):\n    def _conv(x_idx, x):\n        x = tf.keras.layers.ZeroPadding2D(padding=1, name=f'pad_1_{x_idx}_FPN')(x)\n        x = tf.keras.layers.Conv2D(output_channels * 2, 3, padding='SAME', kernel_initializer='he_normal', activation='relu', name=f'Conv2D_1_{x_idx}_FPN')(x)\n        x = tf.keras.layers.BatchNormalization(name=f'bn_1_{x_idx}_FPN')(x)\n        x = tf.keras.layers.ZeroPadding2D(padding=1, name=f'pad_2_{x_idx}_FPN')(x)\n        x = tf.keras.layers.Conv2D(output_channels, 3, padding='SAME', kernel_initializer='he_normal', name=f'Conv2D_2_{x_idx}_FPN')(x)\n        x = tf.image.resize(x, size=target_size, method=tf.image.ResizeMethod.BILINEAR)\n        x = tf.nn.relu(x)\n        return x\n\n    target_size = last_layer.shape[1:3]\n    xs = tf.keras.layers.Concatenate(name='concat_xs_FPN')([_conv(x_idx, x) for x_idx, x in enumerate(xs)])\n    x = tf.keras.layers.Concatenate(name='concat_x_FPN')([xs, last_layer])\n\n    if debug:\n        return x, xs\n    else:\n        return x","metadata":{"papermill":{"duration":0.069931,"end_time":"2022-08-29T16:30:59.961694","exception":false,"start_time":"2022-08-29T16:30:59.891763","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.381255Z","iopub.execute_input":"2022-08-30T20:27:49.381481Z","iopub.status.idle":"2022-08-30T20:27:49.393826Z","shell.execute_reply.started":"2022-08-30T20:27:49.381456Z","shell.execute_reply":"2022-08-30T20:27:49.392887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ASPP","metadata":{"papermill":{"duration":0.060547,"end_time":"2022-08-29T16:31:00.081315","exception":false,"start_time":"2022-08-29T16:31:00.020768","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def ASPP(x, mid_c=320, dilations=[1, 2, 3, 4], out_c=640, debug=False):\n    def _aspp_module(x, filters, kernel_size, padding, dilation, groups=1):\n        x = tf.keras.layers.ZeroPadding2D(padding=padding)(x)\n        x = tf.keras.layers.Conv2D(\n                filters=filters,\n                kernel_size=kernel_size,\n                dilation_rate=dilation,\n                groups=1,\n                kernel_initializer='he_uniform',\n            )(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.nn.relu(x)\n        \n        return x\n    \n    x0 = tf.math.reduce_max(x, axis=(1,2), keepdims=True)\n    x0 = tf.keras.layers.Conv2D(filters=mid_c, kernel_size=1, strides=1, kernel_initializer='he_uniform', use_bias=False)(x0)\n    x0 = tf.keras.layers.BatchNormalization(gamma_initializer=tf.constant_initializer(value=0.25))(x0)\n    x0 = tf.nn.relu(x0)\n                                  \n                                  \n    xs = (\n        [_aspp_module(x, mid_c, 1, padding=0, dilation=1)] +\n        [_aspp_module(x, mid_c, 3, padding=d, dilation=d, groups=4) for d in dilations]\n    )\n    \n    x0= tf.image.resize(x0, size=xs[0].shape[1:3])\n    x = tf.keras.layers.Concatenate()([x0] + xs)\n    x = tf.keras.layers.Conv2D(filters=out_c, kernel_size=1, kernel_initializer='he_uniform', use_bias=False)(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = tf.nn.relu(x)\n                       \n    if debug:\n        return x, x0, xs\n    else:\n        return x","metadata":{"papermill":{"duration":0.075922,"end_time":"2022-08-29T16:31:00.215517","exception":false,"start_time":"2022-08-29T16:31:00.139595","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.395883Z","iopub.execute_input":"2022-08-30T20:27:49.396297Z","iopub.status.idle":"2022-08-30T20:27:49.411281Z","shell.execute_reply.started":"2022-08-30T20:27:49.396254Z","shell.execute_reply":"2022-08-30T20:27:49.410531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Upsample","metadata":{"papermill":{"duration":0.064935,"end_time":"2022-08-29T16:31:00.344672","exception":false,"start_time":"2022-08-29T16:31:00.279737","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def PixelShuffle(x, upscale_factor=2):\n    _, w, h, c = x.shape\n    n = -1\n\n    c_out = c // upscale_factor ** 2\n    w_out = w * upscale_factor\n    h_out = h * upscale_factor\n\n    x = tf.reshape(x, [-1, upscale_factor, upscale_factor, w, h, c_out])\n    x = tf.transpose(x, [0, 3, 1, 4, 2, 5])\n    x = tf.reshape(x, [-1, w_out, h_out, c_out])\n\n    return x","metadata":{"papermill":{"duration":0.082804,"end_time":"2022-08-29T16:31:00.502806","exception":false,"start_time":"2022-08-29T16:31:00.420002","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.412296Z","iopub.execute_input":"2022-08-30T20:27:49.413190Z","iopub.status.idle":"2022-08-30T20:27:49.427298Z","shell.execute_reply.started":"2022-08-30T20:27:49.413130Z","shell.execute_reply":"2022-08-30T20:27:49.426299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Inspiration: https://www.tensorflow.org/tutorials/generative/pix2pix#build_an_input_pipeline_with_tfdata\ndef upsample(x, concat, target_filters, name, conv2dt_kernel_init_max, relu=True, dropout=0, debug=False):\n    x_up = PixelShuffle(x)\n    \n    concat = tf.keras.layers.BatchNormalization(\n        gamma_initializer=tf.constant_initializer(value=0.25),\n        name=f'BatchNormalization_{name}'\n    )(concat)\n    x = tf.keras.layers.Concatenate(name=f'Concatenate_{name}')([x_up, concat])\n    x = tf.nn.relu(x)\n    \n        \n    x = tf.keras.layers.Conv2D(target_filters, 3, padding='SAME', kernel_initializer='he_uniform', activation='relu', name=f'Conv2D_1_{name}')(x)\n    x = tf.keras.layers.Conv2D(target_filters, 3, padding='SAME', kernel_initializer='he_uniform', name=f'Conv2D_2_{name}')(x)\n    \n    if relu:\n        x = tf.nn.relu(x)\n    \n    x = tf.keras.layers.Dropout(dropout, name=f'Dropout_{name}')(x)\n\n    if debug:\n        return x, x_up, concat\n    else:\n        return x","metadata":{"papermill":{"duration":0.069943,"end_time":"2022-08-29T16:31:00.630581","exception":false,"start_time":"2022-08-29T16:31:00.560638","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.428874Z","iopub.execute_input":"2022-08-30T20:27:49.429598Z","iopub.status.idle":"2022-08-30T20:27:49.440215Z","shell.execute_reply.started":"2022-08-30T20:27:49.429555Z","shell.execute_reply":"2022-08-30T20:27:49.439514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"papermill":{"duration":0.060665,"end_time":"2022-08-29T16:31:00.752268","exception":false,"start_time":"2022-08-29T16:31:00.691603","status":"completed"},"tags":[]}},{"cell_type":"code","source":"GCS_WEIGHTS_PATH = KaggleDatasets().get_gcs_path('efficientnetv2-head-1x1-endpoint')","metadata":{"papermill":{"duration":0.458901,"end_time":"2022-08-29T16:31:01.269633","exception":false,"start_time":"2022-08-29T16:31:00.810732","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.441625Z","iopub.execute_input":"2022-08-30T20:27:49.442557Z","iopub.status.idle":"2022-08-30T20:27:49.833552Z","shell.execute_reply.started":"2022-08-30T20:27:49.442492Z","shell.execute_reply":"2022-08-30T20:27:49.832625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"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, debug=DEBUG):\n    # enable XLA optmizations\n    tf.config.optimizer.set_jit(True)\n    # Set seed for deterministic weights initialization\n    seed_everything()\n    \n    with strategy.scope():\n        # EfficientNetV2 Backbone # \n        cnn = effnetv2_model.get_model(f'efficientnet-{EFN_SIZE}', include_top=False, weights=None if IS_TPU else 'jft', model_config={ 'conv_dropout': dropout_cnn })\n        cnn.trainable = cnn_trainable\n        if IS_TPU:\n            WEIGHT_PATH = f'{GCS_WEIGHTS_PATH}/noisy_student_efficientnet-{EFN_SIZE}'\n            ckpt = tf.train.latest_checkpoint(WEIGHT_PATH)\n            cnn.load_weights(ckpt)\n\n        # Inputs, note the names are equal to the dictionary keys in the dataset\n        image = tf.keras.layers.Input(INPUT_SHAPE, name='image', dtype=tf.float32)\n        image_norm = tf.cast(image, tf.float32) / 255\n        image_norm = tf.keras.layers.experimental.preprocessing.Normalization(mean=MEAN, variance=STD, dtype=tf.float32)(image_norm)\n\n        embedding, up6, up5, up4, up3, up2, up1 = cnn(image_norm, with_endpoints=True)\n        \n        if debug:\n            print(f'embedding shape: {embedding.shape} up1 shape: {up1.shape}, up2 shape: {up2.shape}')\n            print(f'up3 shape: {up3.shape}, up4 shape: {up4.shape}, up5 shape: {up5.shape}, up6 shape: {up6.shape}')\n                \n        if debug:\n            dec0, x0, (xs0, xs1, xs2, xs3, xs4) = ASPP(up2, debug=True)\n        else:\n            dec0 = ASPP(up2)\n        \n        dec0 = tf.keras.layers.Dropout(0.30)(dec0)\n\n        if debug:\n            dec1, dec1_up, dec1_concat = upsample(dec0, up3, up4.shape[-1] * 4, 'upsample1', 0.02, dropout=dropout_decoder, debug=True)\n            dec2, dec2_up, dec2_concat = upsample(dec1, up4, up5.shape[-1] * 2, 'upsample2', 0.02, dropout=dropout_decoder, debug=True)\n            dec3, dec3_up, dec3_concat = upsample(dec2, up5, up6.shape[-1] * 2, 'upsample3', 0.02, debug=True)\n            dec4, dec4_up, dec4_concat = upsample(dec3, up6, 48, 'upsample4', 0.02, debug=True)\n        else:\n            dec1 = upsample(dec0, up3, up4.shape[-1] * 4, 'upsample1', 0.02, dropout=dropout_decoder)\n            dec2 = upsample(dec1, up4, up5.shape[-1] * 2, 'upsample2', 0.02, dropout=dropout_decoder)\n            dec3 = upsample(dec2, up5, up6.shape[-1] * 2, 'upsample3', 0.02)\n            dec4 = upsample(dec3, up6, 48, 'upsample4', 0.02)\n        \n        if debug:\n            dec_fpn, dec_fpn_xs = FPN([dec0, dec1, dec2, dec3], 32, dec4, debug=True)\n        else:\n            dec_fpn = FPN([dec0, dec1, dec2, dec3], 32, dec4)\n        \n        if debug:\n            print(f'dec0 shape: {dec0.shape}, dec1 shape: {dec1.shape}, dec2 shape: {dec2.shape}, dec3 shape: {dec3.shape}, dec4 shape: {dec4.shape}')\n            print(f'dec_fpn shape: {dec_fpn.shape}')\n\n        # Head\n        x = tf.keras.layers.Conv2D(dec_fpn.shape[-1] , 3, padding='SAME', kernel_initializer='he_uniform', activation='relu', name=f'Conv2D_1_head')(dec_fpn)\n        x = tf.keras.layers.Dropout(0.10)(x)\n        x = tf.keras.layers.Conv2D(\n            filters=1,\n            kernel_size=1,\n            padding='SAME',\n            kernel_initializer=tf.random_normal_initializer(0.00, 0.05),\n            activation=None if debug else 'sigmoid',\n            name='Conv2D_3_head'\n        )(x)\n        output = tf.image.resize(x, size=[IMG_SIZE, IMG_SIZE], method=tf.image.ResizeMethod.BILINEAR)\n\n        # We will use the famous Adam optimizer for fast learning\n        optimizer = tf.optimizers.Adam(learning_rate=lr, epsilon=eps, clipnorm=clipnorm)\n\n\n        # Loss\n        loss = tf.keras.losses.BinaryCrossentropy()\n        \n        # Metrics\n        metrics = [\n            SEMANTIC_LOSS_FUNCTIONS.iou,\n            tf.keras.metrics.Precision(),\n            tf.keras.metrics.Recall(),\n            tf.keras.metrics.AUC(),\n            tf.keras.metrics.BinaryAccuracy(),\n        ]\n\n        if debug:\n            model = tf.keras.models.Model(inputs=image, outputs=[\n                image_norm,\n                embedding, up6, up5, up4, up3, up2, up1,\n                dec0, x0, xs0, xs1, xs2, xs3, xs4,\n                dec1, dec1_up, dec1_concat,\n                dec2, dec2_up, dec2_concat,\n                dec3, dec3_up, dec3_concat,\n                dec4, dec4_up, dec4_concat,\n                dec_fpn, dec_fpn_xs, output\n                \n            ])\n        else:\n            model = tf.keras.models.Model(inputs=image, outputs=[output])\n        \n        model.compile(optimizer=optimizer, loss=loss, metrics=metrics)\n\n        if file_path:\n            print('Loading pretrained weights...')\n            model.load_weights(file_path)\n\n        return model","metadata":{"papermill":{"duration":0.089271,"end_time":"2022-08-29T16:31:01.420873","exception":false,"start_time":"2022-08-29T16:31:01.331602","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.837316Z","iopub.execute_input":"2022-08-30T20:27:49.837582Z","iopub.status.idle":"2022-08-30T20:27:49.862905Z","shell.execute_reply.started":"2022-08-30T20:27:49.837552Z","shell.execute_reply":"2022-08-30T20:27:49.862006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Pretrained File Path: '/kaggle/input/sartorius-training-dataset/model.h5'\ntf.keras.backend.clear_session()\ngc.collect()\n\nmodel = get_model(file_path=None, debug=False)","metadata":{"papermill":{"duration":135.357452,"end_time":"2022-08-29T16:33:16.836859","exception":false,"start_time":"2022-08-29T16:31:01.479407","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:27:49.864428Z","iopub.execute_input":"2022-08-30T20:27:49.866745Z","iopub.status.idle":"2022-08-30T20:30:02.804518Z","shell.execute_reply.started":"2022-08-30T20:27:49.866706Z","shell.execute_reply":"2022-08-30T20:30:02.803645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot model summary\nmodel.summary()","metadata":{"papermill":{"duration":0.196601,"end_time":"2022-08-29T16:33:17.094206","exception":false,"start_time":"2022-08-29T16:33:16.897605","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:30:02.805949Z","iopub.execute_input":"2022-08-30T20:30:02.806223Z","iopub.status.idle":"2022-08-30T20:30:02.915370Z","shell.execute_reply.started":"2022-08-30T20:30:02.806190Z","shell.execute_reply":"2022-08-30T20:30:02.914423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf.keras.utils.plot_model(model, show_shapes=True, show_dtype=True, show_layer_names=True, expand_nested=False)","metadata":{"papermill":{"duration":4.615832,"end_time":"2022-08-29T16:33:21.770499","exception":false,"start_time":"2022-08-29T16:33:17.154667","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:30:02.917398Z","iopub.execute_input":"2022-08-30T20:30:02.917731Z","iopub.status.idle":"2022-08-30T20:30:07.270849Z","shell.execute_reply.started":"2022-08-30T20:30:02.917692Z","shell.execute_reply":"2022-08-30T20:30:07.269487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading Weights Check","metadata":{"papermill":{"duration":0.093543,"end_time":"2022-08-29T16:33:21.957840","exception":false,"start_time":"2022-08-29T16:33:21.864297","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# from skimage.data import chelsea\n# def check_loading_imagenet_weights():\n#     image = tf.keras.applications.imagenet_utils.preprocess_input(chelsea(), mode='torch') # Chelsea the cat\n#     display(pd.Series(imm.flatten()).describe().to_frame('Value'))\n\n#     pred = model(tf.expand_dims(tf.image.resize(imm, model.input_shape[1:3]), 0))[-1].numpy()\n#     print(tf.keras.applications.imagenet_utils.decode_predictions(pred)[0])\n    \n# check_loading_imagenet_weights()","metadata":{"papermill":{"duration":0.10437,"end_time":"2022-08-29T16:33:22.159869","exception":false,"start_time":"2022-08-29T16:33:22.055499","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:30:07.273050Z","iopub.execute_input":"2022-08-30T20:30:07.273915Z","iopub.status.idle":"2022-08-30T20:30:07.277581Z","shell.execute_reply.started":"2022-08-30T20:30:07.273866Z","shell.execute_reply":"2022-08-30T20:30:07.276788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.092664,"end_time":"2022-08-29T16:33:22.347366","exception":false,"start_time":"2022-08-29T16:33:22.254702","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def benchmark_dataset(dataset, num_epochs=3, n_steps_per_epoch=10, bs=BATCH_SIZE):\n    start_time = time.perf_counter()\n    dataset_iter = iter(dataset)\n    for epoch_num in range(num_epochs):\n        for idx in range(n_steps_per_epoch):\n            images, labels = next(dataset_iter)\n            if idx == 0:\n                epoch_start = time.perf_counter()\n            elif idx == 1 and epoch_num == 0:\n                print(f'image shape: {images.shape}, image dtype: {images.dtype}')\n            else:\n                pass\n        epoch_t = time.perf_counter() - epoch_start\n        mean_step_t = round(epoch_t / n_steps_per_epoch * 1000, 1)\n        n_imgs_per_s = int(1 / (mean_step_t / 1000) * bs)\n        print(f'epoch {epoch_num} took: {round(epoch_t, 2)} sec, mean step duration: {mean_step_t}ms, images/s: {n_imgs_per_s}')","metadata":{"papermill":{"duration":0.107731,"end_time":"2022-08-29T16:33:22.549965","exception":false,"start_time":"2022-08-29T16:33:22.442234","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:30:07.278864Z","iopub.execute_input":"2022-08-30T20:30:07.279121Z","iopub.status.idle":"2022-08-30T20:30:07.294467Z","shell.execute_reply.started":"2022-08-30T20:30:07.279093Z","shell.execute_reply":"2022-08-30T20:30:07.293424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plots a batch of images\ndef show_batch(dataset, cols=6):\n    imgs, lbls, orgs = next(iter(dataset))\n    imgs = imgs.numpy()\n    rows = imgs.shape[0] // cols\n    # Plot\n    fig, axes = plt.subplots(nrows=rows, ncols=cols, figsize=(cols*4, rows*4))\n    for r in range(rows):\n        for c in range(cols // 3):\n            img = imgs[r*cols+c]\n            lbl = lbls[r*cols+c]\n            org = orgs[r*cols+c].numpy().decode()\n            \n            axes[r, c*3].imshow(img)\n            \n            axes[r, c*3+1].imshow(img)\n            axes[r, c*3+1].imshow((lbl * np.array([255, 0, 0])), alpha=0.50)\n            axes[r, c*3+1].set_title(f'std: {img.std():.1f}, organ: {org}')\n            \n            axes[r, c*3+2].imshow(lbl)","metadata":{"papermill":{"duration":0.108888,"end_time":"2022-08-29T16:33:22.754998","exception":false,"start_time":"2022-08-29T16:33:22.646110","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:58:14.384770Z","iopub.execute_input":"2022-08-30T20:58:14.385363Z","iopub.status.idle":"2022-08-30T20:58:14.394638Z","shell.execute_reply.started":"2022-08-30T20:58:14.385324Z","shell.execute_reply":"2022-08-30T20:58:14.393640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_image(record_bytes, val):\n    features = tf.io.parse_single_example(record_bytes, {\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'mask': tf.io.FixedLenFeature([], tf.string),\n        'organ': tf.io.FixedLenFeature([], tf.string),\n    })\n\n    image = tf.io.parse_tensor(features['image'], out_type=tf.uint8)\n    image = tf.reshape(image, [IMG0_SIZE, IMG0_SIZE, N_CHANNELS])\n    \n    mask = tf.io.parse_tensor(features['mask'], out_type=tf.uint8)\n    mask = tf.reshape(mask, [MASK0_SIZE, MASK0_SIZE, 1])\n    mask = tf.cast(mask, tf.float32)\n    \n    # Resize if validation\n    if val:\n        image = tf.image.resize(image, [IMG_SIZE, IMG_SIZE], method=tf.image.ResizeMethod.BICUBIC)\n        image = tf.cast(image, tf.uint8)\n            \n        mask = tf.image.resize(mask, [MASK_SIZE, MASK_SIZE], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n        mask = tf.cast(mask, tf.uint8)\n    \n    # Ogran\n    organ = features['organ']\n    \n    return image, mask, organ","metadata":{"papermill":{"duration":0.106667,"end_time":"2022-08-29T16:33:22.961715","exception":false,"start_time":"2022-08-29T16:33:22.855048","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:30:07.307134Z","iopub.execute_input":"2022-08-30T20:30:07.307504Z","iopub.status.idle":"2022-08-30T20:30:07.322495Z","shell.execute_reply.started":"2022-08-30T20:30:07.307472Z","shell.execute_reply":"2022-08-30T20:30:07.321550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_hubmap_scale(organ):\n    if organ == 'kidney':\n        scale = 0.5 / 0.4\n    elif organ == 'largeintestine':\n        scale = 0.2290 / 0.4\n    elif organ == 'spleen':\n        scale = 0.4945 / 0.4\n    elif organ == 'lung':\n        scale = 0.7562 / 0.4\n    elif organ == 'prostate':\n        scale = 6.263 / 0.4\n    else:\n        scale = 1.0\n        \n    return scale","metadata":{"execution":{"iopub.status.busy":"2022-08-30T20:42:36.057612Z","iopub.execute_input":"2022-08-30T20:42:36.058008Z","iopub.status.idle":"2022-08-30T20:42:36.064846Z","shell.execute_reply.started":"2022-08-30T20:42:36.057971Z","shell.execute_reply":"2022-08-30T20:42:36.063921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_image(image, mask, organ, val, return_organ):    \n    if not val:\n        rotations = tf_rand_int(0, 4, dtype=tf.int32)\n\n        image = tf.image.rot90(image, rotations)\n        mask = tf.image.rot90(mask, rotations)\n        \n        if one_in(2): # Transpose\n            image = tf.image.transpose(image)\n            mask = tf.image.transpose(mask)\n        \n        # Pixel Level Augmentations\n        if one_in(2): # Random HUE\n            image = tf.image.random_hue(image, 0.20)\n        if one_in(2): # Random Saturation\n            image = tf.image.random_saturation(image, 0.80, 1.20)\n        if one_in(2): # Random Contrast\n            image = tf.image.random_contrast(image, 0.80, 1.20)\n        if one_in(2): # Random Brightness\n            image = tf.image.random_brightness(image, 0.10)\n        if one_in(2): # Random JPEG Quality\n            image = tf.image.random_jpeg_quality(image, 75, 100)\n        if one_in(2): # Random Noise\n            noise = tf.random.uniform(INPUT_SHAPE0, -8, 8, dtype=tf.int32)\n            image = tf.cast(image, tf.int32) + noise\n            image = tf.cast(image / tf.reduce_max(image) * 255, tf.uint8)\n            \n        if one_in(2): # Rotate\n            angle = tf.random.uniform([], -45 * math.pi / 180, 45 * math.pi / 180, dtype=tf.float32)\n            y, idx, count = tf.unique_with_counts(tf.reshape(image, [-1]))\n            image = tfa.image.rotate(image, angle, interpolation='bilinear', fill_value=tf.cast(y[0], tf.float32))\n            mask = tfa.image.rotate(mask, angle, interpolation='nearest', fill_value=0)\n        \n        if one_in(2): # HuBMAP Scaling\n            scale = get_hubmap_scale(organ)\n            image_hubmap_size = tf.cast(IMG0_SIZE / scale, tf.int64)\n            mask_hubmap_size = tf.cast(MASK0_SIZE / scale, tf.int64)\n            \n            if scale > 1 and scale < 3: # pad [skip prostate]\n                # Downscale Image\n                image = tf.image.resize(image, [image_hubmap_size, image_hubmap_size], method=tf.image.ResizeMethod.BICUBIC)\n                image = tf.cast(image, tf.uint8)\n                mask = tf.image.resize(mask, [mask_hubmap_size, mask_hubmap_size], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n                \n                # padding is half scale, as padding as added on both sides\n                pad_x_l = tf_rand_int(0, IMG0_SIZE - image_hubmap_size)\n                pad_x_r = tf.cast(IMG0_SIZE - image_hubmap_size - pad_x_l, tf.int64)\n                pad_y_t = tf_rand_int(0, MASK0_SIZE - mask_hubmap_size)\n                pad_y_b = tf.cast(MASK0_SIZE - image_hubmap_size - pad_y_t, tf.int64)\n                \n                paddings = tf.reshape([[pad_x_l, pad_x_r], [pad_y_t, pad_y_b], [0,0]], [3,2])\n                paddings_mask = tf.cast(paddings * MASK0_SIZE / IMG0_SIZE, tf.int32)\n\n                y, idx, count = tf.unique_with_counts(tf.reshape(image, [-1]))\n                image = tf.pad(image, paddings, mode='CONSTANT', constant_values=y[0])\n                mask = tf.pad(mask, paddings_mask, mode='CONSTANT', constant_values=0)\n            elif scale < 1: # crop\n                image = tf.image.resize(image, [image_hubmap_size, image_hubmap_size], method=tf.image.ResizeMethod.BICUBIC)\n                image = tf.cast(image / tf.reduce_max(image) * 255, tf.uint8)\n                mask = tf.image.resize(mask, [mask_hubmap_size, mask_hubmap_size], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n                \n                offset_x = tf_rand_int(0, image_hubmap_size - IMG0_SIZE)\n                offset_y = tf_rand_int(0, image_hubmap_size - IMG0_SIZE)\n                img_size_crop = IMG0_SIZE - offset_x\n\n                offset_x_mask = tf.cast(offset_x * image_hubmap_size / mask_hubmap_size, tf.int32)\n                offset_y_mask = tf.cast(offset_y * image_hubmap_size / mask_hubmap_size, tf.int32)\n\n                image = tf.slice(image, [offset_x, offset_y, 0], [IMG0_SIZE, IMG0_SIZE, N_CHANNELS])\n                mask = tf.slice(mask, [offset_x_mask, offset_y_mask, 0], [MASK0_SIZE, MASK0_SIZE, 1])\n            \n        # Resize and Cast\n        image = tf.image.resize(image, [IMG_SIZE, IMG_SIZE], method=tf.image.ResizeMethod.BICUBIC)\n        image = tf.cast(image / tf.reduce_max(image) * 255, tf.uint8)\n            \n        mask = tf.image.resize(mask, [MASK_SIZE, MASK_SIZE], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n        mask = tf.cast(mask, tf.uint8)\n        \n        # Explicit Reshape for TPU\n        image = tf.reshape(image, [IMG_SIZE, IMG_SIZE, N_CHANNELS])\n        mask = tf.reshape(mask, [MASK_SIZE, MASK_SIZE, 1])\n\n    if return_organ:\n        return image, mask, organ\n    else:\n        return image, mask","metadata":{"papermill":{"duration":0.121018,"end_time":"2022-08-29T16:33:23.175895","exception":false,"start_time":"2022-08-29T16:33:23.054877","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:55:38.748744Z","iopub.execute_input":"2022-08-30T20:55:38.749416Z","iopub.status.idle":"2022-08-30T20:55:38.775441Z","shell.execute_reply.started":"2022-08-30T20:55:38.749377Z","shell.execute_reply":"2022-08-30T20:55:38.774607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_TPU:\n    GCS_DS_PATH = KaggleDatasets().get_gcs_path('hubmap-hpa-hacking-the-human-body-1024x1024')\n    print(f'GCS_DS_PATH: {GCS_DS_PATH}')\n\n    TFRECORDS_FILE_PATHS = np.array(tf.io.gfile.glob(f'{GCS_DS_PATH}/*.tfrecords'))\nelse:\n    TFRECORDS_FILE_PATHS = np.array(tf.io.gfile.glob('/kaggle/input/hubmap-hpa-hacking-the-human-body-1024x1024/*.tfrecords'))\n    \nTFRECORDS_FILE_PATHS = np.array(\n    sorted(TFRECORDS_FILE_PATHS, key=lambda fp: int(fp.split('.')[-2].split('_')[-1]))\n)\nprint(f'Found {len(TFRECORDS_FILE_PATHS)} TFRecords')","metadata":{"papermill":{"duration":0.642648,"end_time":"2022-08-29T16:33:23.911872","exception":false,"start_time":"2022-08-29T16:33:23.269224","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:55:39.493562Z","iopub.execute_input":"2022-08-30T20:55:39.494108Z","iopub.status.idle":"2022-08-30T20:55:39.950780Z","shell.execute_reply.started":"2022-08-30T20:55:39.494073Z","shell.execute_reply":"2022-08-30T20:55:39.949844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check whether file paths are correctly ordered\npd.options.display.max_colwidth = 999\ndisplay(pd.DataFrame(TFRECORDS_FILE_PATHS, columns=['File Path']).sample(10, random_state=SEED))","metadata":{"papermill":{"duration":0.119538,"end_time":"2022-08-29T16:33:24.129017","exception":false,"start_time":"2022-08-29T16:33:24.009479","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:55:39.952402Z","iopub.execute_input":"2022-08-30T20:55:39.952660Z","iopub.status.idle":"2022-08-30T20:55:39.964600Z","shell.execute_reply.started":"2022-08-30T20:55:39.952633Z","shell.execute_reply":"2022-08-30T20:55:39.963664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_SAMPLES_PER_TFRECORD = np.load('/kaggle/input/hubmap-hpa-hacking-the-human-body-1024x1024/N_SAMPLES_PER_TFRECORD.npy')\nprint(f'N_SAMPLES_PER_TFRECORD shape: {N_SAMPLES_PER_TFRECORD.shape}, dtype: {N_SAMPLES_PER_TFRECORD.dtype}')","metadata":{"papermill":{"duration":0.113377,"end_time":"2022-08-29T16:33:24.338370","exception":false,"start_time":"2022-08-29T16:33:24.224993","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:55:39.966325Z","iopub.execute_input":"2022-08-30T20:55:39.967127Z","iopub.status.idle":"2022-08-30T20:55:39.978333Z","shell.execute_reply.started":"2022-08-30T20:55:39.967082Z","shell.execute_reply":"2022-08-30T20:55:39.977617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataset(tfrecord_idxs=None, bs=BATCH_SIZE, idxs=None, return_steps=False, val=False, return_organ=False, debug=False, benchmark=False):\n    ignore_order = tf.data.Options()\n    ignore_order.experimental_deterministic = False\n    if IS_TPU:\n        npr = 1 if val else AUTO\n    else:\n        npr = 1 if val else cpu_count()\n    \n    if tfrecord_idxs is None:\n        dataset = tf.data.TFRecordDataset(\n            TFRECORDS_FILE_PATHS,\n            num_parallel_reads=npr,\n            compression_type='GZIP',\n        )\n    else:\n        dataset = tf.data.TFRecordDataset(\n            TFRECORDS_FILE_PATHS[tfrecord_idxs],\n            num_parallel_reads=npr,\n            compression_type='GZIP',\n        )\n    \n    # Decode Raw TFRecords\n    dataset = dataset.map(\n        lambda record_bytes: decode_image(record_bytes, val), num_parallel_calls=AUTO if IS_TPU else 1\n    )\n    \n    # Cache TFRecords\n    dataset = dataset.cache()\n    \n    if not val and not debug:\n        dataset = dataset.with_options(ignore_order)\n        dataset = dataset.shuffle(128)\n        dataset = dataset.repeat()\n        \n    if benchmark:\n        dataset = dataset.with_options(ignore_order)\n        dataset = dataset.repeat()\n    \n    # Augment\n    dataset = dataset.map(\n        lambda image, mask, organ: augment_image(image, mask, organ, val, return_organ), num_parallel_calls=npr\n    )\n\n    dataset = dataset.batch(bs, drop_remainder=True)\n    dataset = dataset.prefetch(AUTO)\n    \n    if return_steps:\n        return dataset, math.ceil(N_SAMPLES_PER_TFRECORD[tfrecord_idxs].sum() / bs)\n    else:\n        return dataset","metadata":{"papermill":{"duration":0.111914,"end_time":"2022-08-29T16:33:24.544010","exception":false,"start_time":"2022-08-29T16:33:24.432096","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:55:39.979701Z","iopub.execute_input":"2022-08-30T20:55:39.980698Z","iopub.status.idle":"2022-08-30T20:55:39.992423Z","shell.execute_reply.started":"2022-08-30T20:55:39.980654Z","shell.execute_reply":"2022-08-30T20:55:39.991678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Benchmark Dataset\nn_steps_per_epoch = 10 if IS_TPU else 2\nbenchmark_dataset(get_dataset(benchmark=True), n_steps_per_epoch=n_steps_per_epoch)","metadata":{"papermill":{"duration":38.572753,"end_time":"2022-08-29T16:34:03.213010","exception":false,"start_time":"2022-08-29T16:33:24.640257","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:53:21.164049Z","iopub.execute_input":"2022-08-30T20:53:21.164699Z","iopub.status.idle":"2022-08-30T20:53:41.658853Z","shell.execute_reply.started":"2022-08-30T20:53:21.164650Z","shell.execute_reply":"2022-08-30T20:53:41.657940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, labels = next(iter(get_dataset(debug=True)))\nprint(f'images shape: {images.shape}, labels shape: {labels.shape}')\nprint(f'images dtype: {images.dtype}, labels dtype: {labels.dtype}')\n\npercentiles = [0.01, 0.05, 0.10, 0.25, 0.40 ,0.50, 0.60, 0.75, 0.90, 0.95, 0.99]\ndisplay(pd.Series(images.numpy().flatten()).describe(percentiles=percentiles).to_frame(name='images').T)\ndisplay(pd.Series(labels.numpy().flatten()).value_counts().to_frame(name='label value counts').T)\ndisplay(pd.Series(labels.numpy().flatten()).value_counts(normalize=True).to_frame(name='label normalised value counts').T)","metadata":{"papermill":{"duration":7.024493,"end_time":"2022-08-29T16:34:10.339758","exception":false,"start_time":"2022-08-29T16:34:03.315265","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:53:41.661127Z","iopub.execute_input":"2022-08-30T20:53:41.661481Z","iopub.status.idle":"2022-08-30T20:53:47.997291Z","shell.execute_reply.started":"2022-08-30T20:53:41.661447Z","shell.execute_reply":"2022-08-30T20:53:47.996413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_batch(get_dataset(bs=32, debug=True, return_organ=True))","metadata":{"papermill":{"duration":15.684006,"end_time":"2022-08-29T16:34:26.141544","exception":false,"start_time":"2022-08-29T16:34:10.457538","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:58:18.143501Z","iopub.execute_input":"2022-08-30T20:58:18.143818Z","iopub.status.idle":"2022-08-30T20:58:31.469362Z","shell.execute_reply.started":"2022-08-30T20:58:18.143788Z","shell.execute_reply":"2022-08-30T20:58:31.467614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Weights Initialization","metadata":{"papermill":{"duration":0.27403,"end_time":"2022-08-29T16:34:26.716441","exception":false,"start_time":"2022-08-29T16:34:26.442411","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def weight_init_analysis():\n    # Get large batch\n    images, _ = next(iter(get_dataset(bs=8, debug=True)))\n    print(f'images shape: {images.shape}')\n    \n    model = get_model(file_path=None, debug=True)\n    \n    (\n        image_norm,\n        embedding, up6, up5, up4, up3, up2, up1,\n        dec0, x0, xs0, xs1, xs2, xs3, xs4,\n        dec1, dec1_up, dec1_concat,\n        dec2, dec2_up, dec2_concat,\n        dec3, dec3_up, dec3_concat,\n        dec4, dec4_up, dec4_concat,\n        dec_fpn, dec_fpn_xs, output\n    ) = model(images, training=False) \n\n    percentiles = [0.01, 0.05, 0.10, 0.25, 0.40, 0.50, 0.60, 0.75, 0.90, 0.95, 0.99]\n\n    for v, v_name in zip(\n        [image_norm, embedding, up6, up5, up4, up3, up2, up1, dec0, x0, xs0, xs1, xs2, xs3, xs4, dec1, dec1_up, dec1_concat, dec2, dec2_up, dec2_concat, dec3, dec3_up, dec3_concat, dec4, dec4_up, dec4_concat, dec_fpn, dec_fpn_xs, output],\n        ['image_norm', 'embedding', 'up6', 'up5', 'up4', 'up3', 'up2', 'up1', 'dec0', 'x0', 'xs0', 'xs1', 'xs2', 'xs3', 'xs4', 'dec1', 'dec1_up', 'dec1_concat', 'dec2', 'dec2_up', 'dec2_concat', 'dec3', 'dec3_up', 'dec3_concat', 'dec4', 'dec4_up', 'dec4_concat', 'dec_fpn', 'dec_fpn_xs', 'output'],\n    ):\n        display(pd.Series(v.numpy().flatten()).describe(percentiles=percentiles).astype(float).round(2).to_frame(name=v_name).T)\n        print('=' * 50)\n\n    # Histogram\n    plt.figure(figsize=(15,5))\n    plt.title('Logits Output', size=24)\n    pd.Series(output.numpy().flatten()).plot(kind='hist', bins=100)\n    plt.grid()\n    plt.show()\n\n    # Histogram\n    plt.figure(figsize=(15,5))\n    plt.title('Sigmoid Output', size=24)\n    pd.Series(tf.math.sigmoid(output).numpy().flatten()).plot(kind='hist', bins=100)\n    plt.xticks(np.arange(0.0, 1.1, 0.1))\n    plt.grid()\n    plt.show()\n    \ntf.keras.backend.clear_session()\ngc.collect()\n    \nweight_init_analysis()","metadata":{"papermill":{"duration":85.015209,"end_time":"2022-08-29T16:35:51.998146","exception":false,"start_time":"2022-08-29T16:34:26.982937","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:31:11.499289Z","iopub.execute_input":"2022-08-30T20:31:11.499781Z","iopub.status.idle":"2022-08-30T20:32:36.332920Z","shell.execute_reply.started":"2022-08-30T20:31:11.499749Z","shell.execute_reply":"2022-08-30T20:32:36.331815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"papermill":{"duration":2.573265,"end_time":"2022-08-29T16:35:54.962283","exception":false,"start_time":"2022-08-29T16:35:52.389018","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:36.334786Z","iopub.execute_input":"2022-08-30T20:32:36.335350Z","iopub.status.idle":"2022-08-30T20:32:38.772408Z","shell.execute_reply.started":"2022-08-30T20:32:36.335285Z","shell.execute_reply":"2022-08-30T20:32:38.771532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Rate Scheduler","metadata":{"papermill":{"duration":0.29434,"end_time":"2022-08-29T16:35:55.576288","exception":false,"start_time":"2022-08-29T16:35:55.281948","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def lrfn(current_step, num_warmup_steps, lr_max, num_cycles=0.50, num_training_steps=EPOCHS_WHOLE):\n    \n    if current_step < num_warmup_steps:\n        return lr_max * 0.50 ** (num_warmup_steps - current_step)\n    else:\n        progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n\n        return max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress))) * lr_max","metadata":{"papermill":{"duration":0.319434,"end_time":"2022-08-29T16:35:56.198398","exception":false,"start_time":"2022-08-29T16:35:55.878964","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:38.775375Z","iopub.execute_input":"2022-08-30T20:32:38.775653Z","iopub.status.idle":"2022-08-30T20:32:38.781782Z","shell.execute_reply.started":"2022-08-30T20:32:38.775623Z","shell.execute_reply":"2022-08-30T20:32:38.780953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_lr_schedule(lr_schedule, epochs):\n    fig = plt.figure(figsize=(20, 10))\n    plt.plot([None] + lr_schedule + [None])\n    # X Labels\n    x = np.arange(1, epochs + 1)\n    x_axis_labels = [i if epochs <= 40 or i % 5 == 0 or i == 1 else None for i in range(1, epochs + 1)]\n    plt.xlim([1, epochs])\n    plt.xticks(x, x_axis_labels) # set tick step to 1 and let x axis start at 1\n    \n    # Increase y-limit for better readability\n    plt.ylim([0, max(lr_schedule) * 1.1])\n    \n    # Title\n    schedule_info = f'start: {lr_schedule[0]:.1E}, max: {max(lr_schedule):.1E}, final: {lr_schedule[-1]:.1E}'\n    plt.title(f'Step Learning Rate Schedule, {schedule_info}', size=18, pad=12)\n    \n    # Plot Learning Rates\n    for x, val in enumerate(lr_schedule):\n        if epochs <= 40 or x % 5 == 0 or x is epochs - 1:\n            if x < len(lr_schedule) - 1:\n                if lr_schedule[x - 1] < val:\n                    ha = 'right'\n                else:\n                    ha = 'left'\n            elif x == 0:\n                ha = 'right'\n            else:\n                ha = 'left'\n            plt.plot(x + 1, val, 'o', color='black');\n            offset_y = (max(lr_schedule) - min(lr_schedule)) * 0.02\n            plt.annotate(f'{val:.1E}', xy=(x + 1, val + offset_y), size=12, ha=ha)\n    \n    plt.xlabel('Epoch', size=16, labelpad=5)\n    plt.ylabel('Learning Rate', size=16, labelpad=5)\n    plt.grid()\n    plt.show()\n\n# Learning rate for encoder\nLR_SCHEDULE = [lrfn(step, num_warmup_steps=3, lr_max=LR_MAX_WHOLE, num_cycles=0.50) for step in range(EPOCHS_WHOLE)]\nplot_lr_schedule(LR_SCHEDULE, epochs=EPOCHS_WHOLE)","metadata":{"papermill":{"duration":1.20514,"end_time":"2022-08-29T16:35:57.705541","exception":false,"start_time":"2022-08-29T16:35:56.500401","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:38.783246Z","iopub.execute_input":"2022-08-30T20:32:38.783498Z","iopub.status.idle":"2022-08-30T20:32:39.682141Z","shell.execute_reply.started":"2022-08-30T20:32:38.783462Z","shell.execute_reply":"2022-08-30T20:32:39.681146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=1)","metadata":{"papermill":{"duration":0.311862,"end_time":"2022-08-29T16:35:58.325047","exception":false,"start_time":"2022-08-29T16:35:58.013185","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:39.683382Z","iopub.execute_input":"2022-08-30T20:32:39.684203Z","iopub.status.idle":"2022-08-30T20:32:39.688565Z","shell.execute_reply.started":"2022-08-30T20:32:39.684143Z","shell.execute_reply":"2022-08-30T20:32:39.687693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Whole Model","metadata":{"papermill":{"duration":0.300862,"end_time":"2022-08-29T16:35:58.928266","exception":false,"start_time":"2022-08-29T16:35:58.627404","status":"completed"},"tags":[]}},{"cell_type":"code","source":"ORGAN_PER_TFRECORD = np.load('/kaggle/input/hubmap-hpa-hacking-the-human-body-1024x1024/ORGAN_PER_TFRECORD.npy')","metadata":{"papermill":{"duration":0.318013,"end_time":"2022-08-29T16:35:59.547381","exception":false,"start_time":"2022-08-29T16:35:59.229368","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:39.689713Z","iopub.execute_input":"2022-08-30T20:32:39.689962Z","iopub.status.idle":"2022-08-30T20:32:39.713111Z","shell.execute_reply.started":"2022-08-30T20:32:39.689934Z","shell.execute_reply":"2022-08-30T20:32:39.712121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Learning Rate Callback\nlr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=1)","metadata":{"papermill":{"duration":0.310948,"end_time":"2022-08-29T16:36:00.159837","exception":false,"start_time":"2022-08-29T16:35:59.848889","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:39.714662Z","iopub.execute_input":"2022-08-30T20:32:39.714926Z","iopub.status.idle":"2022-08-30T20:32:39.719791Z","shell.execute_reply.started":"2022-08-30T20:32:39.714895Z","shell.execute_reply":"2022-08-30T20:32:39.718874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WeightDecayCallback(tf.keras.callbacks.Callback):\n    def on_epoch_begin(self, epoch, logs=None):\n        model.optimizer.weight_decay = model.optimizer.learning_rate * 1e-2\n        lr = model.optimizer.learning_rate.numpy()\n        wd = model.optimizer.weight_decay.numpy()\n        print(f'learning rate: {lr:.1E}, weight deay: {wd:.1E}')","metadata":{"papermill":{"duration":0.312561,"end_time":"2022-08-29T16:36:00.774655","exception":false,"start_time":"2022-08-29T16:36:00.462094","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:39.721498Z","iopub.execute_input":"2022-08-30T20:32:39.721823Z","iopub.status.idle":"2022-08-30T20:32:39.730604Z","shell.execute_reply.started":"2022-08-30T20:32:39.721792Z","shell.execute_reply":"2022-08-30T20:32:39.729652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create KFOLDS\nFOLDS = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nHISTORIES = dict()\nVAL_PREDS = dict()\nIDXS = { 'train_idxs': [], 'val_idxs': [] }\nfor fold, (train_idxs, val_idxs) in enumerate(FOLDS.split(X=TFRECORDS_FILE_PATHS, y=ORGAN_PER_TFRECORD)):\n    # Only train selected folds\n    if fold not in TRAIN_FOLDS:\n        continue\n    \n    IDXS['train_idxs'].append(train_idxs)\n    IDXS['val_idxs'].append(val_idxs)\n    train_dataset, train_steps_per_epoch = get_dataset(tfrecord_idxs=train_idxs, bs=BATCH_SIZE, return_steps=True)\n    val_dataset, val_steps_per_epoch = get_dataset(tfrecord_idxs=val_idxs, bs=N_PATCHES_PER_IMAGE * REPLICAS, return_steps=True, val=True)\n    print('=' * 80)\n    print(f'FOLD {fold}, train_steps_per_epoch: {train_steps_per_epoch}, val_steps_per_epoch: {val_steps_per_epoch}')\n    print('=' * 80)\n    \n    tf.keras.backend.clear_session()\n    gc.collect()\n    \n    model = get_model(file_path=None, cnn_trainable=True, lr=LR_MAX_WHOLE, eps=1e-6, debug=False)\n    if fold == TRAIN_FOLDS[0]:\n        print(model.summary())\n    \n    HISTORIES[fold] = model.fit(\n        train_dataset,\n        steps_per_epoch = train_steps_per_epoch * COMBINE_EPOCHS,\n        validation_data = val_dataset,\n        epochs = EPOCHS_WHOLE,\n        verbose = 1 if os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive' else 2,\n        callbacks = [\n            lr_callback,\n        ],\n    )\n    \n    model.save_weights(f'model_{fold}.h5')\n    \n    print('\\n' * 3)\n    \n    # Validation Dataset\n    dataset = get_dataset(\n            tfrecord_idxs=val_idxs,\n            bs=len(val_idxs),\n            val=True,\n            return_organ=True,\n            debug=True,\n        )\n\n    # Validation images, masks and organs\n    val_images, val_masks, val_organs = next(iter(dataset))\n    # Cast from Tensorflow to Numpy\n    val_images = val_images.numpy()\n    val_masks = val_masks.numpy().astype(np.uint8)\n    val_organs = val_organs.numpy()\n    print(f'val_images shape: {val_images.shape}, val_masks shape: {val_masks.shape}, val_organs shape: {val_organs.shape}')\n\n    # Pad to Multiple of REPLICAS\n    val_images_len = val_images.shape[0]\n    pad = REPLICAS - (val_images_len % REPLICAS)\n    val_images_pad = np.pad(val_images, [(0,pad),(0,0),(0,0),(0,0)])\n    # Validation predictions\n    val_y_preds = model.predict(val_images_pad, verbose=1, batch_size=len(val_images_pad) if IS_TPU else 1)\n    print(f'val_y_preds shape: {val_y_preds.shape}, val_y_preds dtype: {val_y_preds.dtype}')\n    \n    VAL_PREDS[fold] = dict({\n        'val_images': np.copy(val_images),\n        'val_masks': np.copy(val_masks),\n        'val_organs': np.copy(val_organs),\n        'val_y_preds': np.copy(val_y_preds[:val_images_len]),\n    })\n    \n    # Clean Up\n    del model, train_dataset, val_dataset\n    del dataset, val_images, val_masks, val_organs, val_images_pad, val_y_preds\n    gc.collect()","metadata":{"papermill":{"duration":1613.324762,"end_time":"2022-08-29T17:02:54.400265","exception":false,"start_time":"2022-08-29T16:36:01.075503","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:32:39.732291Z","iopub.execute_input":"2022-08-30T20:32:39.732558Z","iopub.status.idle":"2022-08-30T20:40:43.372768Z","shell.execute_reply.started":"2022-08-30T20:32:39.732531Z","shell.execute_reply":"2022-08-30T20:40:43.371302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training History","metadata":{"papermill":{"duration":0.326244,"end_time":"2022-08-29T17:02:55.078157","exception":false,"start_time":"2022-08-29T17:02:54.751913","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def plot_history_metric(metric, f_best=np.argmax, ylim=None, yscale=None, yticks=None):\n    plt.figure(figsize=(20, 10))\n    \n    for fold, history in HISTORIES.items():\n        values = history.history[metric]\n        N_EPOCHS = len(values)\n        val = 'val' in ''.join(history.history.keys())\n        # Epoch Ticks\n        if N_EPOCHS <= 20:\n            x = np.arange(1, N_EPOCHS + 1)\n        else:\n            x = [1, 5] + [10 + 5 * idx for idx in range((N_EPOCHS - 10) // 5 + 1)]\n        \n        x_ticks = np.arange(1, N_EPOCHS+1)\n\n        # Validation\n        if val:\n            val_values = history.history[f'val_{metric}']\n            val_argmin = f_best(val_values)\n            plt.plot(x_ticks, val_values, label=f'val_fold_{fold}')\n\n        # summarize history for accuracy\n        plt.plot(x_ticks, values, label=f'train_fold_{fold}')\n        argmin = f_best(values)\n        plt.scatter(argmin + 1, values[argmin], color='red', s=75, marker='o', label=f'train_best_fold_{fold}')\n        if val:\n            plt.scatter(val_argmin + 1, val_values[val_argmin], color='purple', s=75, marker='o', label=f'val_best_fold_{fold}')\n\n    plt.title(f'Model {metric}', fontsize=24, pad=10)\n    plt.ylabel(metric, fontsize=20, labelpad=10)\n\n    if ylim:\n        plt.ylim(ylim)\n\n    if yscale is not None:\n        plt.yscale(yscale)\n        \n    if yticks is not None:\n        plt.yticks(yticks, fontsize=16)\n\n    plt.xlabel('epoch', fontsize=20, labelpad=10)        \n    plt.tick_params(axis='x', labelsize=8)\n    plt.xticks(x, fontsize=16) # set tick step to 1 and let x axis start at 1\n    plt.yticks(fontsize=16)\n    \n    plt.legend(prop={'size': 10})\n    plt.grid()\n    plt.show()","metadata":{"papermill":{"duration":0.341395,"end_time":"2022-08-29T17:02:55.739126","exception":false,"start_time":"2022-08-29T17:02:55.397731","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.373955Z","iopub.status.idle":"2022-08-30T20:40:43.374362Z","shell.execute_reply.started":"2022-08-30T20:40:43.374129Z","shell.execute_reply":"2022-08-30T20:40:43.374146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('loss', f_best=np.argmin, yscale='log')","metadata":{"papermill":{"duration":1.352547,"end_time":"2022-08-29T17:02:57.409488","exception":false,"start_time":"2022-08-29T17:02:56.056941","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.376198Z","iopub.status.idle":"2022-08-30T20:40:43.376565Z","shell.execute_reply.started":"2022-08-30T20:40:43.376373Z","shell.execute_reply":"2022-08-30T20:40:43.376397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('precision', ylim=(0,1), yticks=np.arange(0, 1.1, 0.1))","metadata":{"papermill":{"duration":1.025776,"end_time":"2022-08-29T17:02:58.781608","exception":false,"start_time":"2022-08-29T17:02:57.755832","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.377605Z","iopub.status.idle":"2022-08-30T20:40:43.378309Z","shell.execute_reply.started":"2022-08-30T20:40:43.378078Z","shell.execute_reply":"2022-08-30T20:40:43.378104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('recall', ylim=(0,1), yticks=np.arange(0, 1.1, 0.1))","metadata":{"papermill":{"duration":1.015414,"end_time":"2022-08-29T17:03:00.131952","exception":false,"start_time":"2022-08-29T17:02:59.116538","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.379280Z","iopub.status.idle":"2022-08-30T20:40:43.379926Z","shell.execute_reply.started":"2022-08-30T20:40:43.379717Z","shell.execute_reply":"2022-08-30T20:40:43.379742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('iou', ylim=(0,1), yticks=np.arange(0, 1.1, 0.1))","metadata":{"papermill":{"duration":1.004531,"end_time":"2022-08-29T17:03:01.478958","exception":false,"start_time":"2022-08-29T17:03:00.474427","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.381142Z","iopub.status.idle":"2022-08-30T20:40:43.381529Z","shell.execute_reply.started":"2022-08-30T20:40:43.381339Z","shell.execute_reply":"2022-08-30T20:40:43.381362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('auc', ylim=(0,1), yticks=np.arange(0, 1.1, 0.1))","metadata":{"papermill":{"duration":1.019851,"end_time":"2022-08-29T17:03:02.838029","exception":false,"start_time":"2022-08-29T17:03:01.818178","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.382625Z","iopub.status.idle":"2022-08-30T20:40:43.383292Z","shell.execute_reply.started":"2022-08-30T20:40:43.383053Z","shell.execute_reply":"2022-08-30T20:40:43.383075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('binary_accuracy', ylim=(0,1), yticks=np.arange(0, 1.1, 0.1))","metadata":{"papermill":{"duration":1.068654,"end_time":"2022-08-29T17:03:04.296154","exception":false,"start_time":"2022-08-29T17:03:03.227500","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.384574Z","iopub.status.idle":"2022-08-30T20:40:43.385109Z","shell.execute_reply.started":"2022-08-30T20:40:43.384906Z","shell.execute_reply":"2022-08-30T20:40:43.384928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Predictions","metadata":{"papermill":{"duration":0.339498,"end_time":"2022-08-29T17:03:04.984299","exception":false,"start_time":"2022-08-29T17:03:04.644801","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# # Validation Dataset\n# dataset = get_dataset(\n#         tfrecord_idxs=val_idxs,\n#         bs=len(val_idxs),\n#         val=True,\n#         return_organ=True,\n#         debug=True,\n#     )\n\n# # Validation images, masks and organs\n# val_images, val_masks, val_organs = next(iter(dataset))\n# # Cast from Tensorflow to Numpy\n# val_images = val_images.numpy()\n# val_masks = val_masks.numpy().astype(np.uint8)\n# val_organs = val_organs.numpy()\n# print(f'val_images shape: {val_images.shape}, val_masks shape: {val_masks.shape}, val_organs shape: {val_organs.shape}')\n\n# # Validation Predictions\n# VAL_Y_PREDS = model.predict(val_images, verbose=1, batch_size=len(val_idxs) if IS_TPU else 1)\n# print(f'VAL_Y_PREDS shape: {VAL_Y_PREDS.shape}, VAL_Y_PREDS dtype: {VAL_Y_PREDS.dtype}')","metadata":{"papermill":{"duration":0.376303,"end_time":"2022-08-29T17:03:05.705477","exception":false,"start_time":"2022-08-29T17:03:05.329174","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.386262Z","iopub.status.idle":"2022-08-30T20:40:43.386906Z","shell.execute_reply.started":"2022-08-30T20:40:43.386700Z","shell.execute_reply":"2022-08-30T20:40:43.386725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mean Intersection Over Union","metadata":{"papermill":{"duration":0.347772,"end_time":"2022-08-29T17:03:06.397886","exception":false,"start_time":"2022-08-29T17:03:06.050114","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def iou(y_true, y_pred):\n    intersection = np.count_nonzero(y_true * y_pred)\n    union = np.count_nonzero(y_true + y_pred)\n    return intersection / union","metadata":{"papermill":{"duration":0.35178,"end_time":"2022-08-29T17:03:07.102265","exception":false,"start_time":"2022-08-29T17:03:06.750485","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.388654Z","iopub.status.idle":"2022-08-30T20:40:43.389000Z","shell.execute_reply.started":"2022-08-30T20:40:43.388819Z","shell.execute_reply":"2022-08-30T20:40:43.388842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Predictions and true labels of validation dataset\ndef get_y_true_y_pred(val_pred):\n    thresholds = np.arange(0, 1.01, 0.01)\n    IoUs = {}\n    for t in thresholds:\n        IoUs[t] = []\n        \n    IoUsOrgans = {}\n    for o in ORGANS:\n        IoUsOrgans[o] = {}\n        for t in thresholds:\n            IoUsOrgans[o][t] = []\n    \n    for idx, image in enumerate(tqdm(val_pred['val_images'])):\n        y_true = val_pred['val_masks'][idx]\n        organ = val_pred['val_organs'][idx]\n        y_pred = val_pred['val_y_preds'][idx]\n        \n        if idx == 0:\n            print(f'image shape: {image.shape}, y_true shape: {y_true.shape}')\n            print(f'organs: {organ.decode()}, y_pred shape: {y_pred.shape}')\n        \n        # Compute IoU for each threshold\n        o = organ.decode()\n        for t in thresholds:\n            IoU = iou(y_true, (y_pred > t).astype(np.int8))\n            IoUs[t].append(IoU)\n            IoUsOrgans[o][t].append(IoU)\n    \n    return IoUs, IoUsOrgans\n\nIoU_Folds = dict()\nfor fold, v in VAL_PREDS.items():\n    IoUs, IoUsOrgans = get_y_true_y_pred(v)\n    IoU_Folds[fold] = {\n        'IoUs': IoUs,\n        'IoUsOrgans': IoUsOrgans,\n    }","metadata":{"papermill":{"duration":16.206749,"end_time":"2022-08-29T17:03:23.658326","exception":false,"start_time":"2022-08-29T17:03:07.451577","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.390220Z","iopub.status.idle":"2022-08-30T20:40:43.390571Z","shell.execute_reply.started":"2022-08-30T20:40:43.390380Z","shell.execute_reply":"2022-08-30T20:40:43.390405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mean IoU at Threshold","metadata":{"papermill":{"duration":0.352763,"end_time":"2022-08-29T17:03:24.368013","exception":false,"start_time":"2022-08-29T17:03:24.015250","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def plot_iou_by_threshold(ious, name):\n    thresholds = list(ious.keys())\n    MeanIoUs = [np.mean(v)for v in ious.values()]\n\n    plt.figure(figsize=(12,8))\n    plt.title(f'Mean IoU by Threshold {name}', size=24)\n    plt.plot(thresholds, MeanIoUs)\n    plt.grid()\n    plt.xlabel('Threshold', size=16)\n    plt.ylabel('Mean IoU', size=16)\n    plt.xticks(size=12)\n    plt.yticks(size=12)\n    plt.ylim(0,1)\n\n    # Best Threshold\n    arg_best = np.argmax(MeanIoUs)\n    threshold_best = thresholds[arg_best]\n    mean_iou_best = MeanIoUs[arg_best]\n    plt.scatter(threshold_best, mean_iou_best, color='red', s=100, marker='o', label=f'Best Mean IoU ({mean_iou_best:.3f}) at Threshold {threshold_best:.3f}')\n    plt.legend(prop={'size': 16})\n\n    plt.show()\n    \n    # Save Best Threshold\n    np.save(f'threshold_best_{name}.npy', threshold_best)\n    \n    return threshold_best","metadata":{"papermill":{"duration":0.360155,"end_time":"2022-08-29T17:03:25.074366","exception":false,"start_time":"2022-08-29T17:03:24.714211","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.391792Z","iopub.status.idle":"2022-08-30T20:40:43.392142Z","shell.execute_reply.started":"2022-08-30T20:40:43.391960Z","shell.execute_reply":"2022-08-30T20:40:43.391984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Global Mean Intersection over Union at Threshold\nfor fold, v in IoU_Folds.items():\n    print('=' * 80)\n    print(f'FOLD {fold}')\n    print('=' * 80)\n    v['threshold_best'] = plot_iou_by_threshold(v['IoUs'], f'all_{fold}')","metadata":{"papermill":{"duration":0.799037,"end_time":"2022-08-29T17:03:26.212376","exception":false,"start_time":"2022-08-29T17:03:25.413339","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.393377Z","iopub.status.idle":"2022-08-30T20:40:43.393732Z","shell.execute_reply.started":"2022-08-30T20:40:43.393546Z","shell.execute_reply":"2022-08-30T20:40:43.393571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Per Organ Mean Intersection over Union at Threshold\nfor fold, v in IoU_Folds.items():\n    print('=' * 80)\n    print(f'FOLD {fold}')\n    print('=' * 80)\n    for organ, ious in v['IoUsOrgans'].items():\n        plot_iou_by_threshold(ious, f'{organ}_{fold}')","metadata":{"papermill":{"duration":2.326485,"end_time":"2022-08-29T17:03:28.896841","exception":false,"start_time":"2022-08-29T17:03:26.570356","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.394939Z","iopub.status.idle":"2022-08-30T20:40:43.395325Z","shell.execute_reply.started":"2022-08-30T20:40:43.395103Z","shell.execute_reply":"2022-08-30T20:40:43.395126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# IoU Distribution Best Threshold","metadata":{"papermill":{"duration":0.359241,"end_time":"2022-08-29T17:03:29.647069","exception":false,"start_time":"2022-08-29T17:03:29.287828","status":"completed"},"tags":[]}},{"cell_type":"code","source":"for fold, v in IoU_Folds.items():\n    print('=' * 80)\n    print(f'FOLD {fold}')\n    print('=' * 80)\n    \n    percentiles = [0.01, 0.05, 0.10, 0.25, 0.40, 0.50, 0.60, 0.75, 0.90, 0.95, 0.99]\n    s = v['IoUs'][v['threshold_best']]\n    \n    display(pd.Series(s).describe(percentiles=percentiles).apply(lambda v: f'{v:.2f}').to_frame(name='Value').T)","metadata":{"papermill":{"duration":0.389245,"end_time":"2022-08-29T17:03:30.397172","exception":false,"start_time":"2022-08-29T17:03:30.007927","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.396499Z","iopub.status.idle":"2022-08-30T20:40:43.397029Z","shell.execute_reply.started":"2022-08-30T20:40:43.396835Z","shell.execute_reply":"2022-08-30T20:40:43.396859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold, v in IoU_Folds.items():\n    print('=' * 80)\n    print(f'FOLD {fold}')\n    print('=' * 80)\n    plt.figure(figsize=(12,8))\n    pd.Series(v['IoUs'][v['threshold_best']]).plot(kind='hist')\n    plt.title('IoU Distribution at Best Threshold', size=24)\n    plt.grid()\n    plt.xlabel('Threshold', size=16)\n    plt.ylabel('Count', size=16)\n    plt.xticks(size=12)\n    plt.yticks(size=12)\n    plt.xlim(0,1)\n    plt.show()","metadata":{"papermill":{"duration":0.708725,"end_time":"2022-08-29T17:03:31.462186","exception":false,"start_time":"2022-08-29T17:03:30.753461","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.398189Z","iopub.status.idle":"2022-08-30T20:40:43.398547Z","shell.execute_reply.started":"2022-08-30T20:40:43.398355Z","shell.execute_reply":"2022-08-30T20:40:43.398380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation Prediction Visualization","metadata":{"papermill":{"duration":0.360464,"end_time":"2022-08-29T17:03:32.188744","exception":false,"start_time":"2022-08-29T17:03:31.828280","status":"completed"},"tags":[]}},{"cell_type":"code","source":"gc.collect()","metadata":{"papermill":{"duration":2.350333,"end_time":"2022-08-29T17:03:34.896402","exception":false,"start_time":"2022-08-29T17:03:32.546069","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.399633Z","iopub.status.idle":"2022-08-30T20:40:43.399978Z","shell.execute_reply.started":"2022-08-30T20:40:43.399799Z","shell.execute_reply":"2022-08-30T20:40:43.399820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def merge_patches(patches):\n    if len(patches.shape) == 3:\n        return patches\n    \n    image = np.zeros(shape=[IMG_SIZE_FULL, IMG_SIZE_FULL, patches.shape[-1]], dtype=patches.dtype)\n    s = int(N_PATCHES_PER_IMAGE ** 0.50)\n    for r in range(s):\n        for c in range(s):\n            start_x = r * IMG_SIZE\n            end_x = (r + 1) * IMG_SIZE\n            start_y = c * IMG_SIZE\n            end_y = (c + 1) * IMG_SIZE\n            image[start_x:end_x, start_y:end_y] = patches[r * s + c]\n            \n    return image","metadata":{"papermill":{"duration":0.369675,"end_time":"2022-08-29T17:03:35.624793","exception":false,"start_time":"2022-08-29T17:03:35.255118","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.401029Z","iopub.status.idle":"2022-08-30T20:40:43.401381Z","shell.execute_reply.started":"2022-08-30T20:40:43.401205Z","shell.execute_reply":"2022-08-30T20:40:43.401223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_validation_predictions(v, threshold_best, n):\n    for idx in range(n):\n        image = v['val_images'][idx]\n        y_true = v['val_masks'][idx]\n        organ = v['val_organs'][idx]\n        y_pred = v['val_y_preds'][idx]\n            \n        organ = organ.decode()\n        # Predicted Mask\n        y_pred = merge_patches(y_pred)\n        y_pred_binary = (y_pred > threshold_best).astype(np.uint8)\n        # Merge Image and Label patches\n        image = merge_patches(image)\n        y_true = merge_patches(y_true)\n        # Red = False Positive\n        r = ((y_pred_binary == 1) * (y_true == 0)).astype(np.uint8) \n        # Green = True Positive\n        g = ((y_pred_binary == 1) * (y_true == 1)).astype(np.uint8)\n        # Blue = False Negative\n        b = ((y_pred_binary == 0) * (y_true == 1)).astype(np.uint8)\n        # Error Visualization Using RGB\n        mask_error = np.stack((r, g, b), axis=2).squeeze() * 255\n\n        fig, axes = plt.subplots(1, 5, figsize=(20, 4))\n        axes[0].imshow(image)\n        axes[0].set_title(f'Image {organ} IoU: {IoUs[threshold_best][idx]:.2f}')\n        axes[1].imshow(y_true)\n        axes[1].set_title('Mask True')\n        axes[2].imshow(y_pred)\n        axes[2].set_title('Mask Pred')\n        axes[3].imshow(y_pred_binary)\n        axes[3].set_title('Mask Pred Binary')\n        axes[4].imshow(mask_error)\n        axes[4].set_title('Mask Error')\n\n        plt.show()","metadata":{"papermill":{"duration":0.380373,"end_time":"2022-08-29T17:03:36.368820","exception":false,"start_time":"2022-08-29T17:03:35.988447","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.403772Z","iopub.status.idle":"2022-08-30T20:40:43.404288Z","shell.execute_reply.started":"2022-08-30T20:40:43.404004Z","shell.execute_reply":"2022-08-30T20:40:43.404027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()\nfor fold, v in VAL_PREDS.items():\n    print('=' * 80)\n    print(f'FOLD {fold}')\n    print('=' * 80)\n    threshold_best = IoU_Folds[fold]['threshold_best']\n    plot_validation_predictions(v, threshold_best, 32)","metadata":{"papermill":{"duration":38.599073,"end_time":"2022-08-29T17:04:15.326117","exception":false,"start_time":"2022-08-29T17:03:36.727044","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-30T20:40:43.405820Z","iopub.status.idle":"2022-08-30T20:40:43.406337Z","shell.execute_reply.started":"2022-08-30T20:40:43.406057Z","shell.execute_reply":"2022-08-30T20:40:43.406081Z"},"trusted":true},"execution_count":null,"outputs":[]}]}