{"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":"from tensorflow.keras.layers import Conv2D, BatchNormalization, Activation, MaxPool2D, Conv2DTranspose, Concatenate, Input\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.applications import EfficientNetB0\nimport tensorflow as tf\nfrom tensorflow import keras\nimport keras.backend as K\nfrom keras import layers\nfrom keras.utils import get_file\n\nfrom tensorflow.keras.metrics import Metric\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport PIL.Image as Image\n\nimport random\nimport gc\nimport cv2\nimport time\nfrom tqdm import tqdm\nimport re\nimport math\nfrom collections import namedtuple\nfrom io import StringIO\nimport os\n\ntf.keras.utils.set_random_seed(1234)\n\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\n\nprint(\"TF Version: \", tf.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-26T09:02:01.115064Z","iopub.execute_input":"2023-04-26T09:02:01.115355Z","iopub.status.idle":"2023-04-26T09:02:09.861403Z","shell.execute_reply.started":"2023-04-26T09:02:01.115327Z","shell.execute_reply":"2023-04-26T09:02:09.859769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Keras EfficientNet-based UNet\n\nThis notebook is the training code for [this notebook](https://www.kaggle.com/code/fpeccia/efficientnet-b0-unet-submission). I am providing this code in hope someone is able to find why its overfitting so much.","metadata":{}},{"cell_type":"markdown","source":"## Notebook parameters","metadata":{}},{"cell_type":"code","source":"# Training parameters\nPROD = True # set to False if you have a Weight and Biases api (needs to be added to the notebooks secrets) and want to log your training metrics\nPATCH_SIZE = 128  # e.g. 128x128\nDOWNSAMPLING = 0.5 # Setting this to e.g. 0.5 means images will be loaded as 2x smaller. 1 does nothing.\nZ_DIM = 16   # Number of slices in the z direction. Max value is 65 - Z_START\nZ_START = 24  # Offset of slices in the z direction\nBATCH_SIZE = 8\nTHRESHOLD = 0.5 # Threshold for sigmoid output\nLEARNING_RATE = 0.0001\nWARMUP_EPOCHS = 10 # Learning rate warmup epochs, see function scheduler for implementation\nDECAY_EPOCHS = 75 # Learning rate decay start epochs, see function scheduler for implementation\nEPOCHS = 100\nSTEPS_PER_EPOCH = 100\nOPTIMIZER = 'ADAM' # ADAM or SGD\nLOSS = \"BCE\" # BCE,MODIFIED_DICE,JACCARD,MODIFIED_DICE+BCE\nVAL_FOLD = '2' # 1,2,3,CUSTOM: (all as string). Which volume is used as validation. CUSTOM takes a section of each one as validation.\nMODEL_NAME = 'efficientnetb0-unet' # Can be anything from b0 to b7\nDATA_AUG = True # Enable data augmentation, see function train_augment_fn for implementation\nDIVISOR = 1 # 1,2,4,8: Reduces the amount of filters of the entire UNet by this factor. 1 does not reduce anything.\nPRETRAINED = False # Use pretrained weights for EfficientNet backbone. Can only be used with DIVISOR = 1. First layer is NOT freezed because of the different amount of input channels (pretrained weights expect a 3 channel input).\n\nPATCH_HALFSIZE = PATCH_SIZE // 2","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Config parameters for W&B\nCONFIG = {\n    \"PATCH_SIZE\": PATCH_SIZE,\n    \"PATCH_HALFSIZE\": PATCH_HALFSIZE,\n    \"DOWNSAMPLING\": DOWNSAMPLING,\n    \"Z_DIM\": Z_DIM,\n    \"Z_START\": Z_START,\n    \"BATCH_SIZE\": BATCH_SIZE,\n    \"learning_rate\": LEARNING_RATE,\n    \"epochs\": EPOCHS,\n    \"steps_per_epoch\": STEPS_PER_EPOCH,\n    \"optimizer\": OPTIMIZER,\n    \"THRESHOLD\": THRESHOLD,\n    \"LOSS\": LOSS,\n    \"VAL_FOLD\": VAL_FOLD,\n    \"SIGMOID_OUTPUT\": True,\n    'model_name': MODEL_NAME,\n    'DATA_AUG': DATA_AUG,\n    'WARMUP_EPOCHS': WARMUP_EPOCHS,\n    'DECAY_EPOCHS': DECAY_EPOCHS,\n    'DIVISOR': DIVISOR,\n    'PRETRAINED': PRETRAINED\n}","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG['PRETRAINED']:\n    assert CONFIG['DIVISOR'] == 1","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect TPU, return appropriate distribution strategy\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\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    strategy = tf.distribute.get_strategy() \n\nprint(\"REPLICAS: \", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:09.864461Z","iopub.execute_input":"2023-04-26T09:02:09.865047Z","iopub.status.idle":"2023-04-26T09:02:09.877062Z","shell.execute_reply.started":"2023-04-26T09:02:09.865011Z","shell.execute_reply":"2023-04-26T09:02:09.875966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_kaggle_gpu_enabled():\n    \n    \"\"\"Return whether GPU is enabled in the running Kaggle kernel\"\"\"\n\n    from tensorflow.python.client import device_lib\n\n    # when only CPU is enabled the list shows one CPU entry, otherwise there are more, listing GPU as well\n    return len(device_lib.list_local_devices()) > 1\n\nCONFIG[\"GPU\"] = is_kaggle_gpu_enabled()\nprint(CONFIG[\"GPU\"])","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:09.878849Z","iopub.execute_input":"2023-04-26T09:02:09.879646Z","iopub.status.idle":"2023-04-26T09:02:12.021086Z","shell.execute_reply.started":"2023-04-26T09:02:09.879604Z","shell.execute_reply":"2023-04-26T09:02:12.019919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Losses definitions","metadata":{}},{"cell_type":"code","source":"def jaccard_loss(y_true, y_pred, smooth=100):\n    \"\"\" Calculates mean of Jaccard distance as a loss function \"\"\"\n    y_true_f = K.flatten(y_true)\n    if CONFIG[\"SIGMOID_OUTPUT\"]:\n        y_pred_f = K.flatten(y_pred)\n    else:\n        y_pred_f = tf.keras.activations.sigmoid(K.flatten(y_pred))\n    intersection = tf.reduce_sum(y_true_f * y_pred_f, axis=-1)\n    sum_ = tf.reduce_sum(y_true_f + y_pred_f, axis=-1)\n    jac = (intersection + smooth) / (sum_ - intersection + smooth)\n    jd =  (1 - jac) * smooth\n    return jd\n\ndef bce_loss(y_true, y_pred):\n    BCE = tf.keras.losses.binary_crossentropy(K.flatten(y_true),K.flatten(y_pred),from_logits=False if CONFIG[\"SIGMOID_OUTPUT\"] else True)\n    return BCE\n\ndef binary_fbeta(ytrue , ypred, beta=1, epsilon=1e-7):\n    # epsilon is set so as to avoid division by zero error\n    beta_squared = beta**2 # squaring beta\n\n    # casting ytrue and ypred as float dtype\n    ytrue = tf.cast(ytrue, tf.float32)\n    ypred = tf.cast(ypred, tf.float32)\n\n    tp = tf.reduce_sum(ytrue*ypred) # calculating true positives\n    predicted_positive = tf.reduce_sum(ypred) # calculating predicted positives\n    actual_positive = tf.reduce_sum(ytrue) # calculating actual positives\n    \n    precision = tp/(predicted_positive+epsilon) # calculating precision\n    recall = tp/(actual_positive+epsilon) # calculating recall\n    \n    # calculating fbeta\n    fb = (1+beta_squared)*precision*recall / (beta_squared*precision + recall + epsilon)\n\n    return fb\n\ndef modified_dice_loss(y_true, y_pred, beta=0.5, epsilon=1e-7):\n    \n    beta_squared = beta**2 # squaring beta\n    \n    y_true_f = K.flatten(y_true)\n    if CONFIG[\"SIGMOID_OUTPUT\"]:\n        y_pred_f = K.flatten(y_pred)\n    else:\n        y_pred_f = tf.keras.activations.sigmoid(K.flatten(y_pred))\n    \n    tp = tf.reduce_sum(y_true_f*y_pred_f) # calculating true positives\n    predicted_positive = tf.reduce_sum(y_pred_f) # calculating predicted positives\n    actual_positive = tf.reduce_sum(y_true_f) # calculating actual positives\n    \n    precision = tp/(predicted_positive+epsilon) # calculating precision\n    recall = tp/(actual_positive+epsilon) # calculating recall\n    \n    fb = (1+beta_squared)*precision*recall / (beta_squared*precision + recall + epsilon)\n    \n    return 1-fb\n\ndef combined_loss(y_true, y_pred):\n    BCE = bce_loss(y_true, y_pred)\n    MODIFIED_DICE = modified_dice_loss(y_true, y_pred)\n    return BCE + MODIFIED_DICE","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:12.024406Z","iopub.execute_input":"2023-04-26T09:02:12.025158Z","iopub.status.idle":"2023-04-26T09:02:12.04052Z","shell.execute_reply.started":"2023-04-26T09:02:12.025119Z","shell.execute_reply":"2023-04-26T09:02:12.03968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"losses_map = {\n    \"BCE\": bce_loss,\n    \"MODIFIED_DICE\": modified_dice_loss,\n    \"JACCARD\": jaccard_loss,\n    \"MODIFIED_DICE+BCE\": combined_loss\n}\n\nassert LOSS in losses_map, f\"{LOSS} does not exist in losses_map!\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## FBeta metric for Keras","metadata":{}},{"cell_type":"code","source":"class StatefullBinaryFBeta(Metric):\n    def __init__(self, name='state_full_binary_fbeta', beta=0.5, threshold=THRESHOLD, epsilon=1e-7, **kwargs):\n        # initializing an object of the super class\n        super(StatefullBinaryFBeta, self).__init__(name=name, **kwargs)\n\n        # initializing state variables\n        self.tp = self.add_weight(name='tp', initializer='zeros') # initializing true positives \n        self.actual_positive = self.add_weight(name='fp', initializer='zeros') # initializing actual positives\n        self.predicted_positive = self.add_weight(name='fn', initializer='zeros') # initializing predicted positives\n\n        # initializing other atrributes that wouldn't be changed for every object of this class\n        self.beta_squared = beta**2 \n        self.threshold = threshold\n        self.epsilon = epsilon\n\n    def update_state(self, ytrue, ypred, sample_weight=None):\n        # casting ytrue and ypred as float dtype\n        ytrue = tf.cast(ytrue, tf.float32)\n        if CONFIG[\"SIGMOID_OUTPUT\"]:\n            ypred = tf.cast(ypred, tf.float32)\n        else:\n            ypred = tf.keras.activations.sigmoid(tf.cast(ypred, tf.float32))\n\n        # setting values of ypred greater than the set threshold to 1 while those lesser to 0\n        ypred = tf.cast(tf.greater_equal(ypred, tf.constant(self.threshold)), tf.float32)\n\n        self.tp.assign_add(tf.reduce_sum(ytrue*ypred)) # updating true positives atrribute\n        self.predicted_positive.assign_add(tf.reduce_sum(ypred)) # updating predicted positive atrribute\n        self.actual_positive.assign_add(tf.reduce_sum(ytrue)) # updating actual positive atrribute\n\n    def result(self):\n        self.precision = self.tp/(self.predicted_positive+self.epsilon) # calculates precision\n        self.recall = self.tp/(self.actual_positive+self.epsilon) # calculates recall\n\n        # calculating fbeta\n        self.fb = (1+self.beta_squared)*self.precision*self.recall / (self.beta_squared*self.precision + self.recall + self.epsilon)\n\n        return self.fb\n\n    def reset_state(self):\n        self.tp.assign(0) # resets true positives to zero\n        self.predicted_positive.assign(0) # resets predicted positives to zero\n        self.actual_positive.assign(0) # resets actual positives to zero","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:12.042088Z","iopub.execute_input":"2023-04-26T09:02:12.04286Z","iopub.status.idle":"2023-04-26T09:02:12.056397Z","shell.execute_reply.started":"2023-04-26T09:02:12.042822Z","shell.execute_reply":"2023-04-26T09:02:12.055611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Installation and configuration for W&B","metadata":{}},{"cell_type":"code","source":"if not PROD:\n    !pip install --upgrade -q wandb\n    \n    from kaggle_secrets import UserSecretsClient\n    import wandb\n\n    user_secrets = UserSecretsClient()\n\n    # I have saved my API token with \"wandb_api\" as Label. \n    # If you use some other Label make sure to change the same below. \n    wandb_api = user_secrets.get_secret(\"wandb_api\") \n\n    wandb.login(key=wandb_api)\n    \n    from wandb.keras import WandbCallback","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:12.071274Z","iopub.execute_input":"2023-04-26T09:02:12.071937Z","iopub.status.idle":"2023-04-26T09:02:26.289635Z","shell.execute_reply.started":"2023-04-26T09:02:12.071901Z","shell.execute_reply":"2023-04-26T09:02:26.28842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions definitions","metadata":{}},{"cell_type":"markdown","source":"## Data import","metadata":{}},{"cell_type":"code","source":"def resize(img,downsampling=DOWNSAMPLING):\n    if downsampling != 1.:\n        size = int(img.shape[1] * downsampling), int(img.shape[0] * downsampling)\n        img = cv2.resize(img, size)\n    return img\n\ndef resize_to_original(img,original_mask):\n    size = original_mask.shape[1], original_mask.shape[0]\n    img = cv2.resize(img, size)\n    return img\n\ndef load_mask(split, index):\n    img = cv2.imread(f\"{DATA_DIR}/{split}/{index}/mask.png\", 0)\n    img = resize(img)\n    return img.astype(\"bool\")\n\n\ndef load_labels(split, index):\n    img = cv2.imread(f\"{DATA_DIR}/{split}/{index}/inklabels.png\", 0)\n    img = resize(img)\n    return np.expand_dims(img, axis=-1)\n\n\ndef load_volume(split, index):\n    # A more memory-efficient volune loader\n    fnames = [f\"{DATA_DIR}/{split}/{index}/surface_volume/{i:02}.tif\"\n             for i in range(Z_START, Z_START + Z_DIM)]\n\n    batch_size = 8\n    fname_batches = [fnames[i :i + batch_size] for i in range(0, len(fnames), batch_size)]\n    volumes = []\n    for fname_batch in fname_batches:\n        z_slices = []\n        for fname in tqdm(fname_batch):\n            img = cv2.imread(fname, 0)\n            img = resize(img)\n            z_slices.append(img)\n        volumes.append(np.stack(z_slices, axis=-1))\n        del z_slices\n    return np.concatenate(volumes, axis=-1)\n\n\ndef load_sample(split, index):\n    print(f\"Loading '{split}/{index}'...\")\n    gc.collect()\n    if split == \"train\":\n        return load_volume(split, index), load_mask(split, index), load_labels(split, index)\n    return load_volume(split, index), load_mask(split, index), None","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:29.266512Z","iopub.execute_input":"2023-04-26T09:02:29.266889Z","iopub.status.idle":"2023-04-26T09:02:29.280991Z","shell.execute_reply.started":"2023-04-26T09:02:29.266861Z","shell.execute_reply":"2023-04-26T09:02:29.279966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"volume_1, mask_1, labels_1 = load_sample(split=\"train\", index=1)\nvolume_2, mask_2, labels_2 = load_sample(split=\"train\", index=2)\nvolume_3, mask_3, labels_3 = load_sample(split=\"train\", index=3)\ngc.collect()\nprint(\"Loading complete.\")","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:02:29.283898Z","iopub.execute_input":"2023-04-26T09:02:29.284314Z","iopub.status.idle":"2023-04-26T09:03:28.063465Z","shell.execute_reply.started":"2023-04-26T09:02:29.284278Z","shell.execute_reply":"2023-04-26T09:03:28.062373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Plot the imported data for sanity check","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(2, 3)\nax[0,0].imshow(labels_1, cmap='gray')\nax[0,1].imshow(labels_2, cmap='gray')\nax[0,2].imshow(labels_3, cmap='gray')\nax[1,0].imshow(mask_1, cmap='gray')\nax[1,1].imshow(mask_2, cmap='gray')\nax[1,2].imshow(mask_3, cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:28.065193Z","iopub.execute_input":"2023-04-26T09:03:28.065914Z","iopub.status.idle":"2023-04-26T09:03:31.027069Z","shell.execute_reply.started":"2023-04-26T09:03:28.065874Z","shell.execute_reply":"2023-04-26T09:03:31.026066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Obtain train and validation sets\n\nThis is the CUSTOM validation process, we put a little of each train dataset in the validation set","metadata":{}},{"cell_type":"code","source":"raw_limits = [\n    {\n        \"y_min\": int(mask_1.shape[1]*0.3) + PATCH_HALFSIZE,\n        \"y_max\": int(mask_1.shape[1]*0.7) - PATCH_HALFSIZE,\n        \"x_min\": int(mask_1.shape[0]*0.3) + PATCH_HALFSIZE,\n        \"x_max\": int(mask_1.shape[0]*0.7) - PATCH_HALFSIZE\n    },\n    {\n        \"y_min\": int(mask_2.shape[1]*0.2) + PATCH_HALFSIZE,\n        \"y_max\": int(mask_2.shape[1]*0.5) - PATCH_HALFSIZE,\n        \"x_min\": int(mask_2.shape[0]*0.4) + PATCH_HALFSIZE,\n        \"x_max\": int(mask_2.shape[0]*0.6) - PATCH_HALFSIZE\n    },\n    {\n        \"y_min\": int(mask_3.shape[1]*0.2) + PATCH_HALFSIZE,\n        \"y_max\": int(mask_3.shape[1]*0.5) - PATCH_HALFSIZE,\n        \"x_min\": int(mask_3.shape[0]*0.3) + PATCH_HALFSIZE,\n        \"x_max\": int(mask_3.shape[0]*0.7) - PATCH_HALFSIZE\n    }\n]\n\nraw_volumes = [volume_1, volume_2, volume_3]\nraw_labels = [labels_1, labels_2, labels_3]\nraw_masks = [mask_1, mask_2, mask_3]\n\n#raw_volumes = [volume_2]\n#raw_labels = [labels_2]\n#raw_masks = [mask_2]\n\ntrain_volumes = []\ntrain_labels = []\ntrain_masks = []\nvalidation_volumes = []\nvalidation_labels = []\nvalidation_mask = []\n\nvalidation_coords = []\n\nfor idx, raw_label in enumerate(raw_labels):\n    orig_width = raw_label.shape[0]\n    orig_height = raw_label.shape[1]\n    \n    val_patch_width = int(orig_width*0.2) \n    val_patch_height = int(orig_height*0.2) \n    val_patch_half_width = val_patch_width // 2\n    val_patch_half_height = val_patch_height // 2\n    \n    label_ink_perc = 0\n    while label_ink_perc < 0.2:\n        x = np.random.randint(raw_limits[idx][\"x_min\"], raw_limits[idx][\"x_max\"])\n        y = np.random.randint(raw_limits[idx][\"y_min\"], raw_limits[idx][\"y_max\"])\n        \n        if np.all(raw_masks[idx][x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height]):\n            if np.sum(raw_masks[idx][x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height])/(val_patch_width*val_patch_height) > 0.4:\n                label_patch = raw_label[x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height,:]\n                label_ink_perc = np.sum(label_patch)/np.size(label_patch)\n    \n    # Extract validation volume\n    validation_volumes.append(raw_volumes[idx][x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height,:])\n    validation_labels.append(raw_labels[idx][x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height,:])\n    validation_mask.append(raw_masks[idx][x-val_patch_half_width:x+val_patch_half_width, y-val_patch_half_height:y+val_patch_half_height])\n    \n    validation_coords.append(\n        (y-val_patch_half_height,x-val_patch_half_width,val_patch_height,val_patch_width)\n    )\n    \n    # Extract train volumes\n    x_max = x-val_patch_half_width\n    train_volumes.append(raw_volumes[idx][:x_max,:,:])\n    train_labels.append(raw_labels[idx][:x_max,:,:])\n    train_masks.append(raw_masks[idx][:x_max,:])\n    \n    x_min = x+val_patch_half_width\n    train_volumes.append(raw_volumes[idx][x_min:,:,:])\n    train_labels.append(raw_labels[idx][x_min:,:,:])\n    train_masks.append(raw_masks[idx][x_min:,:])\n    \n    y_max = y - val_patch_half_height\n    x_min = x-val_patch_half_width\n    x_max = x+val_patch_half_width\n    train_volumes.append(raw_volumes[idx][x_min:x_max,:y_max,:])\n    train_labels.append(raw_labels[idx][x_min:x_max,:y_max,:])\n    train_masks.append(raw_masks[idx][x_min:x_max,:y_max])\n    \n    y_min = y + val_patch_half_height\n    x_min = x-val_patch_half_width\n    x_max = x+val_patch_half_width\n    train_volumes.append(raw_volumes[idx][x_min:x_max,y_min:,:])\n    train_labels.append(raw_labels[idx][x_min:x_max,y_min:,:])\n    train_masks.append(raw_masks[idx][x_min:x_max,y_min:])\n    ","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:36.459418Z","iopub.execute_input":"2023-04-26T09:03:36.459735Z","iopub.status.idle":"2023-04-26T09:03:36.488374Z","shell.execute_reply.started":"2023-04-26T09:03:36.459706Z","shell.execute_reply":"2023-04-26T09:03:36.487202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Print training labels, just for a quick sanity check","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(2, len(train_volumes))\nfor idx in range(len(train_volumes)):\n    ax[0,idx].imshow(train_labels[idx], cmap='gray')\n    ax[1,idx].imshow(train_masks[idx], cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:36.490215Z","iopub.execute_input":"2023-04-26T09:03:36.490643Z","iopub.status.idle":"2023-04-26T09:03:38.98117Z","shell.execute_reply.started":"2023-04-26T09:03:36.490604Z","shell.execute_reply":"2023-04-26T09:03:38.980012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Print validation labels, just for a quick sanity check","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(1, len(raw_labels))\nfor i in range(len(raw_labels)):\n    if len(raw_labels) > 1:\n        ax[i].imshow(raw_labels[i], cmap='gray')\n        rect = validation_coords[i]\n        patch = patches.Rectangle((rect[0], rect[1]), rect[2], rect[3], linewidth=2, edgecolor='r', facecolor='none')\n        ax[i].add_patch(patch)\n    else:\n        ax.imshow(raw_labels[i], cmap='gray')\n        rect = validation_coords[i]\n        patch = patches.Rectangle((rect[0], rect[1]), rect[2], rect[3], linewidth=2, edgecolor='r', facecolor='none')\n        ax.add_patch(patch)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:38.982679Z","iopub.execute_input":"2023-04-26T09:03:38.983242Z","iopub.status.idle":"2023-04-26T09:03:41.189916Z","shell.execute_reply.started":"2023-04-26T09:03:38.983206Z","shell.execute_reply":"2023-04-26T09:03:41.188975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(2, len(validation_volumes))\nfor idx in range(len(validation_volumes)):\n    if len(raw_labels) > 1:\n        ax[0,idx].imshow(validation_labels[idx], cmap='gray')\n        ax[1,idx].imshow(validation_mask[idx], cmap='gray')\n    else:\n        ax[0].imshow(validation_labels[idx], cmap='gray')\n        ax[1].imshow(validation_mask[idx], cmap='gray')\n    assert np.any(validation_mask[idx])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:41.193971Z","iopub.execute_input":"2023-04-26T09:03:41.194832Z","iopub.status.idle":"2023-04-26T09:03:42.204293Z","shell.execute_reply.started":"2023-04-26T09:03:41.194791Z","shell.execute_reply":"2023-04-26T09:03:42.203137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dev_folds = {\n    \"1\": {\n        \"train_volumes\": [volume_2, volume_3],\n        \"train_labels\": [labels_2, labels_3],\n        \"train_masks\": [mask_2, mask_3],\n        \"validation_volume\": [volume_1],\n        \"validation_labels\": [labels_1],\n        \"validation_mask\": [mask_1],\n    },\n    \"2\": {\n        \"train_volumes\": [volume_1, volume_3],\n        \"train_labels\": [labels_1, labels_3],\n        \"train_masks\": [mask_1, mask_3],\n        \"validation_volume\": [volume_2],\n        \"validation_labels\": [labels_2],\n        \"validation_mask\": [mask_2],\n    },\n    \"3\": {\n        \"train_volumes\": [volume_1, volume_2],\n        \"train_labels\": [labels_1, labels_2],\n        \"train_masks\": [mask_1, mask_2],\n        \"validation_volume\": [volume_3],\n        \"validation_labels\": [labels_3],\n        \"validation_mask\": [mask_3],\n    },\n    \"CUSTOM\": {\n        \"train_volumes\": train_volumes,\n        \"train_labels\": train_labels,\n        \"train_masks\": train_masks,\n        \"validation_volume\": validation_volumes,\n        \"validation_labels\": validation_labels,\n        \"validation_mask\": validation_mask,\n    }\n}\n\nprod_data  = {\n    \"train_volumes\": [volume_1, volume_2, volume_3],\n    \"train_labels\": [labels_1, labels_2, labels_3],\n    \"train_masks\": [mask_1, mask_2, mask_3],\n}","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:42.205967Z","iopub.execute_input":"2023-04-26T09:03:42.206332Z","iopub.status.idle":"2023-04-26T09:03:42.215375Z","shell.execute_reply.started":"2023-04-26T09:03:42.206296Z","shell.execute_reply":"2023-04-26T09:03:42.214266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data augmentation pipeline","metadata":{}},{"cell_type":"code","source":"@tf.function\ndef train_augment_fn(patch, labels):\n    \n    seed = (tf.random.uniform(shape=(), minval=0, maxval=255, dtype=tf.int32),tf.random.uniform(shape=(), minval=0, maxval=255, dtype=tf.int32))\n\n    # Random flip left/right\n    patch = tf.image.stateless_random_flip_left_right(patch, seed=seed)\n    labels = tf.image.stateless_random_flip_left_right(labels, seed=seed)\n    \n    # Random flip up/down\n    patch = tf.image.stateless_random_flip_up_down(patch, seed=seed)\n    labels = tf.image.stateless_random_flip_up_down(labels, seed=seed)\n    \n    # Random rotation -270/270 with a 90 degree step\n    third_seed = tf.random.uniform(shape=(), minval=-3, maxval=4, dtype=tf.int32)\n    patch = tf.image.rot90(patch,k=third_seed)\n    labels = tf.image.rot90(labels,k=third_seed)\n    \n    # Random contrast\n    # patch = tf.image.stateless_random_contrast(patch, 0.0, 0.1, seed)\n    \n    # Random brightness\n    # patch = tf.image.stateless_random_brightness(patch,0.1,seed)\n    \n    return patch, labels","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:42.216853Z","iopub.execute_input":"2023-04-26T09:03:42.217467Z","iopub.status.idle":"2023-04-26T09:03:42.229905Z","shell.execute_reply.started":"2023-04-26T09:03:42.21743Z","shell.execute_reply":"2023-04-26T09:03:42.228962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utilities to create dataset","metadata":{}},{"cell_type":"code","source":"def sample_random_location(shape):\n    x = random.randint(PATCH_HALFSIZE, shape[0] - PATCH_HALFSIZE - 1)\n    y = random.randint(PATCH_HALFSIZE, shape[1] - PATCH_HALFSIZE - 1)\n    return (x, y)\n\n\ndef list_all_locations(mask, stride=PATCH_SIZE):\n    locations = []\n    for x in range(PATCH_HALFSIZE, mask.shape[0] - PATCH_HALFSIZE, stride):\n        for y in range(PATCH_HALFSIZE, mask.shape[1] - PATCH_HALFSIZE, stride):\n            if mask[x, y]:  \n                locations.append((x, y))\n    return locations\n\n\ndef extract_patch(location, volume):\n    x = location[0]\n    y = location[1]\n    patch = volume[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                   y - PATCH_HALFSIZE :y + PATCH_HALFSIZE, :]\n    return patch.astype(\"float32\") / 255.\n\n\ndef extract_labels(location, labels):\n    x = location[0]\n    y = location[1]\n    \n    label = labels[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                   y - PATCH_HALFSIZE :y + PATCH_HALFSIZE, :]\n    return label.astype(\"float32\") / 255.\n\n\ndef make_random_data_generator(volume, mask, labels):\n    def data_generator():\n        while True:\n            loc = sample_random_location(mask.shape)\n            x = loc[0]\n            y = loc[1]\n            if mask[x, y]:    \n                patch = extract_patch(loc, volume)\n                label = extract_labels(loc, labels)\n                yield patch, label\n    return data_generator\n\n\ndef make_iterated_data_generator(volume, mask, labels=None, return_locations=False):\n    locations = list_all_locations(mask)\n    def data_generator():\n        for loc in locations:\n            patch = extract_patch(loc, volume)\n            if labels is None:\n                if return_locations:\n                    yield patch, loc\n                else:\n                    yield patch\n            else:\n                label = extract_labels(loc, labels)\n                if return_locations:\n                    yield patch, label, loc\n                else:\n                    yield patch, label\n    return data_generator\n\ndef make_random_data_generator_from_list(volume_list, mask_list, labels_list):\n    def data_generator():\n        while True:\n            dataset_idx = random.randint(0,len(volume_list)-1)\n            #print(f\"Returning locations for dataset {dataset_idx}...\")\n            loc = sample_random_location(mask_list[dataset_idx].shape)\n            x = loc[0]\n            y = loc[1]\n            if mask_list[dataset_idx][x, y]:    \n                patch = extract_patch(loc, volume_list[dataset_idx])\n                label = extract_labels(loc, labels_list[dataset_idx])\n                yield patch, label\n    return data_generator\n    \ndef make_iterated_data_generator_from_list(volume_list, mask_list, labels_list=None, return_locations=False):\n    locations = []\n    list_length = len(mask_list)\n    for i in range(list_length):\n        locations.append(list_all_locations(mask_list[i]))\n    def data_generator():\n        for i in range(list_length):\n            for loc in locations[i]:\n                patch = extract_patch(loc, volume_list[i])\n                if labels_list is None:\n                    if return_locations:\n                        yield patch, loc\n                    else:\n                        yield patch\n                else:\n                    label = extract_labels(loc, labels_list[i])\n                    if return_locations:\n                        yield patch, label, loc\n                    else:\n                        yield patch, label\n    return data_generator\n\ndef make_tf_dataset(gen_fn, labeled=True, return_locations=False):\n    if labeled:\n        if return_locations:\n            output_signature = (\n                tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32),\n                tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, 1), dtype=tf.float32),\n                tf.TensorSpec(shape=(2,), dtype=tf.float32),\n            )\n        else:\n            output_signature = (\n                tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32),\n                tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, 1), dtype=tf.float32),\n            )\n    else:\n        if return_locations:\n            output_signature = (\n                tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32),\n                tf.TensorSpec(shape=(2,), dtype=tf.float32),\n            )\n        else:\n            output_signature = tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32)\n    ds = tf.data.Dataset.from_generator(\n        gen_fn,\n        output_signature=output_signature,\n    )\n    return ds.prefetch(tf.data.AUTOTUNE).batch(BATCH_SIZE)\n\ndef make_datasets_for_fold(fold, train_augment_fn=None, return_locations=False):\n    train_volumes = fold[\"train_volumes\"]\n    train_masks = fold[\"train_masks\"]\n    train_labels = fold[\"train_labels\"]\n    \n    include_validation = \"validation_volume\" in fold\n    if include_validation:\n        validation_volume = fold[\"validation_volume\"]\n        validation_mask = fold[\"validation_mask\"]\n        validation_labels = fold[\"validation_labels\"]\n\n    train_ds = make_tf_dataset(\n        make_random_data_generator_from_list(train_volumes, train_masks, train_labels),\n        labeled=True\n    )\n    \n    if train_augment_fn:\n        train_ds = train_ds.map(train_augment_fn, num_parallel_calls=tf.data.AUTOTUNE)\n    train_ds = train_ds.prefetch(tf.data.AUTOTUNE)\n\n    if not include_validation:\n        return train_ds\n\n    val_ds = make_tf_dataset(\n            make_iterated_data_generator_from_list(validation_volume, validation_mask, validation_labels, return_locations=return_locations),\n            labeled=True,\n            return_locations=return_locations\n        )\n    \n    return train_ds, val_ds","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:42.231444Z","iopub.execute_input":"2023-04-26T09:03:42.231822Z","iopub.status.idle":"2023-04-26T09:03:42.258645Z","shell.execute_reply.started":"2023-04-26T09:03:42.231786Z","shell.execute_reply":"2023-04-26T09:03:42.25774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create actual train and validation datasets","metadata":{}},{"cell_type":"code","source":"train_ds, val_ds = make_datasets_for_fold(dev_folds[CONFIG[\"VAL_FOLD\"]], train_augment_fn=train_augment_fn if CONFIG['DATA_AUG'] else None, return_locations=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:42.260177Z","iopub.execute_input":"2023-04-26T09:03:42.260711Z","iopub.status.idle":"2023-04-26T09:03:42.64319Z","shell.execute_reply.started":"2023-04-26T09:03:42.260662Z","shell.execute_reply":"2023-04-26T09:03:42.642181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sanity check\n\nHow big is each metric when we predict a complete zero (no ink) mask?","metadata":{}},{"cell_type":"code","source":"accuracy_metric = tf.keras.metrics.BinaryAccuracy(name=\"binary_accuracy\",threshold=THRESHOLD)\niou_metric = tf.keras.metrics.BinaryIoU(name=\"binary_iou\",threshold=THRESHOLD)\nfbeta_metric = StatefullBinaryFBeta(name=\"fbeta_score\",beta=0.5,threshold=THRESHOLD)\n\nfor (patch_batch, y_true) in val_ds:\n    y_pred_zero = np.zeros(patch_batch.shape[:3] + (1,))\n    accuracy_metric.update_state(y_true, y_pred_zero)\n    iou_metric.update_state(y_true, y_pred_zero)\n    fbeta_metric.update_state(y_true, y_pred_zero)\n\nprint(f\"If prediction is all zero:\")\nprint(f\"\\tBinary accuracy: {accuracy_metric.result()}\")\nprint(f\"\\tBinary IoU: {iou_metric.result()}\")\nprint(f\"\\tFBeta score: {fbeta_metric.result()}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:42.644512Z","iopub.execute_input":"2023-04-26T09:03:42.644882Z","iopub.status.idle":"2023-04-26T09:03:44.111849Z","shell.execute_reply.started":"2023-04-26T09:03:42.644845Z","shell.execute_reply":"2023-04-26T09:03:44.109954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"How big is each metric when we predict a complete ones (all ink) mask?","metadata":{}},{"cell_type":"code","source":"accuracy_metric = tf.keras.metrics.BinaryAccuracy(name=\"binary_accuracy\",threshold=THRESHOLD)\niou_metric = tf.keras.metrics.BinaryIoU(name=\"binary_iou\",threshold=THRESHOLD)\nfbeta_metric = StatefullBinaryFBeta(name=\"fbeta_score\",beta=0.5,threshold=THRESHOLD)\n\nfor (patch_batch, y_true) in val_ds:\n    y_pred_zero = np.ones(patch_batch.shape[:3] + (1,))\n    accuracy_metric.update_state(y_true, y_pred_zero)\n    iou_metric.update_state(y_true, y_pred_zero)\n    fbeta_metric.update_state(y_true, y_pred_zero)\n\nprint(f\"If prediction is all ones:\")\nprint(f\"\\tBinary accuracy: {accuracy_metric.result()}\")\nprint(f\"\\tBinary IoU: {iou_metric.result()}\")\nprint(f\"\\tFBeta score: {fbeta_metric.result()}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:44.113511Z","iopub.execute_input":"2023-04-26T09:03:44.113865Z","iopub.status.idle":"2023-04-26T09:03:45.290187Z","shell.execute_reply.started":"2023-04-26T09:03:44.113829Z","shell.execute_reply":"2023-04-26T09:03:45.289173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Get the total validation samples, we will use it in the .fit() method.","metadata":{}},{"cell_type":"code","source":"total_validation_samples = 0\nfor patch_batch, y_true_batch in tqdm(val_ds):\n    total_validation_samples += patch_batch.shape[0]\nprint(f\"Total validation samples = {total_validation_samples}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:47.089271Z","iopub.execute_input":"2023-04-26T09:03:47.09107Z","iopub.status.idle":"2023-04-26T09:03:47.618131Z","shell.execute_reply.started":"2023-04-26T09:03:47.09103Z","shell.execute_reply":"2023-04-26T09:03:47.617075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model definition\n\nModel is a UNet with and EfficientNet backbone. \n\nNote: I commented out the ReLu activation in function DoubleConv, because I was having problems with gradients dissapearing during training. Just keep that in mind, you can add them again and check if you are able to get better results than me.","metadata":{}},{"cell_type":"code","source":"GlobalParams = namedtuple('GlobalParams', ['batch_norm_momentum', 'batch_norm_epsilon', 'dropout_rate', 'num_classes',\n                                           'width_coefficient', 'depth_coefficient', 'depth_divisor', 'min_depth',\n                                           'drop_connect_rate'])\nGlobalParams.__new__.__defaults__ = (None,) * len(GlobalParams._fields)\n\nBlockArgs = namedtuple('BlockArgs', ['kernel_size', 'num_repeat', 'input_filters', 'output_filters', 'expand_ratio',\n                                     'id_skip', 'strides', 'se_ratio'])\nBlockArgs.__new__.__defaults__ = (None,) * len(BlockArgs._fields)\n\nIMAGENET_WEIGHTS = {\n\n    'efficientnet-b0': {\n        'name': 'efficientnet-b0_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b0_imagenet_1000.h5',\n        'md5': 'bca04d16b1b8a7c607b1152fe9261af7',\n    },\n\n    'efficientnet-b1': {\n        'name': 'efficientnet-b1_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b1_imagenet_1000.h5',\n        'md5': 'bd4a2b82f6f6bada74fc754553c464fc',\n    },\n\n    'efficientnet-b2': {\n        'name': 'efficientnet-b2_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b2_imagenet_1000.h5',\n        'md5': '45b28b26f15958bac270ab527a376999',\n    },\n\n    'efficientnet-b3': {\n        'name': 'efficientnet-b3_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b3_imagenet_1000.h5',\n        'md5': 'decd2c8a23971734f9d3f6b4053bf424',\n    },\n\n    'efficientnet-b4': {\n        'name': 'efficientnet-b4_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b4_imagenet_1000.h5',\n        'md5': '01df77157a86609530aeb4f1f9527949',\n    },\n\n    'efficientnet-b5': {\n        'name': 'efficientnet-b5_imagenet_1000.h5',\n        'url': 'https://github.com/qubvel/efficientnet/releases/download/v0.0.1/efficientnet-b5_imagenet_1000.h5',\n        'md5': 'c31311a1a38b5111e14457145fccdf32',\n    }\n\n}\n\n\ndef round_filters(filters, global_params):\n    \"\"\"Round number of filters.\"\"\"\n    multiplier = global_params.width_coefficient\n    divisor = global_params.depth_divisor\n    min_depth = global_params.min_depth\n    if not multiplier:\n        return filters\n\n    filters *= multiplier\n    min_depth = min_depth or divisor\n    new_filters = max(min_depth, int(filters + divisor / 2) // divisor * divisor)\n    # Make sure that round down does not go down by more than 10%.\n    if new_filters < 0.9 * filters:\n        new_filters += divisor\n    return int(new_filters)\n\n\ndef round_repeats(repeats, global_params):\n    \"\"\"Round number of repeats.\"\"\"\n    multiplier = global_params.depth_coefficient\n    if not multiplier:\n        return repeats\n    return int(math.ceil(multiplier * repeats))\n\n\ndef get_efficientnet_params(model_name, override_params=None):\n    \"\"\"Get efficientnet params based on model name.\"\"\"\n    params_dict = {\n        # (width_coefficient, depth_coefficient, resolution, dropout_rate)\n        # Note: the resolution here is just for reference, its values won't be used.\n        'efficientnet-b0': (1.0, 1.0, 224, 0.2),\n        'efficientnet-b1': (1.0, 1.1, 240, 0.2),\n        'efficientnet-b2': (1.1, 1.2, 260, 0.3),\n        'efficientnet-b3': (1.2, 1.4, 300, 0.3),\n        'efficientnet-b4': (1.4, 1.8, 380, 0.4),\n        'efficientnet-b5': (1.6, 2.2, 456, 0.4),\n        'efficientnet-b6': (1.8, 2.6, 528, 0.5),\n        'efficientnet-b7': (2.0, 3.1, 600, 0.5),\n    }\n    if model_name not in params_dict.keys():\n        raise KeyError('There is no model named {}.'.format(model_name))\n\n    width_coefficient, depth_coefficient, _, dropout_rate = params_dict[model_name]\n\n    if CONFIG['DIVISOR'] == 1:\n        blocks_args = [\n            'r1_k3_s11_e1_i32_o16_se0.25', 'r2_k3_s22_e6_i16_o24_se0.25',\n            'r2_k5_s22_e6_i24_o40_se0.25', 'r3_k3_s22_e6_i40_o80_se0.25',\n            'r3_k5_s11_e6_i80_o112_se0.25', 'r4_k5_s22_e6_i112_o192_se0.25',\n            'r1_k3_s11_e6_i192_o320_se0.25',\n        ]\n    elif CONFIG['DIVISOR'] == 2:\n        blocks_args = [\n            'r1_k3_s11_e1_i16_o8_se0.25', 'r2_k3_s22_e6_i8_o12_se0.25',\n            'r2_k5_s22_e6_i12_o20_se0.25', 'r3_k3_s22_e6_i20_o40_se0.25',\n            'r3_k5_s11_e6_i40_o56_se0.25', 'r4_k5_s22_e6_i56_o96_se0.25',\n            'r1_k3_s11_e6_i96_o160_se0.25',\n        ]\n    elif CONFIG['DIVISOR'] == 4:\n        blocks_args = [\n            'r1_k3_s11_e1_i8_o4_se0.25', 'r2_k3_s22_e6_i4_o6_se0.25',\n            'r2_k5_s22_e6_i6_o10_se0.25', 'r3_k3_s22_e6_i10_o20_se0.25',\n            'r3_k5_s11_e6_i20_o28_se0.25', 'r4_k5_s22_e6_i28_o48_se0.25',\n            'r1_k3_s11_e6_i48_o80_se0.25',\n        ]\n    elif CONFIG['DIVISOR'] == 8:\n        blocks_args = [\n            'r1_k3_s11_e1_i4_o2_se0.25', 'r2_k3_s22_e6_i2_o3_se0.25',\n            'r2_k5_s22_e6_i3_o5_se0.25', 'r3_k3_s22_e6_i5_o10_se0.25',\n            'r3_k5_s11_e6_i10_o14_se0.25', 'r4_k5_s22_e6_i14_o24_se0.25',\n            'r1_k3_s11_e6_i24_o40_se0.25',\n        ]\n    global_params = GlobalParams(\n        batch_norm_momentum=0.99,\n        batch_norm_epsilon=1e-3,\n        dropout_rate=dropout_rate,\n        drop_connect_rate=0.2,\n        num_classes=1000,\n        width_coefficient=width_coefficient,\n        depth_coefficient=depth_coefficient,\n        depth_divisor=8,\n        min_depth=None)\n\n    if override_params:\n        global_params = global_params._replace(**override_params)\n\n    decoder = BlockDecoder()\n    return decoder.decode(blocks_args), global_params\n\n\nclass BlockDecoder(object):\n    \"\"\"Block Decoder for readability.\"\"\"\n\n    @staticmethod\n    def _decode_block_string(block_string):\n        \"\"\"Gets a block through a string notation of arguments.\"\"\"\n        assert isinstance(block_string, str)\n        ops = block_string.split('_')\n        options = {}\n        for op in ops:\n            splits = re.split(r'(\\d.*)', op)\n            if len(splits) >= 2:\n                key, value = splits[:2]\n                options[key] = value\n\n        if 's' not in options or len(options['s']) != 2:\n            raise ValueError('Strides options should be a pair of integers.')\n\n        return BlockArgs(\n            kernel_size=int(options['k']),\n            num_repeat=int(options['r']),\n            input_filters=int(options['i']),\n            output_filters=int(options['o']),\n            expand_ratio=int(options['e']),\n            id_skip=('noskip' not in block_string),\n            se_ratio=float(options['se']) if 'se' in options else None,\n            strides=[int(options['s'][0]), int(options['s'][1])]\n        )\n\n    @staticmethod\n    def _encode_block_string(block):\n        \"\"\"Encodes a block to a string.\"\"\"\n        args = [\n            'r%d' % block.num_repeat,\n            'k%d' % block.kernel_size,\n            's%d%d' % (block.strides[0], block.strides[1]),\n            'e%s' % block.expand_ratio,\n            'i%d' % block.input_filters,\n            'o%d' % block.output_filters\n        ]\n        if 0 < block.se_ratio <= 1:\n            args.append('se%s' % block.se_ratio)\n        if block.id_skip is False:\n            args.append('noskip')\n        return '_'.join(args)\n\n    def decode(self, string_list):\n        \"\"\"Decodes a list of string notations to specify blocks inside the network.\n        Args:\n          string_list: a list of strings, each string is a notation of block.\n        Returns:\n          A list of namedtuples to represent blocks arguments.\n        \"\"\"\n        assert isinstance(string_list, list)\n        blocks_args = []\n        for block_string in string_list:\n            blocks_args.append(self._decode_block_string(block_string))\n        return blocks_args\n\n    def encode(self, blocks_args):\n        \"\"\"Encodes a list of Blocks to a list of strings.\n        Args:\n          blocks_args: A list of namedtuples to represent blocks arguments.\n        Returns:\n          a list of strings, each string is a notation of block.\n        \"\"\"\n        block_strings = []\n        for block in blocks_args:\n            block_strings.append(self._encode_block_string(block))\n        return block_strings\n\n\nclass Swish(layers.Layer):\n    def __init__(self, name=None, **kwargs):\n        super().__init__(name=name, **kwargs)\n\n    def call(self, inputs, **kwargs):\n        return tf.nn.swish(inputs)\n\n    def get_config(self):\n        config = super().get_config()\n        config['name'] = self.name\n        return config\n\n\ndef SEBlock(block_args, **kwargs):\n    num_reduced_filters = max(\n        1, int(block_args.input_filters * block_args.se_ratio))\n    filters = block_args.input_filters * block_args.expand_ratio\n\n    spatial_dims = [1, 2]\n\n    try:\n        block_name = kwargs['block_name']\n    except KeyError:\n        block_name = ''\n\n    def block(inputs):\n        x = inputs\n        x = layers.Lambda(lambda a: K.mean(a, axis=spatial_dims, keepdims=True))(x)\n        x = layers.Conv2D(\n            num_reduced_filters,\n            kernel_size=[1, 1],\n            strides=[1, 1],\n            kernel_initializer='glorot_uniform',\n            padding='same',\n            name=block_name + 'se_reduce_conv2d',\n            use_bias=True\n        )(x)\n\n        x = Swish(name=block_name + 'se_swish')(x)\n\n        x = layers.Conv2D(\n            filters,\n            kernel_size=[1, 1],\n            strides=[1, 1],\n            kernel_initializer='glorot_uniform',\n            padding='same',\n            name=block_name + 'se_expand_conv2d',\n            use_bias=True\n        )(x)\n\n        x = layers.Activation('sigmoid')(x)\n        out = layers.Multiply()([x, inputs])\n        return out\n\n    return block\n\n\nclass DropConnect(layers.Layer):\n\n    def __init__(self, drop_connect_rate, **kwargs):\n        super().__init__(**kwargs)\n        self.drop_connect_rate = drop_connect_rate\n\n    def call(self, inputs, **kwargs):\n        def drop_connect():\n            keep_prob = 1.0 - self.drop_connect_rate\n\n            # Compute drop_connect tensor\n            batch_size = tf.shape(inputs)[0]\n            random_tensor = keep_prob\n            random_tensor += tf.random.uniform([batch_size, 1, 1, 1], dtype=inputs.dtype)\n            binary_tensor = tf.floor(random_tensor)\n            output = tf.math.divide(inputs, keep_prob) * binary_tensor\n            return output\n\n        return K.in_train_phase(drop_connect(), inputs, training=None)\n\n    def get_config(self):\n        config = super().get_config()\n        config['drop_connect_rate'] = self.drop_connect_rate\n        return config\n\n\ndef conv_kernel_initializer(shape=None, dtype=K.floatx()):\n    \"\"\"Initialization for convolutional kernels.\n    The main difference with tf.variance_scaling_initializer is that\n    tf.variance_scaling_initializer uses a truncated normal with an uncorrected\n    standard deviation, whereas here we use a normal distribution. Similarly,\n    tf.contrib.layers.variance_scaling_initializer uses a truncated normal with\n    a corrected standard deviation.\n    Args:\n        shape: shape of variable\n        dtype: dtype of variable\n    Returns:\n        an initialization for the variable\n    \"\"\"\n    return tf.keras.initializers.GlorotUniform()\n    #kernel_height, kernel_width, _, out_filters = shape\n    #fan_out = int(kernel_height * kernel_width * out_filters)\n    #return tf.random.normal(\n    #    shape, mean=0.0, stddev=np.sqrt(2.0 / fan_out), dtype=dtype)\n\n\ndef dense_kernel_initializer(shape, dtype=K.floatx()):\n    init_range = 1.0 / np.sqrt(shape[1])\n    return tf.random.uniform(shape, -init_range, init_range, dtype=dtype)\n\n\ndef MBConvBlock(block_args, global_params, idx, drop_connect_rate=None):\n    filters = block_args.input_filters * block_args.expand_ratio\n    batch_norm_momentum = global_params.batch_norm_momentum\n    batch_norm_epsilon = global_params.batch_norm_epsilon\n    has_se = (block_args.se_ratio is not None) and (0 < block_args.se_ratio <= 1)\n\n    block_name = 'blocks_' + str(idx) + '_'\n\n    def block(inputs):\n        x = inputs\n\n        # Expansion phase\n        if block_args.expand_ratio != 1:\n            expand_conv = layers.Conv2D(filters,\n                                        kernel_size=[1, 1],\n                                        strides=[1, 1],\n                                        kernel_initializer='glorot_uniform',\n                                        padding='same',\n                                        use_bias=False,\n                                        name=block_name + 'expansion_conv2d'\n                                        )(x)\n            bn0 = layers.BatchNormalization(momentum=batch_norm_momentum,\n                                            epsilon=batch_norm_epsilon,\n                                            name=block_name + 'expansion_batch_norm')(expand_conv)\n\n            x = Swish(name=block_name + 'expansion_swish')(bn0)\n\n        # Depth-wise convolution phase\n        kernel_size = block_args.kernel_size\n        depthwise_conv = layers.DepthwiseConv2D(\n            [kernel_size, kernel_size],\n            strides=block_args.strides,\n            depthwise_initializer='glorot_uniform',\n            padding='same',\n            use_bias=False,\n            name=block_name + 'depthwise_conv2d'\n        )(x)\n        bn1 = layers.BatchNormalization(momentum=batch_norm_momentum,\n                                        epsilon=batch_norm_epsilon,\n                                        name=block_name + 'depthwise_batch_norm'\n                                        )(depthwise_conv)\n        x = Swish(name=block_name + 'depthwise_swish')(bn1)\n\n        if has_se:\n            x = SEBlock(block_args, block_name=block_name)(x)\n\n        # Output phase\n        project_conv = layers.Conv2D(\n            block_args.output_filters,\n            kernel_size=[1, 1],\n            strides=[1, 1],\n            kernel_initializer='glorot_uniform',\n            padding='same',\n            name=block_name + 'output_conv2d',\n            use_bias=False)(x)\n        x = layers.BatchNormalization(momentum=batch_norm_momentum,\n                                      epsilon=batch_norm_epsilon,\n                                      name=block_name + 'output_batch_norm'\n                                      )(project_conv)\n        if block_args.id_skip:\n            if all(\n                    s == 1 for s in block_args.strides\n            ) and block_args.input_filters == block_args.output_filters:\n                # only apply drop_connect if skip presents.\n                if drop_connect_rate:\n                    x = DropConnect(drop_connect_rate)(x)\n                x = layers.add([x, inputs])\n\n        return x\n\n    return block\n\n\ndef freeze_efficientunet_first_n_blocks(model, n):\n    mbblock_nr = 0\n    while True:\n        try:\n            model.get_layer('blocks_{}_output_batch_norm'.format(mbblock_nr))\n            mbblock_nr += 1\n        except ValueError:\n            break\n\n    all_block_names = ['blocks_{}_output_batch_norm'.format(i) for i in range(mbblock_nr)]\n    all_block_index = []\n    for idx, layer in enumerate(model.layers):\n        if layer.name == all_block_names[0]:\n            all_block_index.append(idx)\n            all_block_names.pop(0)\n            if len(all_block_names) == 0:\n                break\n    n_blocks = len(all_block_index)\n\n    if n <= 0:\n        print('n is less than or equal to 0, therefore no layer will be frozen.')\n        return\n    if n > n_blocks:\n        raise ValueError(\"There are {} blocks in total, n cannot be greater than {}.\".format(n_blocks, n_blocks))\n\n    idx_of_last_block_to_be_frozen = all_block_index[n - 1]\n    for layer in model.layers[:idx_of_last_block_to_be_frozen + 1]:\n        layer.trainable = False\n\n\ndef freeze_efficientnet(model):\n    for layer in model.layers:\n        layer.trainable = False\n        \ndef unfreeze_efficientnet_first_conv(model):\n    for layer in model.layers:\n        if layer.name == 'stem_conv2d_replaced' or 'batch' in layer.name:\n            layer.trainable = True\n        \ndef unfreeze_efficientunet(model):\n    for layer in model.layers:\n        layer.trainable = True\n        \ndef _efficientnet(input_shape, blocks_args_list, global_params):\n    batch_norm_momentum = global_params.batch_norm_momentum\n    batch_norm_epsilon = global_params.batch_norm_epsilon\n\n    # Stem part\n    model_input = layers.Input(shape=input_shape)\n    #x = tf.expand_dims(model_input, -1)\n    #x = layers.Conv3D(\n    #    filters=CONV3D_FILTERS,\n    #    kernel_size=3,\n    #    strides=1,\n    #    #kernel_initializer='glorot_uniform',\n    #    padding=\"same\",\n    #    name=\"3d_to_2d\"\n    #)(x)\n    #x = tf.squeeze(x)\n    #x = tf.reshape(x,[BATCH_SIZE,PATCH_SIZE, PATCH_SIZE, Z_DIM*CONV3D_FILTERS])\n    #x = x[:,:,:,:,0]\n    \n    x = layers.Conv2D(\n        #filters=round_filters(32, global_params),\n        filters=round_filters(int(32/CONFIG['DIVISOR']), global_params),\n        kernel_size=[3, 3],\n        strides=[2, 2],\n        kernel_initializer='glorot_uniform',\n        padding='same',\n        use_bias=False,\n        name='stem_conv2d'\n    )(model_input)\n\n    x = layers.BatchNormalization(\n        momentum=batch_norm_momentum,\n        epsilon=batch_norm_epsilon,\n        name='stem_batch_norm'\n    )(x)\n\n    x = Swish(name='stem_swish')(x)\n\n    # Blocks part\n    idx = 0\n    drop_rate = global_params.drop_connect_rate\n    n_blocks = sum([blocks_args.num_repeat for blocks_args in blocks_args_list])\n    drop_rate_dx = drop_rate / n_blocks\n\n    for blocks_args in blocks_args_list:\n        assert blocks_args.num_repeat > 0\n        # Update block input and output filters based on depth multiplier.\n        blocks_args = blocks_args._replace(\n            input_filters=round_filters(blocks_args.input_filters, global_params),\n            output_filters=round_filters(blocks_args.output_filters, global_params),\n            num_repeat=round_repeats(blocks_args.num_repeat, global_params)\n        )\n\n        # The first block needs to take care of stride and filter size increase.\n        x = MBConvBlock(blocks_args, global_params, idx, drop_connect_rate=drop_rate_dx * idx)(x)\n        idx += 1\n\n        if blocks_args.num_repeat > 1:\n            blocks_args = blocks_args._replace(input_filters=blocks_args.output_filters, strides=[1, 1])\n\n        for _ in range(blocks_args.num_repeat - 1):\n            x = MBConvBlock(blocks_args, global_params, idx, drop_connect_rate=drop_rate_dx * idx)(x)\n            idx += 1\n\n    # Head part\n    x = layers.Conv2D(\n        filters=round_filters(int(1280/CONFIG['DIVISOR']), global_params),\n        kernel_size=[1, 1],\n        strides=[1, 1],\n        kernel_initializer='glorot_uniform',\n        padding='same',\n        use_bias=False,\n        name='head_conv2d'\n    )(x)\n\n    x = layers.BatchNormalization(\n        momentum=batch_norm_momentum,\n        epsilon=batch_norm_epsilon,\n        name='head_batch_norm'\n    )(x)\n\n    x = Swish(name='head_swish')(x)\n\n    x = layers.GlobalAveragePooling2D(name='global_average_pooling2d')(x)\n\n    if global_params.dropout_rate > 0:\n        x = layers.Dropout(global_params.dropout_rate)(x)\n\n    x = layers.Dense(\n        global_params.num_classes,\n        kernel_initializer='glorot_uniform',\n        activation='softmax',\n        name='head_dense'\n    )(x)\n\n    model = Model(model_input, x)\n\n    return model\n\n\ndef get_model_by_name(model_name, input_shape, classes=1000, pretrained=False):\n    \"\"\"Get an EfficientNet model by its name.\n    \"\"\"\n    blocks_args, global_params = get_efficientnet_params(model_name, override_params={'num_classes': classes})\n    model = _efficientnet(input_shape, blocks_args, global_params)\n\n    try:\n        if pretrained:\n            weights = IMAGENET_WEIGHTS[model_name]\n            weights_path = get_file(\n                weights['name'],\n                weights['url'],\n                cache_subdir='models',\n                md5_hash=weights['md5'],\n            )\n            model.load_weights(weights_path)\n    except KeyError as e:\n        print(\"NOTE: Currently model {} doesn't have pretrained weights, therefore a model with randomly initialized\"\n              \" weights is returned.\".format(e))\n\n    return model\n\ncustom_objects = {\n    \"Swish\": Swish, \n    \"DropConnect\": DropConnect, \n    \"bce_loss\": bce_loss, \n    \"modified_dice_loss\": modified_dice_loss,\n    \"StatefullBinaryFBeta\": StatefullBinaryFBeta\n}\n    \ndef replace_first_conv(model, input_shape):\n    confs = model.get_config()\n    kept_layers = set()\n    for i, l in enumerate(confs['layers']):\n        if i == 0:\n            confs['layers'][0]['config']['batch_input_shape'] = (None,) + input_shape#model.layers[start].input_shape\n            #if i != start:\n                #confs['layers'][0]['name'] += str(random.randint(0, 100000000)) # rename the input layer to avoid conflicts on merge\n            #    confs['layers'][0]['config']['name'] = confs['layers'][0]['name']\n        elif l['name'] == \"stem_conv2d\":\n            confs['layers'][i]['config']['name'] = 'stem_conv2d_replaced'\n        #kept_layers.add(l['name'])\n    # filter layers\n    #layers = [l for l in confs['layers'] if l['name'] in kept_layers]\n    #layers[1]['inbound_nodes'][0][0][0] = layers[0]['name']\n    # set conf\n    #print(confs['layers'][0])\n    #print(confs['layers'][1])\n    #print(confs['layers'][2])\n    #confs['layers'] = layers\n    #confs['input_layers'][0][0] = layers[0]['name']\n    #confs['output_layers'][0][0] = layers[-1]['name']\n    # create new model\n    with keras.utils.custom_object_scope(custom_objects):\n        submodel = tf.keras.Model.from_config(confs)\n    for l in submodel.layers:\n        if l.name == 'stem_conv2d_replaced':\n            continue\n        orig_l = model.get_layer(l.name)\n        if orig_l is not None:\n            l.set_weights(orig_l.get_weights())\n    return submodel\n\ndef _get_efficientnet_encoder(model_name, input_shape, pretrained=False):\n    model = get_model_by_name(model_name, input_shape, pretrained=pretrained)\n    encoder = Model(model.input, model.get_layer('global_average_pooling2d').output)\n    encoder.layers.pop()  # remove GAP layer\n    return encoder\n\ndef _get_efficientnet_encoder_v2(model_name, input_shape, pretrained=False):\n    model = get_model_by_name(model_name, input_shape[:2] + (3,) if pretrained else input_shape, pretrained=pretrained)\n    \n    if pretrained:\n        model = replace_first_conv(model,input_shape)\n        freeze_efficientnet(model)\n        unfreeze_efficientnet_first_conv(model)\n    \n    encoder = Model(model.input, model.get_layer('global_average_pooling2d').output)\n    encoder.layers.pop()  # remove GAP layer\n    return encoder\n\n\ndef get_efficientnet_b0_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b0', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b1_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b1', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b2_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b2', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b3_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b3', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b4_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b4', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b5_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b5', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b6_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b6', input_shape, pretrained=pretrained)\n\n\ndef get_efficientnet_b7_encoder(input_shape, pretrained=False):\n    return _get_efficientnet_encoder_v2('efficientnet-b7', input_shape, pretrained=pretrained)\n\ndef get_blocknr_of_skip_candidates(encoder, verbose=False):\n    \"\"\"\n    Get block numbers of the blocks which will be used for concatenation in the Unet.\n    :param encoder: the encoder\n    :param verbose: if set to True, the shape information of all blocks will be printed in the console\n    :return: a list of block numbers\n    \"\"\"\n    shapes = []\n    candidates = []\n    mbblock_nr = 0\n    while True:\n        try:\n            mbblock = encoder.get_layer('blocks_{}_output_batch_norm'.format(mbblock_nr)).output\n            shape = int(mbblock.shape[1]), int(mbblock.shape[2])\n            if shape not in shapes:\n                shapes.append(shape)\n                candidates.append(mbblock_nr)\n            if verbose:\n                print('blocks_{}_output_shape: {}'.format(mbblock_nr, shape))\n            mbblock_nr += 1\n        except ValueError:\n            break\n    return candidates\n\ndef DoubleConv(filters, kernel_size, initializer='glorot_uniform'):\n\n    def layer(x):\n\n        x = Conv2D(filters, kernel_size, padding='same', use_bias=False, kernel_initializer=initializer)(x)\n        x = BatchNormalization()(x)\n        #x = Activation('relu')(x)\n        x = Conv2D(filters, kernel_size, padding='same', use_bias=False, kernel_initializer=initializer)(x)\n        x = BatchNormalization()(x)\n        #x = Activation('relu')(x)\n\n        return x\n\n    return layer\n\n\ndef UpSampling2D_block(filters, kernel_size=(3, 3), upsample_rate=(2, 2), interpolation='bilinear',\n                       initializer='glorot_uniform', skip=None):\n    def layer(input_tensor):\n\n        x = UpSampling2D(size=upsample_rate, interpolation=interpolation)(input_tensor)\n\n        if skip is not None:\n            x = Concatenate()([x, skip])\n\n        x = DoubleConv(filters, kernel_size, initializer=initializer)(x)\n\n        return x\n    return layer\n\n\ndef Conv2DTranspose_block(filters, kernel_size=(3, 3), transpose_kernel_size=(2, 2), upsample_rate=(2, 2),\n                          initializer='glorot_uniform', skip=None):\n    def layer(input_tensor):\n\n        x = Conv2DTranspose(filters, transpose_kernel_size, strides=upsample_rate, padding='same')(input_tensor)\n\n        if skip is not None:\n            x = Concatenate()([x, skip])\n\n        x = DoubleConv(filters, kernel_size, initializer=initializer)(x)\n\n        return x\n\n    return layer\n\n\n# noinspection PyTypeChecker\ndef _get_efficient_unet(encoder, out_channels=2, block_type='upsampling', concat_input=True):\n    MBConvBlocks = []\n\n    skip_candidates = get_blocknr_of_skip_candidates(encoder)\n\n    for mbblock_nr in skip_candidates:\n        mbblock = encoder.get_layer('blocks_{}_output_batch_norm'.format(mbblock_nr)).output\n        MBConvBlocks.append(mbblock)\n\n    # delete the last block since it won't be used in the process of concatenation\n    MBConvBlocks.pop()\n\n    input_ = encoder.input\n    head = encoder.get_layer('head_swish').output\n    blocks = [input_] + MBConvBlocks + [head]\n\n    if block_type == 'upsampling':\n        UpBlock = UpSampling2D_block\n    else:\n        UpBlock = Conv2DTranspose_block\n\n    o = blocks.pop()\n    o = UpBlock(int(512/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=blocks.pop())(o)\n    o = UpBlock(int(256/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=blocks.pop())(o)\n    o = UpBlock(int(128/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=blocks.pop())(o)\n    o = UpBlock(int(64/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=blocks.pop())(o)\n    if concat_input:\n        o = UpBlock(int(32/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=blocks.pop())(o)\n    else:\n        o = UpBlock(int(32/CONFIG['DIVISOR']), initializer='glorot_uniform', skip=None)(o)\n    o = Conv2D(out_channels, \n               (1, 1), \n               padding='same', \n               kernel_initializer='glorot_uniform',\n               activation=\"sigmoid\" if CONFIG[\"SIGMOID_OUTPUT\"] else None\n              )(o)\n\n    model = Model(encoder.input, o)\n\n    return model\n\n\ndef get_efficient_unet_b0(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B0 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B0 model\n    \"\"\"\n    encoder = get_efficientnet_b0_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\ndef get_efficient_unet_b1(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B1 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B1 model\n    \"\"\"\n    encoder = get_efficientnet_b1_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b2(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B2 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B2 model\n    \"\"\"\n    encoder = get_efficientnet_b2_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b3(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B3 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B3 model\n    \"\"\"\n    encoder = get_efficientnet_b3_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b4(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B4 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B4 model\n    \"\"\"\n    encoder = get_efficientnet_b4_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b5(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B5 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B5 model\n    \"\"\"\n    encoder = get_efficientnet_b5_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b6(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B6 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B6 model\n    \"\"\"\n    encoder = get_efficientnet_b6_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model\n\n\ndef get_efficient_unet_b7(input_shape, out_channels=2, pretrained=False, block_type='transpose', concat_input=True):\n    \"\"\"Get a Unet model with Efficient-B7 encoder\n    :param input_shape: shape of input (cannot have None element)\n    :param out_channels: the number of output channels\n    :param pretrained: True for ImageNet pretrained weights\n    :param block_type: \"upsampling\" to use UpSampling layer, otherwise use Conv2DTranspose layer\n    :param concat_input: if True, input image will be concatenated with the last conv layer\n    :return: an EfficientUnet_B7 model\n    \"\"\"\n    encoder = get_efficientnet_b7_encoder(input_shape, pretrained=pretrained)\n    model = _get_efficient_unet(encoder, out_channels, block_type=block_type, concat_input=concat_input)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:48.178487Z","iopub.execute_input":"2023-04-26T09:03:48.178799Z","iopub.status.idle":"2023-04-26T09:03:48.278522Z","shell.execute_reply.started":"2023-04-26T09:03:48.17877Z","shell.execute_reply":"2023-04-26T09:03:48.27756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models_dict = {\n    'efficientnetb0-unet': get_efficient_unet_b0,\n    'efficientnetb1-unet': get_efficient_unet_b1,\n    'efficientnetb2-unet': get_efficient_unet_b2,\n    'efficientnetb3-unet': get_efficient_unet_b3,\n    'efficientnetb4-unet': get_efficient_unet_b4,\n    'efficientnetb5-unet': get_efficient_unet_b5,\n    'efficientnetb6-unet': get_efficient_unet_b6,\n    'efficientnetb7-unet': get_efficient_unet_b7,\n}","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:48.280055Z","iopub.execute_input":"2023-04-26T09:03:48.28057Z","iopub.status.idle":"2023-04-26T09:03:48.286075Z","shell.execute_reply.started":"2023-04-26T09:03:48.280529Z","shell.execute_reply":"2023-04-26T09:03:48.284986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_shape = (PATCH_SIZE, PATCH_SIZE, Z_DIM)\nwith strategy.scope():\n    model = models_dict[CONFIG['model_name']](input_shape,out_channels=1,pretrained=CONFIG['PRETRAINED'])\n    print([layer.trainable for layer in model.layers])\n    \n    if CONFIG['optimizer'] == 'ADAM':\n        opt = keras.optimizers.Adam(learning_rate=CONFIG[\"learning_rate\"]/10,clipnorm=1.)\n    elif CONFIG['optimizer'] == 'SGD':\n        opt = keras.optimizers.SGD(learning_rate=CONFIG[\"learning_rate\"]/10,clipnorm=1.)\n\n    model.compile(optimizer=opt, \n              loss=losses_map[CONFIG[\"LOSS\"]],\n              metrics=[\n                tf.keras.metrics.BinaryAccuracy(name=\"binary_accuracy\",threshold=THRESHOLD),\n                tf.keras.metrics.BinaryIoU(name=\"binary_iou\",threshold=THRESHOLD),\n                StatefullBinaryFBeta(name=\"fbeta_score\",beta=0.5,threshold=THRESHOLD)\n                ]\n             )","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:48.287555Z","iopub.execute_input":"2023-04-26T09:03:48.288211Z","iopub.status.idle":"2023-04-26T09:03:50.424305Z","shell.execute_reply.started":"2023-04-26T09:03:48.288164Z","shell.execute_reply":"2023-04-26T09:03:50.423206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training process\n\nInitialize W&B run.","metadata":{}},{"cell_type":"code","source":"if not PROD:\n    run = wandb.init(project='vesuvius', \n                     config=CONFIG,\n                     group='EfficientNet-UNet', \n                     job_type='train')","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:03:50.425935Z","iopub.execute_input":"2023-04-26T09:03:50.42629Z","iopub.status.idle":"2023-04-26T09:04:21.874771Z","shell.execute_reply.started":"2023-04-26T09:03:50.426253Z","shell.execute_reply":"2023-04-26T09:04:21.873747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Create learning rate scheduler.","metadata":{}},{"cell_type":"code","source":"def scheduler(epoch, lr):\n    if epoch < WARMUP_EPOCHS:\n        return LEARNING_RATE*(epoch+1)/WARMUP_EPOCHS\n    elif epoch < DECAY_EPOCHS:\n        return LEARNING_RATE\n    else:\n        return lr * tf.math.exp(-0.1)\n\nlr_scheduler = tf.keras.callbacks.LearningRateScheduler(scheduler)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:04:21.879137Z","iopub.execute_input":"2023-04-26T09:04:21.881872Z","iopub.status.idle":"2023-04-26T09:04:21.891244Z","shell.execute_reply.started":"2023-04-26T09:04:21.881827Z","shell.execute_reply":"2023-04-26T09:04:21.890222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Create callbacks.","metadata":{}},{"cell_type":"code","source":"callbacks = []\n    \ncallbacks.append(lr_scheduler)\n\nif not PROD:\n    callbacks.append(\n        WandbCallback(\n            log_weights=True,\n            save_model=False,\n            save_graph=False,\n            #log_gradients=True,\n            #log_batch_frequency=100,\n            #training_data=train_ds.take(BATCH_SIZE)\n        )\n    )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Train for EPOCHS.","metadata":{}},{"cell_type":"code","source":"model.fit(\n    x=train_ds, \n    validation_data=val_ds, \n    epochs=CONFIG[\"epochs\"], \n    steps_per_epoch=CONFIG[\"steps_per_epoch\"],\n    validation_steps=total_validation_samples // BATCH_SIZE,\n    callbacks=callbacks,\n    use_multiprocessing=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:04:21.895855Z","iopub.execute_input":"2023-04-26T09:04:21.898847Z","iopub.status.idle":"2023-04-26T09:08:51.282921Z","shell.execute_reply.started":"2023-04-26T09:04:21.898808Z","shell.execute_reply":"2023-04-26T09:08:51.281608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Save trained model.","metadata":{}},{"cell_type":"code","source":"save_locally = tf.saved_model.SaveOptions(experimental_io_device='/job:localhost')\nmodel.save(\"model.keras\", options=save_locally)\n\nif not PROD:\n    # Save model as Model Artifact\n    artifact = wandb.Artifact(name=CONFIG['model_name'], type='model')\n    artifact.add_file(\"model.keras\")\n    run.log_artifact(artifact)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation","metadata":{}},{"cell_type":"markdown","source":"## Helper functions","metadata":{}},{"cell_type":"code","source":"def select_best_threshold(predicted_labels,true_labels):\n    beta_scores = []\n    thresholds = []\n    for th in np.arange(0.1,1,0.05):\n        if CONFIG[\"SIGMOID_OUTPUT\"]:\n            predicted_labels_th = np.where(predicted_labels > th, 1, 0).astype(np.uint8)\n        else:\n            predicted_labels_th = np.where(tf.keras.activations.sigmoid(predicted_labels) > th, 1, 0).astype(np.uint8)\n        beta_score = binary_fbeta(ytrue=true_labels,ypred=predicted_labels_th,beta=0.5)\n        beta_scores.append(beta_score)\n        thresholds.append(th)\n    best_beta_score_idx = np.argmax(beta_scores)\n    \n    return thresholds[best_beta_score_idx], beta_scores, thresholds\n\ndef classical_morphological_postprocess(predicted_labels,true_labels):\n    beta_scores = []\n    kernel_sizes = []\n    for k in range(1,50,2):\n        kernelSize = (k,k)\n        kernel = cv2.getStructuringElement(cv2.MORPH_RECT, kernelSize)\n        closing = cv2.morphologyEx(predicted_labels, cv2.MORPH_CLOSE, kernel)\n        beta_score = binary_fbeta(ytrue=true_labels,ypred=np.expand_dims(closing,axis=-1),beta=0.5)\n        beta_scores.append(beta_score)\n        kernel_sizes.append(k)\n    best_beta_score_idx = np.argmax(beta_scores)\n    \n    return kernel_sizes[best_beta_score_idx]","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:08:51.29596Z","iopub.execute_input":"2023-04-26T09:08:51.300263Z","iopub.status.idle":"2023-04-26T09:08:51.318054Z","shell.execute_reply.started":"2023-04-26T09:08:51.300177Z","shell.execute_reply":"2023-04-26T09:08:51.316671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Print prediction on validation set","metadata":{}},{"cell_type":"code","source":"best_val_thresholds = []\nbest_val_kernel_sizes = []\n\nfor i in range(len(dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_volume\"])):\n    val_ds = make_tf_dataset(\n        make_iterated_data_generator(\n            dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_volume\"][i], \n            dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_mask\"][i],\n            dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_labels\"][i],\n            return_locations=True),\n        labeled=True,\n        return_locations=True\n    )\n    \n    predicted_labels = np.zeros(dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_volume\"][i].shape[:2] + (1,), dtype=\"float32\")\n    true_labels = np.zeros(dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_labels\"][i].shape[:2] + (1,), dtype=\"float32\")\n    predictions_map_counts = np.zeros(dev_folds[CONFIG[\"VAL_FOLD\"]][\"validation_volume\"][i].shape[:2] + (1,), dtype=\"int8\")\n    \n    for patch_batch, y_true_batch, loc_batch in tqdm(val_ds):\n        predictions = model.predict_on_batch(patch_batch)\n        #print(loc_batch)\n        for (x, y), pred, y_true in zip(loc_batch, predictions, y_true_batch):\n            x = int(x.numpy())\n            y = int(y.numpy())\n            predicted_labels[x-PATCH_HALFSIZE:x+PATCH_HALFSIZE,y-PATCH_HALFSIZE:y+PATCH_HALFSIZE,:] += pred\n            predictions_map_counts[x - PATCH_HALFSIZE : x + PATCH_HALFSIZE, y - PATCH_HALFSIZE : y + PATCH_HALFSIZE, :] += 1  \n            true_labels[x-PATCH_HALFSIZE:x+PATCH_HALFSIZE,y-PATCH_HALFSIZE:y+PATCH_HALFSIZE,:] += y_true\n            \n    del val_ds\n    \n    predicted_labels /= (predictions_map_counts + 1e-7)\n    true_labels /= (predictions_map_counts + 1e-7)\n    true_labels = np.where(true_labels > THRESHOLD, 1, 0)\n    \n    print(f\"Min/Max predicted_labels : {np.min(predicted_labels)}/{np.max(predicted_labels)}\")\n    train_labels_min = np.min(true_labels)\n    train_labels_max = np.max(true_labels)\n    print(f\"Min/Max train_labels : {train_labels_min}/{train_labels_max}\")\n    fig, (ax1, ax2, ax3, ax4) = plt.subplots(1, 4)\n    ax1.set_title(f\"True\")\n    ax1.imshow(true_labels, cmap='gray')\n    ax2.set_title(f\"Pred\")\n    ax2.imshow(predicted_labels, cmap='gray')\n    \n    best_threshold, beta_scores, thresholds = select_best_threshold(predicted_labels,true_labels)\n    best_val_thresholds.append(best_threshold)\n    \n    ax3.set_title(f\"th = {best_threshold:.2f}\")\n    if CONFIG[\"SIGMOID_OUTPUT\"]:\n        predicted_labels_th = np.where(predicted_labels > best_threshold, 1, 0).astype(np.uint8)\n    else:\n        predicted_labels_th = np.where(tf.keras.activations.sigmoid(predicted_labels) > best_threshold, 1, 0).astype(np.uint8)\n    print(f\"Min/Max predicted_labels_th : {np.min(predicted_labels_th)}/{np.max(predicted_labels_th)}\")\n    ax3.imshow(predicted_labels_th, cmap='gray')\n    \n    \n    beta_score = binary_fbeta(ytrue=true_labels,ypred=predicted_labels_th,beta=0.5)\n    print(f\"Beta score (DOWNSAMPLING = {DOWNSAMPLING}): {beta_score}\")\n    \n    best_kernel_size = classical_morphological_postprocess(predicted_labels_th,true_labels)\n    best_val_kernel_sizes.append(best_kernel_size)\n    kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (best_kernel_size,best_kernel_size))\n    predicted_labels_closed = cv2.morphologyEx(predicted_labels_th, cv2.MORPH_CLOSE, kernel)\n    \n    beta_score = binary_fbeta(ytrue=true_labels,ypred=np.expand_dims(predicted_labels_closed,axis=-1),beta=0.5)\n    print(f\"Beta score (CLOSED): {beta_score}\")\n    \n    ax4.set_title(f\"k = {best_kernel_size}\")\n    ax4.imshow(predicted_labels_closed, cmap='gray')\n    plt.show()\n    \n    del predicted_labels\n    del predictions_map_counts\n    del predicted_labels_th\n    gc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:08:51.325458Z","iopub.execute_input":"2023-04-26T09:08:51.329603Z","iopub.status.idle":"2023-04-26T09:09:29.045472Z","shell.execute_reply.started":"2023-04-26T09:08:51.329552Z","shell.execute_reply":"2023-04-26T09:09:29.044478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Select best post-prediction parameters\n\nSelects the best threshold and kernel size to maximize validation score.","metadata":{}},{"cell_type":"code","source":"print(\"Selecting best threshold and kernel size for morphological closing...\")\nfor i in range(len(best_val_thresholds)):\n    print(f\"\\tth = {best_val_thresholds[i]}\")\n    print(f\"\\tkernel_size = {best_val_kernel_sizes[i]}\")\n    \nbest_threshold = np.mean(best_val_thresholds)\nbest_kernel_size = int(np.mean(best_val_kernel_sizes))\nprint(f\"th = {best_threshold}\")\nprint(f\"kernel_size = {best_kernel_size}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:09:29.052876Z","iopub.execute_input":"2023-04-26T09:09:29.055296Z","iopub.status.idle":"2023-04-26T09:09:29.10474Z","shell.execute_reply.started":"2023-04-26T09:09:29.055252Z","shell.execute_reply":"2023-04-26T09:09:29.103632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Print prediction on entire input volumes\n\nJust as a sanity check, see the predicted masks for the entire available volumes.","metadata":{}},{"cell_type":"code","source":"for i in range(len(prod_data[\"train_volumes\"])):\n    print(\"===============================\")\n    print(f\"Training volume {i}:\")\n    train_ds = make_tf_dataset(\n        make_iterated_data_generator(\n            prod_data[\"train_volumes\"][i], \n            prod_data[\"train_masks\"][i],\n            prod_data[\"train_labels\"][i],\n            return_locations=True),\n        labeled=True,\n        return_locations=True\n    )\n    \n    predicted_labels = np.zeros(prod_data[\"train_volumes\"][i].shape[:2] + (1,), dtype=\"float32\")\n    predictions_map_counts = np.zeros(prod_data[\"train_volumes\"][i].shape[:2] + (1,), dtype=\"int8\")\n    true_labels = np.zeros(prod_data[\"train_labels\"][i].shape[:2] + (1,), dtype=\"float32\")\n    for patch_batch, y_true_batch, loc_batch in tqdm(train_ds):\n        predictions = model.predict_on_batch(patch_batch)\n        for (x, y), pred, y_true in zip(loc_batch, predictions, y_true_batch):\n            x = int(x.numpy())\n            y = int(y.numpy())\n            predicted_labels[x-PATCH_HALFSIZE:x+PATCH_HALFSIZE,y-PATCH_HALFSIZE:y+PATCH_HALFSIZE,:] += pred\n            predictions_map_counts[x - PATCH_HALFSIZE : x + PATCH_HALFSIZE, y - PATCH_HALFSIZE : y + PATCH_HALFSIZE, :] += 1  \n            true_labels[x-PATCH_HALFSIZE:x+PATCH_HALFSIZE,y-PATCH_HALFSIZE:y+PATCH_HALFSIZE,:] += y_true\n            \n    print(f\"\\tMin/Max predictions_map_counts : {np.min(predictions_map_counts)}/{np.max(predictions_map_counts)}\")\n    \n    predicted_labels /= (predictions_map_counts + 1e-7)\n    true_labels /= (predictions_map_counts + 1e-7)\n    true_labels = np.where(true_labels > THRESHOLD, 1, 0)\n    \n    print(f\"\\tMin/Max predicted_labels : {np.min(predicted_labels)}/{np.max(predicted_labels)}\")\n    fig, (ax1, ax2, ax3, ax4) = plt.subplots(1, 4)\n    ax1.set_title(f\"True \")\n    ax1.imshow(true_labels, cmap='gray')\n    ax2.set_title(f\"Pred\")\n    ax2.imshow(predicted_labels, cmap='gray')\n    ax3.set_title(f\"th = {best_threshold:.2f}\")\n    if CONFIG[\"SIGMOID_OUTPUT\"]:\n        predicted_labels_th = np.where(predicted_labels > best_threshold, 1, 0).astype(np.uint8)\n    else:\n        predicted_labels_th = np.where(tf.keras.activations.sigmoid(predicted_labels) > best_threshold, 1, 0).astype(np.uint8)\n    print(f\"\\tMin/Max predicted_labels_th : {np.min(predicted_labels_th)}/{np.max(predicted_labels_th)}\")\n    ax3.imshow(predicted_labels_th, cmap='gray')\n    \n    beta_score = binary_fbeta(ytrue=true_labels,ypred=predicted_labels_th,beta=0.5)\n    \n    kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (best_kernel_size,best_kernel_size))\n    predicted_labels_closed = cv2.morphologyEx(predicted_labels_th, cv2.MORPH_CLOSE, kernel)\n    beta_score_closed = binary_fbeta(ytrue=true_labels,ypred=np.expand_dims(predicted_labels_closed,axis=-1),beta=0.5)\n    \n    ax4.set_title(f\"k = {best_kernel_size}\")\n    ax4.imshow(predicted_labels_closed, cmap='gray')\n    \n    orig_labels = cv2.imread(DATA_DIR + f\"/train/{i+1}/inklabels.png\", 0)\n    orig_labels = np.asarray(orig_labels) / 255.0\n    predicted_labels_closed_orig_size = resize_to_original(predicted_labels_closed,orig_labels)\n    beta_score_orig_size = binary_fbeta(ytrue=orig_labels,ypred=predicted_labels_closed_orig_size,beta=0.5)\n    \n    \n    print(f\"Beta score (DOWNSAMPLING = {DOWNSAMPLING}): {beta_score}\")\n    print(f\"Beta score (CLOSED): {beta_score_closed}\")\n    print(f\"Beta score (ORIGINAL SIZE): {beta_score_orig_size}\")\n    \n    plt.show()\n    \n    if not PROD:\n        wandb.log({f\"y_true_{i}\" : wandb.Image(prod_data[\"train_labels\"][i])})\n        wandb.log({f\"y_pred_{i}\" : wandb.Image(predicted_labels)})\n        wandb.log({f\"y_pred_{i}_th\" : wandb.Image(predicted_labels_th)})\n        wandb.log({f\"y_pred_{i}_closed\" : wandb.Image(predicted_labels_closed)})\n","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:09:29.109604Z","iopub.execute_input":"2023-04-26T09:09:29.112065Z","iopub.status.idle":"2023-04-26T09:10:03.773549Z","shell.execute_reply.started":"2023-04-26T09:09:29.112024Z","shell.execute_reply":"2023-04-26T09:10:03.771465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not PROD:\n    # Close W&B run\n    run.finish()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:10:03.779671Z","iopub.execute_input":"2023-04-26T09:10:03.782164Z","iopub.status.idle":"2023-04-26T09:10:11.501725Z","shell.execute_reply.started":"2023-04-26T09:10:03.782122Z","shell.execute_reply":"2023-04-26T09:10:11.500729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#del train_ds\n#if not PROD:\n#    del val_ds\ndel volume_1\ndel volume_2\ndel volume_3\ndel mask_1\ndel mask_2\ndel mask_3\ndel labels_1\ndel labels_2\ndel labels_3\n\n# Manually trigger garbage collection\nkeras.backend.clear_session()\nimport gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T09:10:11.503211Z","iopub.execute_input":"2023-04-26T09:10:11.503568Z","iopub.status.idle":"2023-04-26T09:10:12.190111Z","shell.execute_reply.started":"2023-04-26T09:10:11.50353Z","shell.execute_reply":"2023-04-26T09:10:12.188932Z"},"trusted":true},"execution_count":null,"outputs":[]}]}