{"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":"## Introduction\nIn this kernel we will tackle the problem of ink detection with [Keras](https://keras.io/). We cover the following:\n\n* Loading the data using a fast `tf.data` pipeline.\n* Building and training a Keras Unet.\n* Saving the trained model in the `model.keras` format.\n* Inference on the test dataset.","metadata":{"id":"OXFHiSFraKF9"}},{"cell_type":"markdown","source":"## Imports","metadata":{"id":"rN07HUMzaWTP"}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\n\nimport random\nimport numpy as np\nfrom tqdm import tqdm\nfrom glob import glob\nimport PIL.Image as Image\nfrom matplotlib import pyplot as plt","metadata":{"id":"xVJx_bshXMax"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration","metadata":{"id":"V7wcxTsdaYBZ"}},{"cell_type":"code","source":"SHARED_HEIGHT = 4096               # Height to resize all papyri\nPATCH_SIZE = 128                   # e.g. 128x128\nPATCH_HALFSIZE = PATCH_SIZE // 2\nZ_DIM = 16                         # Number of slices in the z direction. Max value is 65 - Z_START\nZ_START = 25                       # Offset of slices in the z direction\nBATCH_SIZE = 32\n\nEPOCHS = 20\nLEARNING_RATE = 1e-5\nSTEPS_PER_EPOCH = 500\n\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\n\nIS_PROD = True\nTHRESHOLD = 0.4","metadata":{"id":"wiSwUH7kZOhE"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building the Unet\n\nHere we use the [functional api](https://keras.io/guides/functional_api/) of keras to define our entire model. It is not only idiomatic but also does away with a lot of boiler plate code. Feel free to change the hyperparameters of the model and train the model.","metadata":{"id":"N9R2JpsUab5E"}},{"cell_type":"code","source":"def double_conv(x, dims):\n    x = layers.Conv2D(dims, kernel_size=3, padding=\"SAME\")(x)\n    x = layers.BatchNormalization(epsilon=1e-5, momentum=0.1)(x)\n    x = layers.Activation(\"relu\")(x)\n    x = layers.Conv2D(dims, kernel_size=3, padding=\"SAME\")(x)\n    x = layers.BatchNormalization(epsilon=1e-5, momentum=0.1)(x)\n    return layers.Activation(\"relu\")(x)\n\ndef down_block(x, dims):\n    skip_out = double_conv(x, dims)\n    down_out = layers.MaxPool2D()(skip_out)\n    return (down_out, skip_out)\n\ndef up_block(down_input, skip_input, dims):\n    x = layers.Conv2DTranspose(dims, kernel_size=2, strides=2)(down_input)\n    x = layers.Concatenate(axis=-1)([x, skip_input])\n    return double_conv(x, dims)\n\ndef get_model(input_shape, out_classes):\n    inputs = keras.Input(input_shape)\n\n    x, skip1_out = down_block(inputs, 64)\n    x, skip2_out = down_block(x, 128)\n    x, skip3_out = down_block(x, 256)\n    x, skip4_out = down_block(x, 512)\n\n    x = double_conv(x, 1024)\n\n    x = up_block(x, skip4_out, 1024)\n    x = up_block(x, skip3_out, 512)\n    x = up_block(x, skip2_out, 256)\n    x = up_block(x, skip1_out, 128)\n\n    outputs = layers.Conv2D(out_classes, kernel_size=3, padding=\"SAME\")(x)\n    return keras.Model(inputs, outputs)","metadata":{"id":"sQEOY3ZNZOev"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building the data pipeline\n\nIn this section, we first load the dataset in RAM, and then build an efficient `tf.data` pipeline.\nIt is worth noting that the data is too large for our RAM.\nThis is the reason why we cannot load all the 64 channels of the data at once (refer to the configuration to change this setting).\nTo utilize the RAM as efficiently as we can, we also resize the papyri to a fixed spatial resolution.","metadata":{"id":"bM5vDC0ub3YC"}},{"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    return img.resize(new_size)\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 np.array(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 np.array(img, dtype=\"bool\")\n\ndef load_volume(split, index):\n    # Load the 3d x-ray scan, one slice at a time\n    z_slices_fnames = sorted(glob(f\"{DATA_DIR}/{split}/{index}/surface_volume/*.tif\"))[Z_START : Z_START+Z_DIM]\n    z_slices = []\n    for filename in  tqdm(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 np.stack(z_slices, axis=-1)\n\ndef load_sample(split, index):\n    print(f\"Loading '{split}/{index}'...\")\n    if split == \"train\":\n        return load_volume(split, index), load_mask(split, index), load_labels(split, index)\n    return load_volume(split, index), load_mask(split, index), None","metadata":{"id":"JaJMIyUrZOct"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the volume, mask, and the labels\nvolume_1, mask_1, labels_1 = load_sample(split=\"train\", index=1)\nvolume_2, mask_2, labels_2 = load_sample(split=\"train\", index=2)\nvolume_3, mask_3, labels_3 = load_sample(split=\"train\", index=3)\nprint(\"Loading complete.\")","metadata":{"id":"MRJsKKmdZOaB"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Make dev and prod folds\n\nLet's build the development and production data folds. While developing an algorithm it is important to have a validation split (to validation your algorithm).\nHere we hold out an entire papyrus for validation purpose. You can choose any developement fold (dev_1, dev_2, or dev_3) for your model training.\n\nOnce you are happy with your model, you can switch to the production dataset, which includes all the training data available, and train your model again.","metadata":{"id":"IoVjZpA0cD4y"}},{"cell_type":"code","source":"dev_folds = {\n    \"dev_1\": {\n        \"train_volumes\": [volume_1, volume_2],\n        \"train_labels\": [labels_1, labels_2],\n        \"train_masks\": [mask_1, mask_2],\n        \"validation_volume\": volume_3,\n        \"validation_labels\": labels_3,\n        \"validation_mask\": mask_3,\n    },\n    \"dev_2\": {\n        \"train_volumes\": [volume_1, volume_3],\n        \"train_labels\": [labels_1, labels_3],\n        \"train_masks\": [mask_1, mask_3],\n        \"validation_volume\": volume_2,\n        \"validation_labels\": labels_2,\n        \"validation_mask\": mask_2,\n    },\n    \"dev_3\": {\n        \"train_volumes\": [volume_2, volume_3],\n        \"train_labels\": [labels_2, labels_3],\n        \"train_masks\": [mask_2, mask_3],\n        \"validation_volume\": volume_1,\n        \"validation_labels\": labels_1,\n        \"validation_mask\": mask_1,\n    }\n}\n\nprod_data  = {\n    \"train_volumes\": [volume_1, volume_2, volume_3],\n    \"train_labels\": [labels_1, labels_2, labels_3],\n    \"train_masks\": [mask_1, mask_2, mask_3],\n}","metadata":{"id":"xwlZd9pdZOXu"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Build the `tf.data` pipeline","metadata":{"id":"mfAVj0RycHym"}},{"cell_type":"code","source":"def sample_random_location(shape):\n    x = random.randint(PATCH_HALFSIZE, shape[0] - PATCH_HALFSIZE - 1)\n    y = random.randint(PATCH_HALFSIZE, shape[1] - PATCH_HALFSIZE - 1)\n    return (x, y)\n\n\ndef list_all_locations(mask, stride=PATCH_HALFSIZE):\n    locations = []\n    for x in range(PATCH_HALFSIZE, mask.shape[0] - PATCH_HALFSIZE, stride):\n        for y in range(PATCH_HALFSIZE, mask.shape[1] - PATCH_HALFSIZE, stride):\n            if mask[x, y]:\n                locations.append((x, y))\n    return locations\n\n\ndef extract_subvolume(location, volume):\n    x = location[0]\n    y = location[1]\n    subvolume = volume[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                       y - PATCH_HALFSIZE :y + PATCH_HALFSIZE, :]\n    subvolume = subvolume.astype(\"float32\") / 65535.\n    return subvolume\n\n\ndef extract_labels(location, labels):\n    x = location[0]\n    y = location[1]\n    label = labels[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                    y - PATCH_HALFSIZE :y + PATCH_HALFSIZE]\n    label = label.astype(\"float32\")\n    label = np.expand_dims(label, axis=-1)\n    return label\n\n\ndef make_random_data_generator(volume, mask, labels):\n    def data_generator():\n        while True:\n            loc = sample_random_location(mask.shape)\n            if mask[loc[0], loc[1]]:\n                subvolume = extract_subvolume(loc, volume)\n                label = extract_labels(loc, labels)\n                yield (subvolume, label)\n    return data_generator\n\n\ndef make_iterated_data_generator(volume, mask, labels=None):\n    locations = list_all_locations(mask)\n    def data_generator():\n        for loc in locations:\n            subvolume = extract_subvolume(loc, volume)\n            if labels is None:\n                yield subvolume\n            else:\n                label = extract_labels(loc, labels)\n                yield (subvolume, label)\n    return data_generator\n\n\ndef make_tf_dataset(gen_fn, labeled=True):\n    if labeled:\n        output_signature = (\n            tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32),\n            tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, 1), dtype=tf.float32),\n        )\n    else:\n        output_signature = tf.TensorSpec(shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), dtype=tf.float32)\n    ds = tf.data.Dataset.from_generator(\n        gen_fn,\n        output_signature=output_signature,\n    )\n    return ds.batch(BATCH_SIZE)","metadata":{"id":"qfrZlfn6ZOVV"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_datasets_for_fold(fold, train_augment_fn=None):\n    train_volumes = fold[\"train_volumes\"]\n    train_masks = fold[\"train_masks\"]\n    train_labels = fold[\"train_labels\"]\n\n    include_validation = \"validation_volume\" in fold\n    if include_validation:\n        validation_volume = fold[\"validation_volume\"]\n        validation_mask = fold[\"validation_mask\"]\n        validation_labels = fold[\"validation_labels\"]\n\n    all_train_ds = []\n    for volume, mask, labels in zip(train_volumes, train_masks, train_labels):\n        train_ds = make_tf_dataset(\n            make_random_data_generator(volume, mask, labels),\n            labeled=True,\n        )\n        all_train_ds.append(train_ds)\n    train_ds = tf.data.Dataset.sample_from_datasets(all_train_ds)\n\n    if train_augment_fn:\n        train_ds = train_ds.map(train_augment_fn, num_parallel_calls=tf.data.AUTOTUNE)\n    train_ds = train_ds.prefetch(tf.data.AUTOTUNE)\n\n    if not include_validation:\n        return train_ds\n\n    val_ds = make_tf_dataset(\n        make_iterated_data_generator(validation_volume, validation_mask, validation_labels),\n        labeled=True,\n    )\n    return (train_ds, val_ds)","metadata":{"id":"krSzpW1pZOSw"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It is always good to set up a baseline. This helps us know what we have to beat in order to achieve statistical power.","metadata":{"id":"VzBsEuQmcOQ3"}},{"cell_type":"code","source":"def trivial_baseline(dataset):\n    total = 0\n    matches = 0.\n    for _, batch_label in 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\")","metadata":{"id":"uGic8UnfZOQY"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_PROD:\n    train_ds = make_datasets_for_fold(fold=prod_data)\nelse:\n    train_ds, val_ds = make_datasets_for_fold(fold=dev_folds[\"dev_1\"])\n    score = trivial_baseline(val_ds)\n    print(f\"Best validation score achievable trivially [fold 1]: {score * 100:.2f}% accuracy\")","metadata":{"id":"lXFcdqLIZOOB"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train the model","metadata":{"id":"VrnoiS-gcaq5"}},{"cell_type":"code","source":"keras.backend.clear_session()\nmodel = get_model(input_shape=(PATCH_SIZE, PATCH_SIZE, Z_DIM), out_classes=1)\nmodel.compile(\n    optimizer=\"adam\",\n    loss=\"binary_crossentropy\",\n    metrics=[\"accuracy\"]\n)\n\nif IS_PROD:\n    model.fit(train_ds, epochs=EPOCHS, steps_per_epoch=STEPS_PER_EPOCH)\n    model.save(\"prod_model.keras\")\nelse:\n    model.fit(train_ds, validation_data=val_ds, epochs=5, steps_per_epoch=100)\n    model.save(\"dev_model.keras\")","metadata":{"id":"pqPbb84vZOL4"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Once the model is trained, we free up some memory:","metadata":{}},{"cell_type":"code","source":"del volume_1\ndel volume_2\ndel volume_3\n\ndel mask_1\ndel mask_2\ndel mask_3\n\ndel labels_1\ndel labels_2\ndel labels_3\n\ndel train_ds\nif not IS_PROD:\n    del val_ds\n\n# Manually trigger garbage collection\nkeras.backend.clear_session()\nimport gc\ngc.collect()","metadata":{"id":"q9HcKCvvZtdl"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Compute predictions","metadata":{"id":"-MqNz_Mpcdyo"}},{"cell_type":"code","source":"def compute_predictions_map(split, index):\n    print(f\"Load data for {split}/{index}\")\n\n    test_volume, test_mask, _ = load_sample(split=split, index=index)\n\n    test_ds = make_tf_dataset(\n        make_iterated_data_generator(test_volume, test_mask),\n        labeled=False,\n    )\n    locations_ds = tf.data.Dataset.from_tensor_slices(\n        list_all_locations(test_mask, stride=PATCH_SIZE)\n    ).batch(BATCH_SIZE)\n\n    predictions_map = np.zeros(test_volume.shape[:2] + (1,), dtype=\"float32\")\n    predictions_map_counts = np.zeros(test_volume.shape[:2] + (1,), dtype=\"int32\")\n\n    print(f\"Compute predictions\")\n\n    for loc_batch, patch_batch in tqdm(zip(locations_ds, test_ds)):\n        predictions = model.predict_on_batch(patch_batch)\n        for (x, y), pred in zip(loc_batch, predictions):\n            predictions_map[x - PATCH_HALFSIZE : x + PATCH_HALFSIZE, y - PATCH_HALFSIZE : y + PATCH_HALFSIZE, :] += pred\n            predictions_map_counts[x - PATCH_HALFSIZE : x + PATCH_HALFSIZE, y - PATCH_HALFSIZE : y + PATCH_HALFSIZE, :] += 1\n    predictions_map /= (predictions_map_counts + 1e-7)\n    return predictions_map","metadata":{"id":"L4s6CuhjZtau"},"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":{"id":"mvd-BD-6ZtYr"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize prediction maps","metadata":{}},{"cell_type":"code","source":"plt.imshow(predictions_map_a.squeeze() > THRESHOLD, cmap=\"gray\")\nplt.show()","metadata":{"id":"85FBSrynZtWR"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(predictions_map_b.squeeze() > THRESHOLD, cmap=\"gray\")\nplt.show()","metadata":{"id":"AZmgK_kxZtT9"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resizing prediction maps to the size expected by Kaggle","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).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).squeeze()","metadata":{"id":"rdMBDQByZ6d8"},"execution_count":null,"outputs":[]},{"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":{"id":"bI49e6poZ8mZ"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Make a submission","metadata":{"id":"Q9ylDnyYchSW"}},{"cell_type":"code","source":"rle_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":{"id":"lEkQ6N7YZ_P7"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## References\n\n- https://www.kaggle.com/code/fchollet/keras-starter-kit-unet-train-on-full-dataset\n- https://www.kaggle.com/code/fchollet/a-simple-high-performance-tf-data-pipeline\n- https://www.kaggle.com/code/yururoi/pytorch-unet-baseline-with-train-code","metadata":{}}]}