{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-28T09:06:21.171480Z","iopub.execute_input":"2022-08-28T09:06:21.172475Z","iopub.status.idle":"2022-08-28T09:06:21.697458Z","shell.execute_reply.started":"2022-08-28T09:06:21.172360Z","shell.execute_reply":"2022-08-28T09:06:21.696648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:21.698985Z","iopub.execute_input":"2022-08-28T09:06:21.699232Z","iopub.status.idle":"2022-08-28T09:06:21.703140Z","shell.execute_reply.started":"2022-08-28T09:06:21.699204Z","shell.execute_reply":"2022-08-28T09:06:21.702258Z"},"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 *\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'python version: P{sys.version}')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:21.704618Z","iopub.execute_input":"2022-08-28T09:06:21.704926Z","iopub.status.idle":"2022-08-28T09:06:30.263548Z","shell.execute_reply.started":"2022-08-28T09:06:21.704888Z","shell.execute_reply":"2022-08-28T09:06:30.262784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:30.265188Z","iopub.execute_input":"2022-08-28T09:06:30.265467Z","iopub.status.idle":"2022-08-28T09:06:36.010668Z","shell.execute_reply.started":"2022-08-28T09:06:30.265438Z","shell.execute_reply":"2022-08-28T09:06:36.009612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GCS_DS_PATH = KaggleDatasets().get_gcs_path('hubmap-patched-tfrecords-300x300')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.015917Z","iopub.execute_input":"2022-08-28T09:06:36.016354Z","iopub.status.idle":"2022-08-28T09:06:36.415214Z","shell.execute_reply.started":"2022-08-28T09:06:36.016286Z","shell.execute_reply":"2022-08-28T09:06:36.414429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 43\nDEBUG = False\n\n# Image dimensions\nIMG_SIZE_FULL = 640\nIMG_SIZE = 640\nN_PATCHES_PER_IMAGE = (IMG_SIZE_FULL // IMG_SIZE) ** 2\nN_CHANNELS = 3\nINPUT_SHAPE = (IMG_SIZE, IMG_SIZE, N_CHANNELS)\n\n# EfficientNet version, b0/b1/b2/b3/b4/b5/b6/b7\nEFN_SIZE = 'b7'\n# Peak Learning Rate\nLR_MAX_WHOLE = 8e-4\nN_FOLDS = 4\nTRAIN_FOLDS = [0]\n\n# Epochs are run 10 at a time due to low training samples\nCOMBINE_EPOCHS = 10\nEPOCHS_WHOLE = 200 // COMBINE_EPOCHS\n\n# Batch size\nBATCH_SIZE = 4 * REPLICAS","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.416447Z","iopub.execute_input":"2022-08-28T09:06:36.416697Z","iopub.status.idle":"2022-08-28T09:06:36.422758Z","shell.execute_reply.started":"2022-08-28T09:06:36.416669Z","shell.execute_reply":"2022-08-28T09:06:36.421974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dataset Mean and Standard Deviation\nMEAN = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/MEAN.npy')\nSTD = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.423946Z","iopub.execute_input":"2022-08-28T09:06:36.424576Z","iopub.status.idle":"2022-08-28T09:06:36.449602Z","shell.execute_reply.started":"2022-08-28T09:06:36.424539Z","shell.execute_reply":"2022-08-28T09:06:36.448957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'MEAN: {MEAN}, STD: {STD}')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.450818Z","iopub.execute_input":"2022-08-28T09:06:36.451675Z","iopub.status.idle":"2022-08-28T09:06:36.456411Z","shell.execute_reply.started":"2022-08-28T09:06:36.451641Z","shell.execute_reply":"2022-08-28T09:06:36.455595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.457558Z","iopub.execute_input":"2022-08-28T09:06:36.457951Z","iopub.status.idle":"2022-08-28T09:06:36.468778Z","shell.execute_reply.started":"2022-08-28T09:06:36.457919Z","shell.execute_reply":"2022-08-28T09:06:36.468153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.471183Z","iopub.execute_input":"2022-08-28T09:06:36.471607Z","iopub.status.idle":"2022-08-28T09:06:36.870169Z","shell.execute_reply.started":"2022-08-28T09:06:36.471557Z","shell.execute_reply":"2022-08-28T09:06:36.869292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.871510Z","iopub.execute_input":"2022-08-28T09:06:36.872221Z","iopub.status.idle":"2022-08-28T09:06:36.877718Z","shell.execute_reply.started":"2022-08-28T09:06:36.872188Z","shell.execute_reply":"2022-08-28T09:06:36.876903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import 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    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    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    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    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    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)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.879032Z","iopub.execute_input":"2022-08-28T09:06:36.879308Z","iopub.status.idle":"2022-08-28T09:06:36.913486Z","shell.execute_reply.started":"2022-08-28T09:06:36.879280Z","shell.execute_reply":"2022-08-28T09:06:36.912792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEMANTIC_LOSS_FUNCTIONS = Semantic_loss_functions()\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.914580Z","iopub.execute_input":"2022-08-28T09:06:36.914959Z","iopub.status.idle":"2022-08-28T09:06:36.929058Z","shell.execute_reply.started":"2022-08-28T09:06:36.914928Z","shell.execute_reply":"2022-08-28T09:06:36.928224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def FPN(xs, output_channels, last_layer, debug=False):\n    def _conv(x):\n        x = tf.keras.layers.ZeroPadding2D(padding=1)(x)\n        x = tf.keras.layers.Conv2D(output_channels * 2, 3, padding='SAME', kernel_initializer='he_normal', activation='relu')(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.ZeroPadding2D(padding=1)(x)\n        x = tf.keras.layers.Conv2D(output_channels, 3, padding='SAME', kernel_initializer='he_normal')(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()([_conv(x) for x in xs])\n    x = tf.keras.layers.Concatenate()([xs, last_layer])\n\n    if debug:\n        return x, xs\n    else:\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.930418Z","iopub.execute_input":"2022-08-28T09:06:36.930675Z","iopub.status.idle":"2022-08-28T09:06:36.939297Z","shell.execute_reply.started":"2022-08-28T09:06:36.930630Z","shell.execute_reply":"2022-08-28T09:06:36.938696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 if IS_TPU else groups,\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    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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.940522Z","iopub.execute_input":"2022-08-28T09:06:36.940936Z","iopub.status.idle":"2022-08-28T09:06:36.955202Z","shell.execute_reply.started":"2022-08-28T09:06:36.940906Z","shell.execute_reply":"2022-08-28T09:06:36.954516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.956382Z","iopub.execute_input":"2022-08-28T09:06:36.957289Z","iopub.status.idle":"2022-08-28T09:06:36.971520Z","shell.execute_reply.started":"2022-08-28T09:06:36.957253Z","shell.execute_reply":"2022-08-28T09:06:36.970683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def upsample(x, concat, target_filters, name, conv2dt_kernel_init_max, relu=True, dropout=0, debug=False):\n    filters = concat.shape[-1]\n    x_up = tf.keras.layers.Conv2DTranspose(\n            filters, # Number of Convolutional Filters\n            kernel_size=4, # Kernel Size\n            strides=2, # Kernel Steps\n            padding='SAME', # linear scaling\n            name=f'Conv2DTranspose_{name}', # Name of Layer\n            kernel_initializer='he_uniform',\n            use_bias=False,\n        )(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# 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    filters = concat.shape[-1]\n    x_up = tf.keras.layers.Conv2DTranspose(\n            filters, # Number of Convolutional Filters\n            kernel_size=4, # Kernel Size\n            strides=2, # Kernel Steps\n            padding='SAME', # linear scaling\n            name=f'Conv2DTranspose_{name}', # Name of Layer\n            kernel_initializer='he_uniform',\n            use_bias=False,\n        )(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":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.972920Z","iopub.execute_input":"2022-08-28T09:06:36.973182Z","iopub.status.idle":"2022-08-28T09:06:36.986692Z","shell.execute_reply.started":"2022-08-28T09:06:36.973153Z","shell.execute_reply":"2022-08-28T09:06:36.985722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GCS_WEIGHTS_PATH = KaggleDatasets().get_gcs_path('efficientnetv2-head-1x1-endpoint')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:36.988292Z","iopub.execute_input":"2022-08-28T09:06:36.989223Z","iopub.status.idle":"2022-08-28T09:06:37.352337Z","shell.execute_reply.started":"2022-08-28T09:06:36.989177Z","shell.execute_reply":"2022-08-28T09:06:37.351504Z"},"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        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        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, 32, '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, 64, '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.Dropout(0.10)(dec_fpn)\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        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        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\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:37.353623Z","iopub.execute_input":"2022-08-28T09:06:37.353858Z","iopub.status.idle":"2022-08-28T09:06:37.375806Z","shell.execute_reply.started":"2022-08-28T09:06:37.353831Z","shell.execute_reply":"2022-08-28T09:06:37.375137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf.keras.backend.clear_session()\ngc.collect()\n    \nmodel = get_model(file_path=None, debug=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:06:37.376773Z","iopub.execute_input":"2022-08-28T09:06:37.377205Z","iopub.status.idle":"2022-08-28T09:08:48.268664Z","shell.execute_reply.started":"2022-08-28T09:06:37.377175Z","shell.execute_reply":"2022-08-28T09:08:48.267687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:48.271350Z","iopub.execute_input":"2022-08-28T09:08:48.271669Z","iopub.status.idle":"2022-08-28T09:08:48.374696Z","shell.execute_reply.started":"2022-08-28T09:08:48.271630Z","shell.execute_reply":"2022-08-28T09:08:48.373789Z"},"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=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:48.376384Z","iopub.execute_input":"2022-08-28T09:08:48.376693Z","iopub.status.idle":"2022-08-28T09:08:52.298816Z","shell.execute_reply.started":"2022-08-28T09:08:48.376653Z","shell.execute_reply":"2022-08-28T09:08:52.297466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    for epoch_num in range(num_epochs):\n        for idx, (images, labels) in enumerate(dataset.take(n_steps_per_epoch + 1)):\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":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.300576Z","iopub.execute_input":"2022-08-28T09:08:52.300854Z","iopub.status.idle":"2022-08-28T09:08:52.308780Z","shell.execute_reply.started":"2022-08-28T09:08:52.300820Z","shell.execute_reply":"2022-08-28T09:08:52.307810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_batch(dataset, rows=4, cols=4):\n    imgs, lbls = next(iter(dataset))\n    imgs = imgs.numpy()\n    # Plot\n    fig, axes = plt.subplots(nrows=rows, ncols=cols, figsize=(rows*4, cols*4))\n    for r in range(rows):\n        for c in range(cols // 2):\n            img = imgs[r*cols+c]\n            axes[r, c*2].imshow(img)\n            axes[r, c*2].set_title(f'std: {img.std():.1f}')\n            lbl = lbls[r*cols+c]\n            axes[r, c*2+1].imshow(lbl)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.310722Z","iopub.execute_input":"2022-08-28T09:08:52.310925Z","iopub.status.idle":"2022-08-28T09:08:52.325356Z","shell.execute_reply.started":"2022-08-28T09:08:52.310902Z","shell.execute_reply":"2022-08-28T09:08:52.324176Z"},"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, [IMG_SIZE, IMG_SIZE, N_CHANNELS])\n    \n    mask = tf.io.parse_tensor(features['mask'], out_type=tf.uint8)\n    mask = tf.reshape(mask, [IMG_SIZE, IMG_SIZE, 1])\n    \n    # Ogran\n    organ = features['organ']\n    \n    return image, mask, organ","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.326678Z","iopub.execute_input":"2022-08-28T09:08:52.326955Z","iopub.status.idle":"2022-08-28T09:08:52.339833Z","shell.execute_reply.started":"2022-08-28T09:08:52.326872Z","shell.execute_reply":"2022-08-28T09:08:52.338829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_image(image, mask, organ, val, return_organ):\n    \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):\n            image = tf.image.transpose(image)\n            mask = tf.image.transpose(mask)\n        \n        # Pixel Level Augmentations\n        if one_in(2):\n            image = tf.image.random_hue(image, 0.2)\n        if one_in(2):\n            image = tf.image.random_saturation(image, 0.80, 1.20)\n        if one_in(2):\n            image = tf.image.random_contrast(image, 0.80, 1.20)\n        if one_in(2):\n            image = tf.image.random_brightness(image, 0.10)\n        if one_in(2):\n            image = tf.image.random_jpeg_quality(image, 75, 100)\n        \n        # Random Crop\n        offset_x = tf.random.uniform([], 0, tf.cast(IMG_SIZE * 0.50, tf.int32), dtype=tf.int32)\n        img_size_crop = IMG_SIZE - offset_x\n        if offset_x > 0:\n            offset_y = tf.random.uniform([], 0, offset_x, dtype=tf.int32)\n        else:\n            offset_y = tf.constant(0, dtype=tf.int32)\n\n        # Crop\n        if one_in(2):\n            image = tf.slice(image, [offset_x, offset_y, 0], [img_size_crop, img_size_crop, N_CHANNELS])\n            mask = tf.slice(mask, [offset_x, offset_y, 0], [img_size_crop, img_size_crop, 1])\n            \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            mask = tf.image.resize(mask, [IMG_SIZE, IMG_SIZE], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n        \n        # Rotate\n        if one_in(2):\n            angle = tf.random.uniform([], -45 * math.pi / 180, 45 * math.pi / 180, dtype=tf.float32)\n            image = tfa.image.rotate(image, angle, interpolation='bilinear', fill_mode='reflect')\n            mask = tfa.image.rotate(mask, angle, interpolation='nearest', fill_mode='reflect')\n        \n        # Resize\n        image = tf.cast(image / tf.reduce_max(image) * 255, tf.uint8)\n        \n        # Explicit Reshape for TPU\n        image = tf.reshape(image, [IMG_SIZE, IMG_SIZE, N_CHANNELS])\n        mask = tf.reshape(mask, [IMG_SIZE, IMG_SIZE, 1])\n\n    if return_organ:\n        return image, mask, organ\n    else:\n        return image, mask\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.341278Z","iopub.execute_input":"2022-08-28T09:08:52.341508Z","iopub.status.idle":"2022-08-28T09:08:52.358887Z","shell.execute_reply.started":"2022-08-28T09:08:52.341482Z","shell.execute_reply":"2022-08-28T09:08:52.357551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GCS_DS_PATH = KaggleDatasets().get_gcs_path('hubmap-patched-tfrecords-300x300')\nprint(f'GCS_DS_PATH: {GCS_DS_PATH}')\n\nTFRECORDS_FILE_PATHS = np.array(tf.io.gfile.glob(f'{GCS_DS_PATH}/*.tfrecords'))\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":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.360157Z","iopub.execute_input":"2022-08-28T09:08:52.360476Z","iopub.status.idle":"2022-08-28T09:08:52.895939Z","shell.execute_reply.started":"2022-08-28T09:08:52.360445Z","shell.execute_reply":"2022-08-28T09:08:52.895023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.options.display.max_colwidth = 999\ndisplay(pd.DataFrame(TFRECORDS_FILE_PATHS, columns=['File Path']).sample(10, random_state=SEED))","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.900901Z","iopub.execute_input":"2022-08-28T09:08:52.901600Z","iopub.status.idle":"2022-08-28T09:08:52.916119Z","shell.execute_reply.started":"2022-08-28T09:08:52.901566Z","shell.execute_reply":"2022-08-28T09:08:52.915029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_SAMPLES_PER_TFRECORD = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/N_SAMPLES_PER_TFRECORD.npy')\nprint(f'N_SAMPLES_PER_TFRECORD shape: {N_SAMPLES_PER_TFRECORD.shape}, dtype: {N_SAMPLES_PER_TFRECORD.dtype}')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.917361Z","iopub.execute_input":"2022-08-28T09:08:52.917608Z","iopub.status.idle":"2022-08-28T09:08:52.931766Z","shell.execute_reply.started":"2022-08-28T09:08:52.917579Z","shell.execute_reply":"2022-08-28T09:08:52.931176Z"},"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):\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(TFRECORDS_FILE_PATHS, num_parallel_reads=npr)\n    else:\n        dataset = tf.data.TFRecordDataset(TFRECORDS_FILE_PATHS[tfrecord_idxs], num_parallel_reads=npr)\n    \n    dataset = dataset.map(\n        lambda record_bytes: decode_image(record_bytes, val), num_parallel_calls=AUTO if IS_TPU else 1\n    )\n    if not debug:\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    # 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":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.933762Z","iopub.execute_input":"2022-08-28T09:08:52.934398Z","iopub.status.idle":"2022-08-28T09:08:52.945426Z","shell.execute_reply.started":"2022-08-28T09:08:52.934356Z","shell.execute_reply":"2022-08-28T09:08:52.944477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"benchmark_dataset(get_dataset(debug=True))\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:08:52.947338Z","iopub.execute_input":"2022-08-28T09:08:52.947680Z","iopub.status.idle":"2022-08-28T09:09:25.365054Z","shell.execute_reply.started":"2022-08-28T09:08:52.947613Z","shell.execute_reply":"2022-08-28T09:09:25.363911Z"},"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='labels').T)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:09:25.366646Z","iopub.execute_input":"2022-08-28T09:09:25.367051Z","iopub.status.idle":"2022-08-28T09:09:29.772025Z","shell.execute_reply.started":"2022-08-28T09:09:25.367019Z","shell.execute_reply":"2022-08-28T09:09:29.771411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show what we will be training on\nshow_batch(get_dataset(bs=16, debug=True))\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:09:29.772989Z","iopub.execute_input":"2022-08-28T09:09:29.773583Z","iopub.status.idle":"2022-08-28T09:09:35.663151Z","shell.execute_reply.started":"2022-08-28T09:09:29.773554Z","shell.execute_reply":"2022-08-28T09:09:35.662168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    tf.keras.backend.clear_session()\n    gc.collect()\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    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    \nweight_init_analysis()","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:09:35.664531Z","iopub.execute_input":"2022-08-28T09:09:35.664798Z","iopub.status.idle":"2022-08-28T09:10:50.893998Z","shell.execute_reply.started":"2022-08-28T09:09:35.664767Z","shell.execute_reply":"2022-08-28T09:10:50.893012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:10:50.895535Z","iopub.execute_input":"2022-08-28T09:10:50.895781Z","iopub.status.idle":"2022-08-28T09:10:50.902634Z","shell.execute_reply.started":"2022-08-28T09:10:50.895754Z","shell.execute_reply":"2022-08-28T09:10:50.901528Z"},"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()\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":{"execution":{"iopub.status.busy":"2022-08-28T09:10:50.904336Z","iopub.execute_input":"2022-08-28T09:10:50.904677Z","iopub.status.idle":"2022-08-28T09:10:51.687986Z","shell.execute_reply.started":"2022-08-28T09:10:50.904637Z","shell.execute_reply":"2022-08-28T09:10:51.687309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ORGAN_PER_TFRECORD = np.load('/kaggle/input/hubmap-patched-tfrecords-300x300/ORGAN_PER_TFRECORD.npy')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:11:08.988948Z","iopub.execute_input":"2022-08-28T09:11:08.989950Z","iopub.status.idle":"2022-08-28T09:11:09.005148Z","shell.execute_reply.started":"2022-08-28T09:11:08.989883Z","shell.execute_reply":"2022-08-28T09:11:09.003839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_callback = tf.keras.callbacks.LearningRateScheduler(lambda step: LR_SCHEDULE[step], verbose=1)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:11:19.227628Z","iopub.execute_input":"2022-08-28T09:11:19.228710Z","iopub.status.idle":"2022-08-28T09:11:19.233973Z","shell.execute_reply.started":"2022-08-28T09:11:19.228662Z","shell.execute_reply":"2022-08-28T09:11:19.233167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLDS = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nHISTORIES = 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-5)\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 = 2,\n        callbacks = [\n            lr_callback,\n        ],\n    )\n    \n    model.save_weights(f'model_{fold}.h5')\n    \n    print('\\n' * 3)\n    break","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:12:05.951941Z","iopub.execute_input":"2022-08-28T09:12:05.952611Z","iopub.status.idle":"2022-08-28T09:37:35.969665Z","shell.execute_reply.started":"2022-08-28T09:12:05.952575Z","shell.execute_reply":"2022-08-28T09:37:35.968868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_results(dataset, nrows, ncols=4):\n    images, labels = next(iter(dataset))\n    \n    # Predict Masks\n    labels_pred = model(images, training=False)\n    \n    fig, axes = plt.subplots(nrows=nrows, ncols=ncols, figsize=(ncols*8, nrows*6))\n    \n    for r, (img, lbl, lbl_pred) in enumerate(zip(images, labels, labels_pred)):\n        if r > nrows - 1: # zero indexed\n            break\n        # Plot Image\n        axes[r, 0].imshow(img)\n        axes[r, 0].set_title('Image', size=18)\n        axes[r, 0].axis(False)\n        \n        # Mask\n        axes[r, 1].imshow(lbl)\n        axes[r, 1].set_title('Mask', size=18)\n        axes[r, 1].axis(False)\n        \n        # Predicted Mask with Threshold\n        axes[r, 2].imshow(lbl_pred)\n        axes[r, 2].set_title('Mask Predicted', size=18)\n        axes[r, 2].axis(False)\n        \n        # Predicted Mask with Threshold\n        lbl_pred_th50 =tf.cast(lbl_pred > 0.50, tf.uint8)\n        axes[r, 3].imshow(lbl_pred_th50)\n        axes[r, 3].set_title('Mask Predicted Threshold 0.50', size=18)\n        axes[r, 3].axis(False)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:37:35.972222Z","iopub.execute_input":"2022-08-28T09:37:35.972473Z","iopub.status.idle":"2022-08-28T09:37:35.984312Z","shell.execute_reply.started":"2022-08-28T09:37:35.972444Z","shell.execute_reply":"2022-08-28T09:37:35.983385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_results(get_dataset(bs=8, idxs=train_idxs), 8)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:37:35.985239Z","iopub.execute_input":"2022-08-28T09:37:35.985463Z","iopub.status.idle":"2022-08-28T09:37:55.890216Z","shell.execute_reply.started":"2022-08-28T09:37:35.985437Z","shell.execute_reply":"2022-08-28T09:37:55.888500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_results(get_dataset(bs=8, idxs=val_idxs, val=True), 8)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:37:55.893597Z","iopub.execute_input":"2022-08-28T09:37:55.894018Z","iopub.status.idle":"2022-08-28T09:38:13.275854Z","shell.execute_reply.started":"2022-08-28T09:37:55.893962Z","shell.execute_reply":"2022-08-28T09:38:13.274254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:13.277725Z","iopub.execute_input":"2022-08-28T09:38:13.277984Z","iopub.status.idle":"2022-08-28T09:38:13.295381Z","shell.execute_reply.started":"2022-08-28T09:38:13.277955Z","shell.execute_reply":"2022-08-28T09:38:13.294185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history_metric('loss', f_best=np.argmin, yscale='log')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:38:13.297011Z","iopub.execute_input":"2022-08-28T09:38:13.297343Z","iopub.status.idle":"2022-08-28T09:38:14.326964Z","shell.execute_reply.started":"2022-08-28T09:38:13.297281Z","shell.execute_reply":"2022-08-28T09:38:14.325941Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:14.328268Z","iopub.execute_input":"2022-08-28T09:38:14.328540Z","iopub.status.idle":"2022-08-28T09:38:14.954764Z","shell.execute_reply.started":"2022-08-28T09:38:14.328505Z","shell.execute_reply":"2022-08-28T09:38:14.954058Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:14.955885Z","iopub.execute_input":"2022-08-28T09:38:14.956243Z","iopub.status.idle":"2022-08-28T09:38:15.599196Z","shell.execute_reply.started":"2022-08-28T09:38:14.956214Z","shell.execute_reply":"2022-08-28T09:38:15.598341Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:15.600683Z","iopub.execute_input":"2022-08-28T09:38:15.601254Z","iopub.status.idle":"2022-08-28T09:38:16.234125Z","shell.execute_reply.started":"2022-08-28T09:38:15.601211Z","shell.execute_reply":"2022-08-28T09:38:16.233184Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:16.237291Z","iopub.execute_input":"2022-08-28T09:38:16.238141Z","iopub.status.idle":"2022-08-28T09:38:16.843362Z","shell.execute_reply.started":"2022-08-28T09:38:16.238091Z","shell.execute_reply":"2022-08-28T09:38:16.842512Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:16.844608Z","iopub.execute_input":"2022-08-28T09:38:16.845255Z","iopub.status.idle":"2022-08-28T09:38:17.464540Z","shell.execute_reply.started":"2022-08-28T09:38:16.845218Z","shell.execute_reply":"2022-08-28T09:38:17.463650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"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\nval_images, val_masks, val_organs = next(iter(dataset))\n# Cast from Tensorflow to Numpy\nval_images = val_images.numpy()\nval_masks = val_masks.numpy()\nval_organs = val_organs.numpy()\nprint(f'val_images shape: {val_images.shape}, val_masks shape: {val_masks.shape}, val_organs shape: {val_organs.shape}')\n\n# Validation Predictions\nVAL_Y_PREDS = model.predict(val_images, verbose=1, batch_size=len(val_idxs))\nprint(f'VAL_Y_PREDS shape: {VAL_Y_PREDS.shape}, VAL_Y_PREDS dtype: {VAL_Y_PREDS.dtype}')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:38:17.465915Z","iopub.execute_input":"2022-08-28T09:38:17.466176Z","iopub.status.idle":"2022-08-28T09:38:48.508556Z","shell.execute_reply.started":"2022-08-28T09:38:17.466146Z","shell.execute_reply":"2022-08-28T09:38:48.507600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-28T09:38:48.509748Z","iopub.execute_input":"2022-08-28T09:38:48.509998Z","iopub.status.idle":"2022-08-28T09:38:48.515337Z","shell.execute_reply.started":"2022-08-28T09:38:48.509968Z","shell.execute_reply":"2022-08-28T09:38:48.514274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_y_true_y_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, y_true, organ, y_pred) in enumerate(zip(tqdm(val_images), val_masks, val_organs, VAL_Y_PREDS)):\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\nIoUs, IoUsOrgans = get_y_true_y_pred()\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:38:48.516339Z","iopub.execute_input":"2022-08-28T09:38:48.517074Z","iopub.status.idle":"2022-08-28T09:39:07.142571Z","shell.execute_reply.started":"2022-08-28T09:38:48.517020Z","shell.execute_reply":"2022-08-28T09:39:07.141638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    return threshold_best","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:07.144024Z","iopub.execute_input":"2022-08-28T09:39:07.144368Z","iopub.status.idle":"2022-08-28T09:39:07.155792Z","shell.execute_reply.started":"2022-08-28T09:39:07.144327Z","shell.execute_reply":"2022-08-28T09:39:07.154871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"threshold_best = plot_iou_by_threshold(IoUs, 'All')","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:07.157291Z","iopub.execute_input":"2022-08-28T09:39:07.157691Z","iopub.status.idle":"2022-08-28T09:39:07.582890Z","shell.execute_reply.started":"2022-08-28T09:39:07.157650Z","shell.execute_reply":"2022-08-28T09:39:07.581968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for organ, ious in IoUsOrgans.items():\n    plot_iou_by_threshold(ious, organ)","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:07.584125Z","iopub.execute_input":"2022-08-28T09:39:07.584375Z","iopub.status.idle":"2022-08-28T09:39:09.626443Z","shell.execute_reply.started":"2022-08-28T09:39:07.584346Z","shell.execute_reply":"2022-08-28T09:39:09.625531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(pd.Series(IoUs[threshold_best]).describe().apply(lambda v: f'{v:.2f}').to_frame(name='Value'))","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:09.627723Z","iopub.execute_input":"2022-08-28T09:39:09.628227Z","iopub.status.idle":"2022-08-28T09:39:09.642340Z","shell.execute_reply.started":"2022-08-28T09:39:09.628188Z","shell.execute_reply":"2022-08-28T09:39:09.641592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(12,8))\npd.Series(IoUs[threshold_best]).plot(kind='hist')\nplt.title('IoU Distribution at Best Threshold', size=24)\nplt.grid()\nplt.xlabel('Threshold', size=16)\nplt.ylabel('Count', size=16)\nplt.xticks(size=12)\nplt.yticks(size=12)\nplt.xlim(0,1)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:09.643546Z","iopub.execute_input":"2022-08-28T09:39:09.643861Z","iopub.status.idle":"2022-08-28T09:39:10.022169Z","shell.execute_reply.started":"2022-08-28T09:39:09.643822Z","shell.execute_reply":"2022-08-28T09:39:10.021401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:10.023171Z","iopub.execute_input":"2022-08-28T09:39:10.023410Z","iopub.status.idle":"2022-08-28T09:39:11.896414Z","shell.execute_reply.started":"2022-08-28T09:39:10.023385Z","shell.execute_reply":"2022-08-28T09:39:11.895290Z"},"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":{"execution":{"iopub.status.busy":"2022-08-28T09:39:11.898352Z","iopub.execute_input":"2022-08-28T09:39:11.898711Z","iopub.status.idle":"2022-08-28T09:39:11.909309Z","shell.execute_reply.started":"2022-08-28T09:39:11.898669Z","shell.execute_reply":"2022-08-28T09:39:11.908201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_validation_predictions():\n    for idx, (image, y_true, organ, y_pred) in enumerate(zip(tqdm(val_images), val_masks, val_organs, VAL_Y_PREDS)):\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()\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:11.910606Z","iopub.execute_input":"2022-08-28T09:39:11.910859Z","iopub.status.idle":"2022-08-28T09:39:11.922981Z","shell.execute_reply.started":"2022-08-28T09:39:11.910829Z","shell.execute_reply":"2022-08-28T09:39:11.922187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()\nplot_validation_predictions()","metadata":{"execution":{"iopub.status.busy":"2022-08-28T09:39:11.925146Z","iopub.execute_input":"2022-08-28T09:39:11.926258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}