{"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":"markdown","source":"# CIS 3115 Project 10: Vesuvius Ink Detection\n\nIn this notebook we will apply CNNs to the [Vesuvius Ink Detection challenge on Kaggle](https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/overview). We will look at two different CNN configurations, the Encoder-Decoder model and the U-Net model.\n\nMuch of the code in the project is taken from [Keras starter kit: UNet + train on full dataset](https://www.kaggle.com/code/fchollet/keras-starter-kit-unet-train-on-full-dataset) notebook the by fchollet. \n\n## Writeup\n\nMake a copy of the [ Project 10 Writeup](https://docs.google.com/document/d/14jy1Fyj3QMd9qdp3alYVXdta-SZv07u3LXuyUaowcoY/edit?usp=sharing) answer the questions in the writeup. \n\nYou will submit a link to the Writeup in Brightspace.\n\n## Video Walkthrough\n\nhttps://youtu.be/A_GdjQVETLA\n\n## Overview\n\nThis is part of a larger challenge. The goal of the Vesuvius Challenge to resurrect an ancient library from the ashes of a volcano. In this larger competition researchers are tasked with detecting ink from 3D X-ray scans and reading the contents. Thousands of scrolls were part of a library located in a Roman villa in Herculaneum, a town next to Pompeii. This villa was buried by the Vesuvius eruption nearly 2000 years ago. Due to the heat of the volcano, the scrolls were carbonized, and are now impossible to open without breaking them. These scrolls were discovered a few hundred years ago and have been waiting to be read using modern techniques. \n\nThe [Kaggle Ink Detection challenge](https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/overview) is about the sub-problem of detecting ink from 3d x-ray scans of fragments of papyrus which became detached from some of the excavated scrolls. This subcontest is run on Kaggle since it's a more traditional data science / machine learning problem of building a model that can be verified against known ground truth data.\n\nThe ink used in the Herculaneum scrolls does not show up readily in X-ray scans. But we have found that machine learning models can detect it. Luckily, we have ground truth data. Since the discovery of the Herculaneum Papyri almost 300 years ago, people have tried opening them, often with disastrous results. Many scrolls were destroyed in this process, but ink can be seen on some broken-off fragments, especially under infrared light.\n\nThe dataset contains 3d x-ray scans of four such fragments. Our goal is to try to detect the ink in from the 3d scans.\n\n","metadata":{}},{"cell_type":"markdown","source":"# Part 0 - Setup","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.models import Sequential, Model\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.layers import Input, Dense, Dropout, Activation, Lambda, Flatten, LSTM\nfrom tensorflow.keras.layers import Conv2D, Convolution2D, Conv2DTranspose\nfrom tensorflow.keras.layers import MaxPooling2D, AveragePooling2D, GlobalAveragePooling2D, concatenate\nfrom tensorflow.keras.optimizers import Adam, RMSprop\nfrom tensorflow.keras.utils import to_categorical\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 = 16   # Number of slices in the z direction. Max value is 64 - Z_START\nZ_START = 25  # Offset of slices in the z direction\nSHARED_HEIGHT = 4096  # Height to resize all papyrii\n\n# Model config\nBATCH_SIZE = 32\nUSE_MIXED_PRECISION = False\nUSE_JIT_COMPILE = False\n\nif USE_MIXED_PRECISION:\n    keras.mixed_precision.set_global_policy('mixed_float16')","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:21:04.617849Z","iopub.execute_input":"2023-06-27T02:21:04.618601Z","iopub.status.idle":"2023-06-27T02:21:12.205664Z","shell.execute_reply.started":"2023-06-27T02:21:04.618559Z","shell.execute_reply":"2023-06-27T02:21:12.204154Z"},"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":"2023-06-27T02:21:17.084010Z","iopub.execute_input":"2023-06-27T02:21:17.085189Z","iopub.status.idle":"2023-06-27T02:21:20.361377Z","shell.execute_reply.started":"2023-06-27T02:21:17.085142Z","shell.execute_reply":"2023-06-27T02:21:20.360332Z"},"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":"2023-06-27T02:21:36.122968Z","iopub.execute_input":"2023-06-27T02:21:36.123585Z","iopub.status.idle":"2023-06-27T02:21:40.591887Z","shell.execute_reply.started":"2023-06-27T02:21:36.123545Z","shell.execute_reply":"2023-06-27T02:21:40.590820Z"},"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":"2023-06-27T02:21:51.298670Z","iopub.execute_input":"2023-06-27T02:21:51.299183Z","iopub.status.idle":"2023-06-27T02:21:54.068046Z","shell.execute_reply.started":"2023-06-27T02:21:51.299130Z","shell.execute_reply":"2023-06-27T02:21:54.066918Z"},"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":"2023-06-27T02:22:01.056516Z","iopub.execute_input":"2023-06-27T02:22:01.057331Z","iopub.status.idle":"2023-06-27T02:22:02.529863Z","shell.execute_reply.started":"2023-06-27T02:22:01.057287Z","shell.execute_reply":"2023-06-27T02:22:02.528805Z"},"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":"2023-06-27T02:22:07.829379Z","iopub.execute_input":"2023-06-27T02:22:07.830529Z","iopub.status.idle":"2023-06-27T02:22:07.838346Z","shell.execute_reply.started":"2023-06-27T02:22:07.830470Z","shell.execute_reply":"2023-06-27T02:22:07.837243Z"},"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":"2023-06-27T02:22:10.048658Z","iopub.execute_input":"2023-06-27T02:22:10.049405Z","iopub.status.idle":"2023-06-27T02:25:35.899958Z","shell.execute_reply.started":"2023-06-27T02:22:10.049367Z","shell.execute_reply":"2023-06-27T02:25:35.898766Z"},"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":"2023-06-27T02:27:04.301058Z","iopub.execute_input":"2023-06-27T02:27:04.301656Z","iopub.status.idle":"2023-06-27T02:27:04.315435Z","shell.execute_reply.started":"2023-06-27T02:27:04.301607Z","shell.execute_reply":"2023-06-27T02:27:04.314186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Part 1 - Visualization\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":"2023-06-27T02:27:09.634039Z","iopub.execute_input":"2023-06-27T02:27:09.634427Z","iopub.status.idle":"2023-06-27T02:27:14.332740Z","shell.execute_reply.started":"2023-06-27T02:27:09.634392Z","shell.execute_reply":"2023-06-27T02:27:14.331635Z"},"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 = (1300, 1000)\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":"2023-06-27T02:29:06.014176Z","iopub.execute_input":"2023-06-27T02:29:06.015013Z","iopub.status.idle":"2023-06-27T02:29:07.501330Z","shell.execute_reply.started":"2023-06-27T02:29:06.014971Z","shell.execute_reply":"2023-06-27T02:29:07.500159Z"},"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":"2023-06-27T02:29:10.904123Z","iopub.execute_input":"2023-06-27T02:29:10.904523Z","iopub.status.idle":"2023-06-27T02:29:11.454008Z","shell.execute_reply.started":"2023-06-27T02:29:10.904487Z","shell.execute_reply":"2023-06-27T02:29:11.452821Z"},"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(200):\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":"2023-06-27T02:29:15.170589Z","iopub.execute_input":"2023-06-27T02:29:15.171010Z","iopub.status.idle":"2023-06-27T02:29:17.792196Z","shell.execute_reply.started":"2023-06-27T02:29:15.170966Z","shell.execute_reply":"2023-06-27T02:29:17.791001Z"},"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":"2023-06-27T02:29:21.690126Z","iopub.execute_input":"2023-06-27T02:29:21.690891Z","iopub.status.idle":"2023-06-27T02:29:23.850615Z","shell.execute_reply.started":"2023-06-27T02:29:21.690818Z","shell.execute_reply":"2023-06-27T02:29:23.849447Z"},"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":"2023-06-27T02:29:29.361477Z","iopub.execute_input":"2023-06-27T02:29:29.362098Z","iopub.status.idle":"2023-06-27T02:29:32.670605Z","shell.execute_reply.started":"2023-06-27T02:29:29.362057Z","shell.execute_reply":"2023-06-27T02:29:32.669496Z"},"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":"2023-06-27T02:29:37.466765Z","iopub.execute_input":"2023-06-27T02:29:37.467503Z","iopub.status.idle":"2023-06-27T02:29:42.033672Z","shell.execute_reply.started":"2023-06-27T02:29:37.467462Z","shell.execute_reply":"2023-06-27T02:29:42.032559Z"},"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":"2023-06-27T02:29:46.101833Z","iopub.execute_input":"2023-06-27T02:29:46.102591Z","iopub.status.idle":"2023-06-27T02:29:46.179585Z","shell.execute_reply.started":"2023-06-27T02:29:46.102550Z","shell.execute_reply":"2023-06-27T02:29:46.178492Z"},"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":"2023-06-27T02:29:51.876733Z","iopub.execute_input":"2023-06-27T02:29:51.877257Z","iopub.status.idle":"2023-06-27T02:29:56.712172Z","shell.execute_reply.started":"2023-06-27T02:29:51.877217Z","shell.execute_reply":"2023-06-27T02:29:56.710942Z"},"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":"2023-06-27T02:30:00.548998Z","iopub.execute_input":"2023-06-27T02:30:00.549389Z","iopub.status.idle":"2023-06-27T02:30:03.667197Z","shell.execute_reply.started":"2023-06-27T02:30:00.549354Z","shell.execute_reply":"2023-06-27T02:30:03.665900Z"},"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":"2023-06-27T02:31:25.909197Z","iopub.execute_input":"2023-06-27T02:31:25.914531Z","iopub.status.idle":"2023-06-27T02:31:27.077054Z","shell.execute_reply.started":"2023-06-27T02:31:25.914468Z","shell.execute_reply":"2023-06-27T02:31:27.075818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Part 3 - Define and train models","metadata":{}},{"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":"markdown","source":"## Encoder-Decoder Model\n\nThe following model is a basic encoder-decoder model.\n\n<img src=\"https://cds.ismrm.org/protected/17MProceedings/PDFfiles/images/8249/ISMRM2017-008249_Fig1.png\">\n*Image from https://cds.ismrm.org/protected/17MProceedings/PDFfiles/5662.html*\n\n","metadata":{}},{"cell_type":"code","source":"def encode_decode_model(base_filters, window_image_size, input_channels):\n    input_layer = Input((window_image_size, window_image_size, input_channels))\n        \n    # standard size 64 -> 32      \n    conv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(input_layer)\n    conv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(conv1)\n    pool1 = MaxPooling2D((2, 2))(conv1)\n    \n    # standard size 32 -> 16      \n    conv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(pool1)\n    conv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(conv2)\n    pool2 = MaxPooling2D((2, 2))(conv2)\n\n    # standard size 16 -> 8      \n    conv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(pool2)\n    conv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(conv3)\n    pool3 = MaxPooling2D((2, 2))(conv3)\n    \n    # Middle\n    convm = Conv2D(base_filters * 16, (3, 3), activation=\"relu\", padding=\"same\")(pool3)\n    convm = Conv2D(base_filters * 16, (3, 3), activation=\"relu\", padding=\"same\")(convm)\n    \n    # standard size 8 -> 16        \n    uconv3 = Conv2DTranspose(base_filters * 4, (3, 3), strides=(2, 2), padding=\"same\")(convm)\n    uconv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(uconv3)\n    uconv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(uconv3)\n    \n    # standard size 16 -> 32        \n    uconv2 = Conv2DTranspose(base_filters * 2, (3, 3), strides=(2, 2), padding=\"same\")(uconv3)\n    uconv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(uconv2)\n    uconv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(uconv2)\n\n    # standard size 32 -> 64   \n    uconv1 = Conv2DTranspose(base_filters * 1, (3, 3), strides=(2, 2), padding=\"same\")(uconv2)\n    uconv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(uconv1)\n    uconv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(uconv1)\n\n    output_layer = Conv2D(1, (1,1), padding=\"same\", activation=\"sigmoid\")(uconv1)\n    \n    return Model(input_layer, output_layer)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:31:33.442836Z","iopub.execute_input":"2023-06-27T02:31:33.443605Z","iopub.status.idle":"2023-06-27T02:31:33.463610Z","shell.execute_reply.started":"2023-06-27T02:31:33.443563Z","shell.execute_reply":"2023-06-27T02:31:33.462505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## U-Net Model\n\nThe following model is [U-Net](https://en.wikipedia.org/wiki/U-Net) which is a modified version of the encoder-decoder model above. U-Net adds connections from each encoding layer to the corresponding decoding layer. These are sometimes called residual connections\n\n<img src=\"https://upload.wikimedia.org/wikipedia/commons/2/2b/Example_architecture_of_U-Net_for_producing_k_256-by-256_image_masks_for_a_256-by-256_RGB_image.png\">\n*Image from https://en.wikipedia.org/wiki/U-Net*","metadata":{}},{"cell_type":"code","source":"def unet_model(base_filters, window_image_size, input_channels):\n    input_layer = Input((window_image_size, window_image_size, input_channels))\n        \n    # standard size 64 -> 32   \n    conv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(input_layer)\n    conv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(conv1)\n    pool1 = MaxPooling2D((2, 2))(conv1)\n    \n    # standard size 32 -> 16     \n    conv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(pool1)\n    conv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(conv2)\n    pool2 = MaxPooling2D((2, 2))(conv2)\n\n    # standard size 16 -> 8     \n    conv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(pool2)\n    conv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(conv3)\n    pool3 = MaxPooling2D((2, 2))(conv3)\n    \n    # Middle\n    convm = Conv2D(base_filters * 16, (3, 3), activation=\"relu\", padding=\"same\")(pool3)\n    convm = Conv2D(base_filters * 16, (3, 3), activation=\"relu\", padding=\"same\")(convm)\n    \n    # standard size 8 -> 16       \n    deconv3 = Conv2DTranspose(base_filters * 4, (3, 3), strides=(2, 2), padding=\"same\")(convm)\n    uconv3 = concatenate([deconv3, conv3])\n    uconv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(uconv3)\n    uconv3 = Conv2D(base_filters * 4, (3, 3), activation=\"relu\", padding=\"same\")(uconv3)\n    \n    # standard size 16 -> 32       \n    deconv2 = Conv2DTranspose(base_filters * 2, (3, 3), strides=(2, 2), padding=\"same\")(uconv3)\n    uconv2 = concatenate([deconv2, conv2])\n    uconv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(uconv2)\n    uconv2 = Conv2D(base_filters * 2, (3, 3), activation=\"relu\", padding=\"same\")(uconv2)\n\n    # standard size 32 -> 64  \n    deconv1 = Conv2DTranspose(base_filters * 1, (3, 3), strides=(2, 2), padding=\"same\")(uconv2)\n    uconv1 = concatenate([deconv1, conv1])\n    uconv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(uconv1)\n    uconv1 = Conv2D(base_filters * 1, (3, 3), activation=\"relu\", padding=\"same\")(uconv1)\n\n    output_layer = Conv2D(1, (1,1), padding=\"same\", activation=\"sigmoid\")(uconv1)\n    \n    return Model(input_layer, output_layer)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:31:39.010281Z","iopub.execute_input":"2023-06-27T02:31:39.010682Z","iopub.status.idle":"2023-06-27T02:31:39.030959Z","shell.execute_reply.started":"2023-06-27T02:31:39.010645Z","shell.execute_reply":"2023-06-27T02:31:39.029412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compete_unet_model(window_image_size, input_channels):\n    input_layer = Input((window_image_size, window_image_size, input_channels))\n    \n    x = input_layer\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    output_layer = layers.Conv2D(1, 3, activation=\"sigmoid\", padding=\"same\")(x)\n\n    # Define the model\n    model = keras.Model(input_layer, output_layer)\n    return model\n\n","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:31:42.836905Z","iopub.execute_input":"2023-06-27T02:31:42.837928Z","iopub.status.idle":"2023-06-27T02:31:42.854309Z","shell.execute_reply.started":"2023-06-27T02:31:42.837832Z","shell.execute_reply":"2023-06-27T02:31:42.853180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"   \nWINDOW_SIZE = BUFFER * 2             # Renaming some parameters, WINDOW_SIZE is the num of pixels in the moving window sampling the image\nWINDOW_DEPTH = Z_DIM                 # Renaming some parameters, WINDOW_DEPTH is the number of layers sampled in the origina image\n\n# =========== Select the model to use here ===========\n#model = encode_decode_model(16, WINDOW_SIZE, WINDOW_DEPTH)            # Load the basic encoder-decoder model\n#model = unet_model(16, WINDOW_SIZE, WINDOW_DEPTH)                     # load the U-Net model\nmodel = compete_unet_model(WINDOW_SIZE, WINDOW_DEPTH )                 # load the complete U-Net model with batch normalizations\n    \nmodel.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"accuracy\"], jit_compile=USE_JIT_COMPILE)\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:47:10.541728Z","iopub.execute_input":"2023-06-27T02:47:10.542842Z","iopub.status.idle":"2023-06-27T02:47:11.119778Z","shell.execute_reply.started":"2023-06-27T02:47:10.542797Z","shell.execute_reply":"2023-06-27T02:47:11.118936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.callbacks import ReduceLROnPlateau, EarlyStopping, ModelCheckpoint\n\nlearning_rate_reduction = ReduceLROnPlateau(monitor='loss', \n                                            patience=5, \n                                            verbose=2, \n                                            factor=0.5,                                            \n                                            min_lr=0.000001)\n\nearly_stops = EarlyStopping(monitor='loss', \n                            min_delta=0, \n                            patience=6, \n                            verbose=2, \n                            mode='auto')\n\ncheckpointer = ModelCheckpoint(filepath = 'cis3115.{epoch:02d}-{accuracy:.6f}.hdf5',\n                               verbose=2,\n                               save_best_only=True, \n                               save_weights_only = True)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:47:31.504612Z","iopub.execute_input":"2023-06-27T02:47:31.505046Z","iopub.status.idle":"2023-06-27T02:47:31.512759Z","shell.execute_reply.started":"2023-06-27T02:47:31.505007Z","shell.execute_reply":"2023-06-27T02:47:31.511675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training will take about 15 minutes\n\n# Train the model with the images in the folders\nhistory = model.fit(\n        augmented_train_ds,\n        steps_per_epoch=1000,                   # Number of images to process per epoch \n        epochs=20,                              # Number of epochs\n        callbacks=[learning_rate_reduction, early_stops],\n        validation_data=val_ds,\n        validation_steps=20                     # Number of images from validation set to test\n        )\n\nprint (\"Final training accuracy = \",history.history['accuracy'][-1])\nprint (\"Final testing accuracy = \",history.history['val_accuracy'][-1])\n\nmodel.save(\"model.keras\")","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:47:39.809763Z","iopub.execute_input":"2023-06-27T02:47:39.810502Z","iopub.status.idle":"2023-06-27T03:03:09.007779Z","shell.execute_reply.started":"2023-06-27T02:47:39.810460Z","shell.execute_reply":"2023-06-27T03:03:09.006198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We will display the loss and the accuracy of the model for each epoch\n# NOTE: this is a little fancier display than is shown in the textbook\ndef display_training_curves(training, validation, title, subplot):\n    if subplot%10==1: # set up the subplots on the first call\n        plt.subplots(figsize=(10,10), facecolor='#F0F0F0')\n        plt.tight_layout()\n    ax = plt.subplot(subplot)\n    ax.set_facecolor('#F8F8F8')\n    ax.plot(training)\n    ax.plot(validation)\n    ax.set_title('model '+ title)\n    ax.set_ylabel(title)\n    #ax.set_ylim(0.28,1.05)\n    ax.set_xlabel('epoch')\n    ax.legend(['train', 'valid.'])\n    \ndisplay_training_curves(history.history['loss'], history.history['val_loss'], 'loss', 211)\ndisplay_training_curves(history.history['accuracy'], history.history['val_accuracy'], 'accuracy', 212)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:14:11.839802Z","iopub.status.idle":"2023-06-27T02:14:11.840646Z","shell.execute_reply.started":"2023-06-27T02:14:11.840378Z","shell.execute_reply":"2023-06-27T02:14:11.840408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Clear up memory","metadata":{}},{"cell_type":"code","source":"del volume\ndel mask\ndel labels\ndel train_ds\ndel val_ds\n\n# Manually trigger garbage collection\nkeras.backend.clear_session()\nimport gc\ngc.collect()\n\nmodel = keras.models.load_model(\"model.keras\")","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:14:11.842129Z","iopub.status.idle":"2023-06-27T02:14:11.842973Z","shell.execute_reply.started":"2023-06-27T02:14:11.842701Z","shell.execute_reply":"2023-06-27T02:14:11.842728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Part 4 - Prediction","metadata":{}},{"cell_type":"markdown","source":"## Sample Prediction\n\nPick a sample location on the training images and generate a prediction from that.","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":"2023-06-27T01:59:54.110543Z","iopub.execute_input":"2023-06-27T01:59:54.111224Z","iopub.status.idle":"2023-06-27T01:59:54.124669Z","shell.execute_reply.started":"2023-06-27T01:59:54.111178Z","shell.execute_reply":"2023-06-27T01:59:54.123509Z"},"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":"2023-06-27T02:00:20.656435Z","iopub.execute_input":"2023-06-27T02:00:20.657417Z","iopub.status.idle":"2023-06-27T02:07:21.882978Z","shell.execute_reply.started":"2023-06-27T02:00:20.657359Z","shell.execute_reply":"2023-06-27T02:07:21.880243Z"},"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\noriginal_size_b = Image.open(DATA_DIR + \"/test/b/mask.png\").size\npredictions_map_a = resize_ski(predictions_map_a, original_size_a).squeeze()\npredictions_map_a = resize_ski(predictions_map_b, original_size_b).squeeze()","metadata":{"execution":{"iopub.status.busy":"2023-06-27T02:08:10.523248Z","iopub.execute_input":"2023-06-27T02:08:10.524183Z","iopub.status.idle":"2023-06-27T02:08:15.871681Z","shell.execute_reply.started":"2023-06-27T02:08:10.524143Z","shell.execute_reply":"2023-06-27T02:08:15.870533Z"},"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":"2023-06-27T02:08:19.796694Z","iopub.execute_input":"2023-06-27T02:08:19.797172Z","iopub.status.idle":"2023-06-27T02:08:19.810555Z","shell.execute_reply.started":"2023-06-27T02:08:19.797129Z","shell.execute_reply":"2023-06-27T02:08:19.809542Z"},"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":"2023-06-27T02:08:23.293822Z","iopub.execute_input":"2023-06-27T02:08:23.294558Z","iopub.status.idle":"2023-06-27T02:08:23.849305Z","shell.execute_reply.started":"2023-06-27T02:08:23.294519Z","shell.execute_reply":"2023-06-27T02:08:23.848159Z"},"trusted":true},"execution_count":null,"outputs":[]}]}