{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":47317,"databundleVersionId":5799376,"sourceType":"competition"}],"dockerImageVersionId":30408,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Keras starter kit [full training set, UNet]","metadata":{}},{"cell_type":"markdown","source":"## Setup","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers, models\nfrom tensorflow.keras import backend as K\nimport numpy as np\n\nimport glob\nimport time\nimport PIL.Image as Image\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom tqdm import tqdm\n\n# Data config\nDATA_DIR = '/kaggle/input/vesuvius-challenge-ink-detection/'\nBUFFER = 32  # Half-size of papyrus patches we'll use as model inputs\nZ_DIM = 20   # Number of slices in the z direction. Max value is 64 - Z_START\nZ_START = 16  # Offset of slices in the z direction\nSHARED_HEIGHT = 4000  # Height to resize all papyrii\n\n# Model config\nBATCH_SIZE = 32\nUSE_MIXED_PRECISION = False\nUSE_JIT_COMPILE = False","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:10:59.896069Z","iopub.execute_input":"2024-06-07T09:10:59.896761Z","iopub.status.idle":"2024-06-07T09:11:07.906492Z","shell.execute_reply.started":"2024-06-07T09:10:59.896729Z","shell.execute_reply":"2024-06-07T09:11:07.905375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(Image.open(DATA_DIR + \"/train/1/ir.png\"), cmap=\"gray\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:07.908432Z","iopub.execute_input":"2024-06-07T09:11:07.909090Z","iopub.status.idle":"2024-06-07T09:11:11.094508Z","shell.execute_reply.started":"2024-06-07T09:11:07.909054Z","shell.execute_reply":"2024-06-07T09:11:11.093365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load up the training data","metadata":{}},{"cell_type":"code","source":"def resize(img):\n    current_width, current_height = img.size\n    aspect_ratio = current_width / current_height\n    new_width = int(SHARED_HEIGHT * aspect_ratio)\n    new_size = (new_width, SHARED_HEIGHT)\n    img = img.resize(new_size)\n    return img\n\ndef load_mask(split, index):\n    img = Image.open(f\"{DATA_DIR}/{split}/{index}/mask.png\").convert('1')\n    img = resize(img)\n    return tf.convert_to_tensor(img, dtype=\"bool\")\n\ndef load_labels(split, index):\n    img = Image.open(f\"{DATA_DIR}/{split}/{index}/inklabels.png\")\n    img = resize(img)\n    return tf.convert_to_tensor(img, dtype=\"bool\")\n\nmask = load_mask(split=\"train\", index=1)\nlabels = load_labels(split=\"train\", index=1)\n\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.set_title(\"mask.png\")\nax1.imshow(mask, cmap='gray')\nax2.set_title(\"inklabels.png\")\nax2.imshow(labels, cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:11.095909Z","iopub.execute_input":"2024-06-07T09:11:11.096321Z","iopub.status.idle":"2024-06-07T09:11:15.216386Z","shell.execute_reply.started":"2024-06-07T09:11:11.096285Z","shell.execute_reply":"2024-06-07T09:11:15.215197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask_test_a = load_mask(split=\"test\", index=\"a\")\nmask_test_b = load_mask(split=\"test\", index=\"b\")\n\nmask_train_1 = load_mask(split=\"train\", index=1)\nlabels_train_1 = load_labels(split=\"train\", index=1)\n\nmask_train_2 = load_mask(split=\"train\", index=2)\nlabels_train_2 = load_labels(split=\"train\", index=2)\n\nmask_train_3 = load_mask(split=\"train\", index=3)\nlabels_train_3 = load_labels(split=\"train\", index=3)\n\nprint(f\"mask_test_a: {mask_test_a.shape}\")\nprint(f\"mask_test_b: {mask_test_b.shape}\")\nprint(\"-\")\nprint(f\"mask_train_1: {mask_train_1.shape}\")\nprint(f\"labels_train_1: {labels_train_1.shape}\")\nprint(\"-\")\nprint(f\"mask_train_2: {mask_train_2.shape}\")\nprint(f\"labels_train_2: {labels_train_2.shape}\")\nprint(\"-\")\nprint(f\"mask_train_3: {mask_train_3.shape}\")\nprint(f\"labels_train_3: {labels_train_3.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:15.219704Z","iopub.execute_input":"2024-06-07T09:11:15.220169Z","iopub.status.idle":"2024-06-07T09:11:17.058233Z","shell.execute_reply.started":"2024-06-07T09:11:15.220122Z","shell.execute_reply":"2024-06-07T09:11:17.057146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, (ax1, ax2, ax3) = plt.subplots(1, 3)\n\nax1.set_title(\"labels_train_1\")\nax1.imshow(labels_train_1, cmap='gray')\n\nax2.set_title(\"labels_train_2\")\nax2.imshow(labels_train_2, cmap='gray')\n\nax3.set_title(\"labels_train_3\")\nax3.imshow(labels_train_3, cmap='gray')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:17.059479Z","iopub.execute_input":"2024-06-07T09:11:17.059785Z","iopub.status.idle":"2024-06-07T09:11:18.602557Z","shell.execute_reply.started":"2024-06-07T09:11:17.059754Z","shell.execute_reply":"2024-06-07T09:11:18.601485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_volume(split, index):\n    # Load the 3d x-ray scan, one slice at a time\n    z_slices_fnames = sorted(glob.glob(f\"{DATA_DIR}/{split}/{index}/surface_volume/*.tif\"))[Z_START:Z_START + Z_DIM]\n    z_slices = []\n    for z, filename in  tqdm(enumerate(z_slices_fnames)):\n        img = Image.open(filename)\n        img = resize(img)\n        z_slice = np.array(img, dtype=\"float32\")\n        z_slices.append(z_slice)\n    return tf.stack(z_slices, axis=-1)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:18.604024Z","iopub.execute_input":"2024-06-07T09:11:18.604447Z","iopub.status.idle":"2024-06-07T09:11:18.612020Z","shell.execute_reply.started":"2024-06-07T09:11:18.604400Z","shell.execute_reply":"2024-06-07T09:11:18.611001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"volume_train_1 = load_volume(split=\"train\", index=1)\nprint(f\"volume_train_1: {volume_train_1.shape}, {volume_train_1.dtype}\")\n\nvolume_train_2 = load_volume(split=\"train\", index=2)\nprint(f\"volume_train_2: {volume_train_2.shape}, {volume_train_2.dtype}\")\n\nvolume_train_3 = load_volume(split=\"train\", index=3)\nprint(f\"volume_train_3: {volume_train_3.shape}, {volume_train_3.dtype}\")\n\nvolume = tf.concat([volume_train_1, volume_train_2, volume_train_3], axis=1)\nprint(f\"total volume: {volume.shape}\")\n\ndel volume_train_1\ndel volume_train_2\ndel volume_train_3","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:11:18.613388Z","iopub.execute_input":"2024-06-07T09:11:18.613754Z","iopub.status.idle":"2024-06-07T09:15:56.984138Z","shell.execute_reply.started":"2024-06-07T09:11:18.613716Z","shell.execute_reply":"2024-06-07T09:15:56.983075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = tf.concat([labels_train_1, labels_train_2, labels_train_3], axis=1)\nprint(f\"labels: {labels.shape}, {labels.dtype}\")\n\nmask = tf.concat([mask_train_1, mask_train_2, mask_train_3], axis=1)\nprint(f\"mask: {mask.shape}, {mask.dtype}\")\n\n# Free up memory\ndel labels_train_1\ndel labels_train_2\ndel labels_train_3\ndel mask_train_1\ndel mask_train_2\ndel mask_train_3","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:15:56.985616Z","iopub.execute_input":"2024-06-07T09:15:56.985935Z","iopub.status.idle":"2024-06-07T09:15:56.994746Z","shell.execute_reply.started":"2024-06-07T09:15:56.985905Z","shell.execute_reply":"2024-06-07T09:15:56.993681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize the training data\n\nIn this case, not very informative. But remember to always visualize what you're training on, as a sanity check!","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 4, figsize=(15, 3))\nfor z, ax in enumerate(axes):\n    ax.imshow(volume[:, :, z], cmap='gray')\n    ax.set_xticks([]); ax.set_yticks([])\nfig.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:15:56.995909Z","iopub.execute_input":"2024-06-07T09:15:56.996282Z","iopub.status.idle":"2024-06-07T09:16:01.324003Z","shell.execute_reply.started":"2024-06-07T09:15:56.996242Z","shell.execute_reply":"2024-06-07T09:16:01.323085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Selection a validation holdout area\n\nWe set aside some fraction of the input to validate our model on.","metadata":{}},{"cell_type":"code","source":"val_location = (2000, 3500)\nval_zone_size = (600, 2000)\n\nfig, ax = plt.subplots()\nax.imshow(labels)\npatch = patches.Rectangle([val_location[1], val_location[0]], val_zone_size[1], val_zone_size[0], linewidth=2, edgecolor='g', facecolor='none')\nax.add_patch(patch)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:01.329176Z","iopub.execute_input":"2024-06-07T09:16:01.329849Z","iopub.status.idle":"2024-06-07T09:16:02.858125Z","shell.execute_reply.started":"2024-06-07T09:16:01.329812Z","shell.execute_reply":"2024-06-07T09:16:02.857086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create a dataset that samples random locations in the input volume\n\nOur training dataset will grab random patches within the masked area and outside of the validation area.","metadata":{}},{"cell_type":"code","source":"def sample_random_location(shape):\n    random_train_x = tf.random.uniform(shape=(), minval=BUFFER, maxval=shape[0] - BUFFER - 1, dtype=\"int32\")\n    random_train_y = tf.random.uniform(shape=(), minval=BUFFER, maxval=shape[1] - BUFFER - 1, dtype=\"int32\")\n    random_train_location = tf.stack([random_train_x, random_train_y])\n    return random_train_location\n\ndef is_in_masked_zone(location, mask):\n    return mask[location[0], location[1]]\n\nsample_random_location_train = lambda x: sample_random_location(mask.shape)\nis_in_mask_train = lambda x: is_in_masked_zone(x, mask)\n\ndef is_in_val_zone(location, val_location, val_zone_size):\n    x = location[0]\n    y = location[1]\n    x_match = val_location[0] - BUFFER <= x <= val_location[0] + val_zone_size[0] + BUFFER\n    y_match = val_location[1] - BUFFER <= y <= val_location[1] + val_zone_size[1] + BUFFER\n    return x_match and y_match\n\ndef is_proper_train_location(location):\n    return not is_in_val_zone(location, val_location, val_zone_size) and is_in_mask_train(location)\n\ntrain_locations_ds = tf.data.Dataset.from_tensor_slices([0]).repeat().map(sample_random_location_train, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_locations_ds = train_locations_ds.filter(is_proper_train_location)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:02.859830Z","iopub.execute_input":"2024-06-07T09:16:02.860640Z","iopub.status.idle":"2024-06-07T09:16:03.235326Z","shell.execute_reply.started":"2024-06-07T09:16:02.860597Z","shell.execute_reply":"2024-06-07T09:16:03.234087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize some training patches\n\nSanity check visually that our patches are where they should be.","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots()\nax.imshow(labels)\n\nfor x, y in train_locations_ds.take(600):\n    patch = patches.Rectangle([y - BUFFER, x - BUFFER], 2 * BUFFER, 2 * BUFFER, linewidth=2, edgecolor='r', facecolor='none')\n    ax.add_patch(patch)\n\nval_patch = patches.Rectangle([val_location[1], val_location[0]], val_zone_size[1], val_zone_size[0], linewidth=2, edgecolor='g', facecolor='none')\nax.add_patch(val_patch)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:03.237004Z","iopub.execute_input":"2024-06-07T09:16:03.237462Z","iopub.status.idle":"2024-06-07T09:16:07.691724Z","shell.execute_reply.started":"2024-06-07T09:16:03.237415Z","shell.execute_reply":"2024-06-07T09:16:07.690691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create training dataset that yields random subvolumes and their labels","metadata":{}},{"cell_type":"code","source":"def extract_subvolume(location, volume):\n    x = location[0]\n    y = location[1]\n    subvolume = volume[x-BUFFER:x+BUFFER, y-BUFFER:y+BUFFER, :]\n    subvolume = tf.cast(subvolume, dtype=\"float32\") / 65535.\n    return subvolume\n\ndef extract_labels(location, labels):\n    x = location[0]\n    y = location[1]\n    label = labels[x-BUFFER:x+BUFFER, y-BUFFER:y+BUFFER]\n    label = tf.cast(label, dtype=\"float32\")\n    label = tf.expand_dims(label, axis=-1)\n    return label\n\ndef extract_subvolume_and_label(location):\n    subvolume = extract_subvolume(location, volume)\n    label = extract_labels(location, labels)\n    return subvolume, label\n\nshuffle_buffer_size = BATCH_SIZE * 4\n\ntrain_ds = train_locations_ds.map(extract_subvolume_and_label, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.prefetch(tf.data.AUTOTUNE).batch(BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:07.693074Z","iopub.execute_input":"2024-06-07T09:16:07.693461Z","iopub.status.idle":"2024-06-07T09:16:10.085469Z","shell.execute_reply.started":"2024-06-07T09:16:07.693422Z","shell.execute_reply":"2024-06-07T09:16:10.084569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for subvolume_batch, label_batch in train_ds.take(1):\n    print(f\"subvolume shape: {subvolume_batch.shape[1:]}\")\n    print(f\"label_batch shape: {label_batch.shape[1:]}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:10.086837Z","iopub.execute_input":"2024-06-07T09:16:10.087239Z","iopub.status.idle":"2024-06-07T09:16:13.025267Z","shell.execute_reply.started":"2024-06-07T09:16:10.087199Z","shell.execute_reply":"2024-06-07T09:16:13.023998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check dataset throughput\n\nIt's always a good idea to check that your data pipeline is efficient. You don't want to be CPU-bound at training time!","metadata":{}},{"cell_type":"code","source":"t0 = time.time()\nn = 200\nfor _ in train_ds.take(n):\n    pass\nprint(f\"Time per batch: {(time.time() - t0) / n:.4f}s\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:13.027074Z","iopub.execute_input":"2024-06-07T09:16:13.027485Z","iopub.status.idle":"2024-06-07T09:16:17.278913Z","shell.execute_reply.started":"2024-06-07T09:16:13.027443Z","shell.execute_reply":"2024-06-07T09:16:17.277772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create validation dataset","metadata":{}},{"cell_type":"code","source":"val_locations_stride = BUFFER\nval_locations = []\nfor x in range(val_location[0], val_location[0] + val_zone_size[0], val_locations_stride):\n    for y in range(val_location[1], val_location[1] + val_zone_size[1], val_locations_stride):\n        val_locations.append((x, y))\n\nval_locations_ds = tf.data.Dataset.from_tensor_slices(val_locations).filter(is_in_mask_train)\nval_ds = val_locations_ds.map(extract_subvolume_and_label, num_parallel_calls=tf.data.AUTOTUNE)\nval_ds = val_ds.prefetch(tf.data.AUTOTUNE).batch(BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:17.280530Z","iopub.execute_input":"2024-06-07T09:16:17.280882Z","iopub.status.idle":"2024-06-07T09:16:17.355140Z","shell.execute_reply.started":"2024-06-07T09:16:17.280852Z","shell.execute_reply":"2024-06-07T09:16:17.354042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize validation dataset patches\n\nNote that they are partially overlapping, since the stride is half the patch size.","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots()\nax.imshow(labels)\n\nfor x, y in val_locations_ds:\n    patch = patches.Rectangle([y - BUFFER, x - BUFFER], 2 * BUFFER, 2 * BUFFER, linewidth=2, edgecolor='g', facecolor='none')\n    ax.add_patch(patch)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:17.356579Z","iopub.execute_input":"2024-06-07T09:16:17.356961Z","iopub.status.idle":"2024-06-07T09:16:24.700498Z","shell.execute_reply.started":"2024-06-07T09:16:17.356920Z","shell.execute_reply":"2024-06-07T09:16:24.699465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Compute a trivial baseline\n\nThis is the highest validation score you can reach without looking at the inputs.\nThe model can be considered to have statistical power only if it can beat this baseline.","metadata":{}},{"cell_type":"code","source":"def trivial_baseline(dataset):\n    total = 0\n    matches = 0.\n    for _, batch_label in tqdm(dataset):\n        matches += tf.reduce_sum(tf.cast(batch_label, \"float32\"))\n        total += tf.reduce_prod(tf.shape(batch_label))\n    return 1. - matches / tf.cast(total, \"float32\")\n\nscore = trivial_baseline(val_ds).numpy()\nprint(f\"Best validation score achievable trivially: {score * 100:.2f}% accuracy\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:24.701923Z","iopub.execute_input":"2024-06-07T09:16:24.702612Z","iopub.status.idle":"2024-06-07T09:16:27.769074Z","shell.execute_reply.started":"2024-06-07T09:16:24.702560Z","shell.execute_reply":"2024-06-07T09:16:27.768014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augment the training data","metadata":{}},{"cell_type":"code","source":"augmenter = keras.Sequential([\n    layers.RandomContrast(0.2),\n])\n\ndef augment_train_data(data, label):\n    data = augmenter(data)\n    return data, label\n\naugmented_train_ds = train_ds.map(augment_train_data, num_parallel_calls=tf.data.AUTOTUNE).prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:27.770504Z","iopub.execute_input":"2024-06-07T09:16:27.770879Z","iopub.status.idle":"2024-06-07T09:16:28.573682Z","shell.execute_reply.started":"2024-06-07T09:16:27.770839Z","shell.execute_reply":"2024-06-07T09:16:28.572583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train a Keras model\n\nThis model is a U-Net taken from [this segmentation tutorial](https://keras.io/examples/vision/oxford_pets_image_segmentation/).\n\n`model.fit()` goes brrrrr\n\nConceptually it looks like this (animation from [this tutorial](https://www.kaggle.com/code/jpposma/vesuvius-challenge-ink-detection-tutorial)):\n\n![animation](https://user-images.githubusercontent.com/22727759/224853385-ed190d89-f466-469c-82a9-499881759d57.gif)","metadata":{}},{"cell_type":"code","source":"def get_model(input_shape):\n    inputs = keras.Input(input_shape)\n    \n    x = inputs\n    \n    ### [First half of the network: downsampling inputs] ###\n\n    # Entry block\n    x = layers.Conv2D(64, 3, strides=2, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    previous_block_activation = x  # Set aside residual\n\n    # Blocks 1, 2, 3 are identical apart from the feature depth.\n    for filters in [128, 256]:\n        x = layers.Activation(\"relu\")(x)\n        x = layers.SeparableConv2D(filters, 3, padding=\"same\")(x)\n        x = layers.BatchNormalization()(x)\n\n        x = layers.Activation(\"relu\")(x)\n        x = layers.SeparableConv2D(filters, 3, padding=\"same\")(x)\n        x = layers.BatchNormalization()(x)\n\n        x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n        residual = layers.Conv2D(filters, 1, strides=2, padding=\"same\")(\n            previous_block_activation\n        )\n        x = layers.add([x, residual])  # Add back residual\n        previous_block_activation = x  # Set aside next residual\n\n    ### [Second half of the network: upsampling inputs] ###\n\n    for filters in [256, 128, 64]:\n        x = layers.Activation(\"relu\")(x)\n        x = layers.Conv2DTranspose(filters, 3, padding=\"same\")(x)\n        x = layers.BatchNormalization()(x)\n\n        x = layers.Activation(\"relu\")(x)\n        x = layers.Conv2DTranspose(filters, 3, padding=\"same\")(x)\n        x = layers.BatchNormalization()(x)\n\n        x = layers.UpSampling2D(2)(x)\n\n        # Project residual\n        residual = layers.UpSampling2D(2)(previous_block_activation)\n        residual = layers.Conv2D(filters, 1, padding=\"same\")(residual)\n        x = layers.add([x, residual])  # Add back residual\n        previous_block_activation = x  # Set aside next residual\n\n    # Add a per-pixel classification layer\n    outputs = layers.Conv2D(1, 3, activation=\"sigmoid\", padding=\"same\")(x)\n\n    # Define the model\n    model = keras.Model(inputs, outputs)\n    return model\n\ndef gating_signal(input, out_size):\n    x = layers.Conv2D(out_size, (1, 1), padding='same')(input)\n    x = layers.BatchNormalization()(x)\n    x = layers.Activation('relu')(x)\n    return x\n\n\n\n\ndef repeat_elem(tensor, rep):\n    # lambda function to repeat Repeats the elements of a tensor along an axis\n    # by a factor of rep.\n    # If tensor has shape (None, 256, 256, 3), this will return a tensor of shape \n    # (None, 256, 256, 6), if specified axis=3 and rep=2.\n    return layers.Lambda(lambda x: K.repeat_elements(x, rep, axis=3))(tensor)\n\n# def repeat_elem(tensor, rep):\n#     # lambda function to repeat Repeats the elements of a tensor along an axis\n#     #by a factor of rep.\n#     # If tensor has shape (None, 256,256,3), lambda will return a tensor of shape \n#     #(None, 256,256,6), if specified axis=3 and rep=2.\n\n#      return layers.Lambda(lambda x, repnum: K.repeat_elements(x, repnum, axis=3),\n#                           arguments={'repnum': rep})(tensor) \n    \ndef attention_block(x, gating, inter_shape):\n    shape_x = K.int_shape(x)\n    shape_g = K.int_shape(gating)\n\n# Getting the x signal to the same shape as the gating signal\n    theta_x = layers.Conv2D(inter_shape, (2, 2), strides=(2, 2), padding='same')(x)  # 16\n    shape_theta_x = K.int_shape(theta_x)\n\n# Getting the gating signal to the same number of filters as the inter_shape\n    phi_g = layers.Conv2D(inter_shape, (1, 1), padding='same')(gating)\n    upsample_g = layers.Conv2DTranspose(inter_shape, (3, 3),\n                                 strides=(shape_theta_x[1] // shape_g[1], shape_theta_x[2] // shape_g[2]),\n                                 padding='same')(phi_g)  # 16\n\n    concat_xg = layers.add([upsample_g, theta_x])\n    act_xg = layers.Activation('relu')(concat_xg)\n    psi = layers.Conv2D(1, (1, 1), padding='same')(act_xg)\n    sigmoid_xg = layers.Activation('sigmoid')(psi)\n    shape_sigmoid = K.int_shape(sigmoid_xg)\n    upsample_psi = layers.UpSampling2D(size=(shape_x[1] // shape_sigmoid[1], shape_x[2] // shape_sigmoid[2]))(sigmoid_xg)  # 32\n\n    upsample_psi = repeat_elem(upsample_psi, shape_x[3])\n\n    y = layers.multiply([upsample_psi, x])\n\n    result = layers.Conv2D(shape_x[3], (1, 1), padding='same')(y)\n    result_bn = layers.BatchNormalization()(result)\n    return result_bn\n\n\n# if USE_MIXED_PRECISION:\n#     keras.mixed_precision.set_global_policy('mixed_float16')\n    \n\ndef attention_U_net(input_shape, NUM_CLASSES=1):\n   \n    FILTER_SIZE = 3 # size of the convolutional filter\n    UP_SAMP_SIZE = 2 # size of upsampling filters\n    \n    inputs = layers.Input(input_shape)\n    \n    x = inputs\n    \n    inputs = keras.Input(input_shape)\n    \n    x = inputs\n    \n    ### [First half of the network: downsampling inputs] ###\n\n    # Entry block\n    x = layers.Conv2D(64, 3, strides=2, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n    conv_64 = x\n\n    previous_block_activation = x  # Set aside residual\n\n    # Blocks 1, 2, 3 are identical apart from the feature depth.\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n#     residual = layers.Conv2D(128, 1, strides=2, padding=\"same\")(\n#             previous_block_activation\n#         )\n#     x = layers.add([x, residual])  # Add back residual\n#     previous_block_activation = x  # Set aside next residual\n    \n    conv_128 = x\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n#     residual = layers.Conv2D(256, 1, strides=2, padding=\"same\")(\n#             previous_block_activation\n#         )\n#     x = layers.add([x, residual])  # Add back residual\n#     previous_block_activation = x  # Set aside next residual\n    \n    conv_256 = x\n#     up\n    gat_128 =  gating_signal(conv_256,128)\n    at_128 = attention_block(conv_128,gat_128,128)\n    up_256 = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    x = layers.concatenate([up_256, at_128], axis=3)\n    \n    x = layers.Activation(\"relu\")(x)\n    up_256 = layers.Conv2DTranspose(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(up_256)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n#     residual = layers.UpSampling2D(2)(previous_block_activation)\n#     residual = layers.Conv2D(256, 1, padding=\"same\")(residual)\n#     x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x  # Set aside next residual\n    \n    gat_64 =  gating_signal(conv_128,64)\n    at_64 = attention_block(conv_64,gat_64,64)\n    up_128 = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    x = layers.concatenate([up_128, at_64], axis=3)\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n#     residual = layers.UpSampling2D(2)(previous_block_activation)\n#     residual = layers.Conv2D(128, 1, padding=\"same\")(residual)\n#     x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x\n    \n    \n    x = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    \n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(64, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(64, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        \n#     residual = layers.UpSampling2D(2)(previous_block_activation)\n#     residual = layers.Conv2D(64, 1, padding=\"same\")(residual)\n#     x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x\n    \n       \n    conv_final = layers.Conv2D(1, 3, activation=\"sigmoid\", padding=\"same\")(x)\n    \n\n    # Model integration\n    model = models.Model(inputs, conv_final)\n    return model\n\n\n    \ndef Res_attention_U_net(input_shape, NUM_CLASSES=1):\n   \n    FILTER_SIZE = 3 # size of the convolutional filter\n    UP_SAMP_SIZE = 2 # size of upsampling filters\n    \n    inputs = layers.Input(input_shape)\n    \n    x = inputs\n    \n    inputs = keras.Input(input_shape)\n    \n    x = inputs\n    \n    ### [First half of the network: downsampling inputs] ###\n\n    # Entry block\n    x = layers.Conv2D(64, 3, strides=2, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n    conv_64 = x\n\n    previous_block_activation = x  # Set aside residual\n\n    # Blocks 1, 2, 3 are identical apart from the feature depth.\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n    residual = layers.Conv2D(128, 1, strides=2, padding=\"same\")(\n            previous_block_activation\n        )\n    x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x  # Set aside next residual\n    \n    conv_128 = x\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.SeparableConv2D(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n    residual = layers.Conv2D(256, 1, strides=2, padding=\"same\")(\n            previous_block_activation\n        )\n    x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x  # Set aside next residual\n    \n    conv_256 = x\n#     up\n    gat_128 =  gating_signal(conv_256,128)\n    at_128 = attention_block(conv_128,gat_128,128)\n    up_256 = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    x = layers.concatenate([up_256, at_128], axis=3)\n    \n    x = layers.Activation(\"relu\")(x)\n    up_256 = layers.Conv2DTranspose(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(up_256)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(256, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n    residual = layers.UpSampling2D(2)(previous_block_activation)\n    residual = layers.Conv2D(256, 1, padding=\"same\")(residual)\n    x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x  # Set aside next residual\n    \n    gat_64 =  gating_signal(conv_128,64)\n    at_64 = attention_block(conv_64,gat_64,64)\n    up_128 = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    x = layers.concatenate([up_128, at_64], axis=3)\n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(128, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        # Project residual\n    residual = layers.UpSampling2D(2)(previous_block_activation)\n    residual = layers.Conv2D(128, 1, padding=\"same\")(residual)\n    x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x\n    \n    \n    x = layers.UpSampling2D(size=(UP_SAMP_SIZE, UP_SAMP_SIZE), data_format=\"channels_last\")(x)\n    \n    \n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(64, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2DTranspose(64, 3, padding=\"same\")(x)\n    x = layers.BatchNormalization()(x)\n\n    x = layers.UpSampling2D(2)(x)\n    x = layers.MaxPooling2D(3, strides=2, padding=\"same\")(x)\n\n        \n    residual = layers.UpSampling2D(2)(previous_block_activation)\n    residual = layers.Conv2D(64, 1, padding=\"same\")(residual)\n    x = layers.add([x, residual])  # Add back residual\n    previous_block_activation = x\n    \n       \n    conv_final = layers.Conv2D(1, 3, activation=\"sigmoid\", padding=\"same\")(x)\n    \n\n    # Model integration\n    model = models.Model(inputs, conv_final)\n    return model\n    \n    \n    \n    \n\nmodel_U = get_model((BUFFER * 2, BUFFER * 2, Z_DIM))\nmodel_U.summary()\nmodel_U.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"accuracy\"], jit_compile=USE_JIT_COMPILE)\n\nmodel_U.fit(augmented_train_ds, validation_data=val_ds, epochs=10, steps_per_epoch=1500)\nmodel_U.save(\"model_U.keras\")\n\nmodel_A = attention_U_net((BUFFER * 2, BUFFER * 2, Z_DIM))\nmodel_A.summary()\nmodel_A.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"accuracy\"], jit_compile=USE_JIT_COMPILE)\n\nmodel_A.fit(augmented_train_ds, validation_data=val_ds, epochs=10, steps_per_epoch=1500)\nmodel_A.save(\"model.keras\")\n\nmodel_R = Res_attention_U_net((BUFFER * 2, BUFFER * 2, Z_DIM))\nmodel_R.summary()\nmodel_R.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"accuracy\"], jit_compile=USE_JIT_COMPILE)\n\nmodel_R.fit(augmented_train_ds, validation_data=val_ds, epochs=10, steps_per_epoch=1500)\nmodel_R.save(\"model_r.keras\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T09:16:28.575747Z","iopub.execute_input":"2024-06-07T09:16:28.576166Z","iopub.status.idle":"2024-06-07T10:00:39.519199Z","shell.execute_reply.started":"2024-06-07T09:16:28.576124Z","shell.execute_reply":"2024-06-07T10:00:39.518087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Clear up memory","metadata":{}},{"cell_type":"code","source":"# del volume\n# del mask\n# del labels\n# del train_ds\n# del val_ds\n\n# Manually trigger garbage collection\nkeras.backend.clear_session()\nimport gc\ngc.collect()\n\nmodel = keras.models.load_model(\"model_R.keras\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:39.521012Z","iopub.execute_input":"2024-06-07T10:00:39.521431Z","iopub.status.idle":"2024-06-07T10:00:39.995722Z","shell.execute_reply.started":"2024-06-07T10:00:39.521390Z","shell.execute_reply":"2024-06-07T10:00:39.994082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Compute predictions on test data","metadata":{}},{"cell_type":"code","source":"def compute_predictions_map(split, index):\n    print(f\"Load data for {split}/{index}\")\n\n    test_volume = load_volume(split=split, index=index)\n    test_mask = load_mask(split=split, index=index)\n\n    test_locations = []\n    stride = BUFFER // 2\n    for x in range(BUFFER, test_volume.shape[0] - BUFFER, stride):\n        for y in range(BUFFER, test_volume.shape[1] - BUFFER, stride):\n            test_locations.append((x, y))\n\n    print(f\"{len(test_locations)} test locations (before filtering by mask)\")\n\n    sample_random_location_test = lambda x: sample_random_location(test_mask.shape)\n    is_in_mask_test = lambda x: is_in_masked_zone(x, test_mask)\n    extract_subvolume_test = lambda x: extract_subvolume(x, test_volume)\n\n    test_locations_ds = tf.data.Dataset.from_tensor_slices(test_locations).filter(is_in_mask_test)\n    test_ds = test_locations_ds.map(extract_subvolume_test, num_parallel_calls=tf.data.AUTOTUNE)\n\n    predictions_map = np.zeros(test_volume.shape[:2] + (1,), dtype=\"float16\")\n    predictions_map_counts = np.zeros(test_volume.shape[:2] + (1,), dtype=\"int8\")\n\n    print(f\"Compute predictions\")\n\n    for loc_batch, patch_batch in tqdm(zip(test_locations_ds.batch(BATCH_SIZE), test_ds.batch(BATCH_SIZE))):\n        predictions = model.predict_on_batch(patch_batch)\n        for (x, y), pred in zip(loc_batch, predictions):\n            predictions_map[x - BUFFER : x + BUFFER, y - BUFFER : y + BUFFER, :] += pred\n            predictions_map_counts[x - BUFFER : x + BUFFER, y - BUFFER : y + BUFFER, :] += 1  \n    predictions_map /= (predictions_map_counts + 1e-7)\n    return predictions_map","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:39.996896Z","iopub.status.idle":"2024-06-07T10:00:39.997308Z","shell.execute_reply.started":"2024-06-07T10:00:39.997114Z","shell.execute_reply":"2024-06-07T10:00:39.997135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions_map_a = compute_predictions_map(split=\"test\", index=\"a\")\npredictions_map_b = compute_predictions_map(split=\"test\", index=\"b\")","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:39.998900Z","iopub.status.idle":"2024-06-07T10:00:39.999410Z","shell.execute_reply.started":"2024-06-07T10:00:39.999152Z","shell.execute_reply":"2024-06-07T10:00:39.999177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resize prediction maps to their original size (for submission)","metadata":{}},{"cell_type":"code","source":"from skimage.transform import resize as resize_ski\n\noriginal_size_a = Image.open(DATA_DIR + \"/test/a/mask.png\").size\npredictions_map_a = resize_ski(predictions_map_a, (original_size_a[1], original_size_a[0])).squeeze()\n\noriginal_size_b = Image.open(DATA_DIR + \"/test/b/mask.png\").size\npredictions_map_b = resize_ski(predictions_map_b, (original_size_b[1], original_size_b[0])).squeeze()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:40.000901Z","iopub.status.idle":"2024-06-07T10:00:40.001674Z","shell.execute_reply.started":"2024-06-07T10:00:40.001393Z","shell.execute_reply":"2024-06-07T10:00:40.001420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Generate submission file","metadata":{}},{"cell_type":"code","source":"def rle(predictions_map, threshold):\n    flat_img = predictions_map.flatten()\n    flat_img = np.where(flat_img > threshold, 1, 0).astype(np.uint8)\n\n    starts = np.array((flat_img[:-1] == 0) & (flat_img[1:] == 1))\n    ends = np.array((flat_img[:-1] == 1) & (flat_img[1:] == 0))\n    starts_ix = np.where(starts)[0] + 2\n    ends_ix = np.where(ends)[0] + 2\n    lengths = ends_ix - starts_ix\n    return \" \".join(map(str, sum(zip(starts_ix, lengths), ())))","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:40.004478Z","iopub.status.idle":"2024-06-07T10:00:40.005360Z","shell.execute_reply.started":"2024-06-07T10:00:40.005110Z","shell.execute_reply":"2024-06-07T10:00:40.005136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"threshold = 0.5\n\nrle_a = rle(predictions_map_a, threshold=threshold)\nrle_b = rle(predictions_map_b, threshold=threshold)\nprint(\"Id,Predicted\\na,\" + rle_a + \"\\nb,\" + rle_b, file=open('submission.csv', 'w'))","metadata":{"execution":{"iopub.status.busy":"2024-06-07T10:00:40.006751Z","iopub.status.idle":"2024-06-07T10:00:40.007577Z","shell.execute_reply.started":"2024-06-07T10:00:40.007302Z","shell.execute_reply":"2024-06-07T10:00:40.007330Z"},"trusted":true},"execution_count":null,"outputs":[]}]}