{"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":"# 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\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":"2023-03-18T23:10:40.298008Z","iopub.execute_input":"2023-03-18T23:10:40.298967Z","iopub.status.idle":"2023-03-18T23:10:49.514216Z","shell.execute_reply.started":"2023-03-18T23:10:40.298921Z","shell.execute_reply":"2023-03-18T23:10:49.513058Z"},"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-03-18T23:10:49.516568Z","iopub.execute_input":"2023-03-18T23:10:49.517657Z","iopub.status.idle":"2023-03-18T23:10:52.651975Z","shell.execute_reply.started":"2023-03-18T23:10:49.517615Z","shell.execute_reply":"2023-03-18T23:10:52.650919Z"},"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-03-18T23:10:52.653064Z","iopub.execute_input":"2023-03-18T23:10:52.653587Z","iopub.status.idle":"2023-03-18T23:10:57.915118Z","shell.execute_reply.started":"2023-03-18T23:10:52.653544Z","shell.execute_reply":"2023-03-18T23:10:57.914028Z"},"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-03-18T23:10:57.921638Z","iopub.execute_input":"2023-03-18T23:10:57.922052Z","iopub.status.idle":"2023-03-18T23:11:00.615559Z","shell.execute_reply.started":"2023-03-18T23:10:57.922012Z","shell.execute_reply":"2023-03-18T23:11:00.614397Z"},"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-03-18T23:11:00.619103Z","iopub.execute_input":"2023-03-18T23:11:00.619816Z","iopub.status.idle":"2023-03-18T23:11:02.849127Z","shell.execute_reply.started":"2023-03-18T23:11:00.619773Z","shell.execute_reply":"2023-03-18T23:11:02.847843Z"},"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-03-18T23:11:02.853840Z","iopub.execute_input":"2023-03-18T23:11:02.854397Z","iopub.status.idle":"2023-03-18T23:11:02.867203Z","shell.execute_reply.started":"2023-03-18T23:11:02.854351Z","shell.execute_reply":"2023-03-18T23:11:02.866021Z"},"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-03-18T23:11:02.872191Z","iopub.execute_input":"2023-03-18T23:11:02.874585Z","iopub.status.idle":"2023-03-18T23:14:11.220165Z","shell.execute_reply.started":"2023-03-18T23:11:02.874541Z","shell.execute_reply":"2023-03-18T23:14:11.218002Z"},"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-03-18T23:14:11.221817Z","iopub.execute_input":"2023-03-18T23:14:11.222178Z","iopub.status.idle":"2023-03-18T23:14:11.231554Z","shell.execute_reply.started":"2023-03-18T23:14:11.222141Z","shell.execute_reply":"2023-03-18T23:14:11.230344Z"},"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":"2023-03-18T23:14:11.233103Z","iopub.execute_input":"2023-03-18T23:14:11.233579Z","iopub.status.idle":"2023-03-18T23:14:16.475332Z","shell.execute_reply.started":"2023-03-18T23:14:11.233541Z","shell.execute_reply":"2023-03-18T23:14:16.474429Z"},"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-03-18T23:14:16.476814Z","iopub.execute_input":"2023-03-18T23:14:16.477954Z","iopub.status.idle":"2023-03-18T23:14:17.956818Z","shell.execute_reply.started":"2023-03-18T23:14:16.477911Z","shell.execute_reply":"2023-03-18T23:14:17.955681Z"},"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-03-18T23:14:17.958495Z","iopub.execute_input":"2023-03-18T23:14:17.958942Z","iopub.status.idle":"2023-03-18T23:14:18.455150Z","shell.execute_reply.started":"2023-03-18T23:14:17.958902Z","shell.execute_reply":"2023-03-18T23:14:18.454126Z"},"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-03-18T23:14:18.456727Z","iopub.execute_input":"2023-03-18T23:14:18.457083Z","iopub.status.idle":"2023-03-18T23:14:21.576187Z","shell.execute_reply.started":"2023-03-18T23:14:18.457045Z","shell.execute_reply":"2023-03-18T23:14:21.575222Z"},"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-03-18T23:14:21.577853Z","iopub.execute_input":"2023-03-18T23:14:21.579061Z","iopub.status.idle":"2023-03-18T23:14:23.602229Z","shell.execute_reply.started":"2023-03-18T23:14:21.579017Z","shell.execute_reply":"2023-03-18T23:14:23.601156Z"},"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-03-18T23:14:23.607154Z","iopub.execute_input":"2023-03-18T23:14:23.607564Z","iopub.status.idle":"2023-03-18T23:14:26.453345Z","shell.execute_reply.started":"2023-03-18T23:14:23.607531Z","shell.execute_reply":"2023-03-18T23:14:26.452345Z"},"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-03-18T23:14:26.455968Z","iopub.execute_input":"2023-03-18T23:14:26.459120Z","iopub.status.idle":"2023-03-18T23:14:31.187401Z","shell.execute_reply.started":"2023-03-18T23:14:26.459076Z","shell.execute_reply":"2023-03-18T23:14:31.186334Z"},"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-03-18T23:14:31.188813Z","iopub.execute_input":"2023-03-18T23:14:31.189186Z","iopub.status.idle":"2023-03-18T23:14:31.263309Z","shell.execute_reply.started":"2023-03-18T23:14:31.189143Z","shell.execute_reply":"2023-03-18T23:14:31.262192Z"},"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-03-18T23:14:31.264916Z","iopub.execute_input":"2023-03-18T23:14:31.265448Z","iopub.status.idle":"2023-03-18T23:14:35.752995Z","shell.execute_reply.started":"2023-03-18T23:14:31.265399Z","shell.execute_reply":"2023-03-18T23:14:35.751869Z"},"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-03-18T23:14:35.754646Z","iopub.execute_input":"2023-03-18T23:14:35.755036Z","iopub.status.idle":"2023-03-18T23:14:38.733818Z","shell.execute_reply.started":"2023-03-18T23:14:35.754998Z","shell.execute_reply":"2023-03-18T23:14:38.732850Z"},"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-03-18T23:14:38.737790Z","iopub.execute_input":"2023-03-18T23:14:38.740343Z","iopub.status.idle":"2023-03-18T23:14:39.032162Z","shell.execute_reply.started":"2023-03-18T23:14:38.740301Z","shell.execute_reply":"2023-03-18T23:14:39.031142Z"},"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\nif USE_MIXED_PRECISION:\n    keras.mixed_precision.set_global_policy('mixed_float16')\n\nmodel = get_model((BUFFER * 2, BUFFER * 2, Z_DIM))\nmodel.summary()\nmodel.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"accuracy\"], jit_compile=USE_JIT_COMPILE)\n\nmodel.fit(augmented_train_ds, validation_data=val_ds, epochs=20, steps_per_epoch=1000)\nmodel.save(\"model.keras\")","metadata":{"execution":{"iopub.status.busy":"2023-03-18T23:14:39.033748Z","iopub.execute_input":"2023-03-18T23:14:39.034124Z","iopub.status.idle":"2023-03-18T23:28:08.685364Z","shell.execute_reply.started":"2023-03-18T23:14:39.034086Z","shell.execute_reply":"2023-03-18T23:28:08.684324Z"},"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-03-18T23:28:08.687065Z","iopub.execute_input":"2023-03-18T23:28:08.687464Z","iopub.status.idle":"2023-03-18T23:28:09.709853Z","shell.execute_reply.started":"2023-03-18T23:28:08.687423Z","shell.execute_reply":"2023-03-18T23:28:09.708782Z"},"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":"2023-03-18T23:28:09.711493Z","iopub.execute_input":"2023-03-18T23:28:09.711867Z","iopub.status.idle":"2023-03-18T23:28:09.724755Z","shell.execute_reply.started":"2023-03-18T23:28:09.711827Z","shell.execute_reply":"2023-03-18T23:28:09.723509Z"},"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-03-18T23:28:09.726481Z","iopub.execute_input":"2023-03-18T23:28:09.727518Z","iopub.status.idle":"2023-03-18T23:35:39.938224Z","shell.execute_reply.started":"2023-03-18T23:28:09.727465Z","shell.execute_reply":"2023-03-18T23:35:39.937169Z"},"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":"2023-03-18T23:35:44.578139Z","iopub.execute_input":"2023-03-18T23:35:44.578704Z","iopub.status.idle":"2023-03-18T23:35:49.822790Z","shell.execute_reply.started":"2023-03-18T23:35:44.578662Z","shell.execute_reply":"2023-03-18T23:35:49.821723Z"},"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-03-18T23:35:49.824818Z","iopub.execute_input":"2023-03-18T23:35:49.825209Z","iopub.status.idle":"2023-03-18T23:35:49.832645Z","shell.execute_reply.started":"2023-03-18T23:35:49.825169Z","shell.execute_reply":"2023-03-18T23:35:49.831406Z"},"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-03-18T23:35:49.834264Z","iopub.execute_input":"2023-03-18T23:35:49.834970Z","iopub.status.idle":"2023-03-18T23:36:04.409422Z","shell.execute_reply.started":"2023-03-18T23:35:49.834933Z","shell.execute_reply":"2023-03-18T23:36:04.408010Z"},"trusted":true},"execution_count":null,"outputs":[]}]}