{"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":"# A simple, high-performance tf.data pipeline, with data augmentation\n\nThis notebook demonstrates a super simple, readable, and high-performance data pipeline written in tf.data.\n\nIt has the following features:\n\n- We can downscale images to save RAM\n- The data is loaded in NumPy arrays, and only a single copy of the data is kept in memory, to avoid OOM\n- We can train on as many papyruses as we want at once (we sample randomly from one to the other)\n- We can iterate over our data either randomly or iteratively (in a grid-like fashion)\n- TF-based data augmentation\n- Extremely fast!","metadata":{}},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\n\nimport random\nimport gc\nimport cv2\nimport time\nfrom tqdm import tqdm\n\n\nDATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nPATCH_SIZE = 128  # e.g. 128x128\nPATCH_HALFSIZE = PATCH_SIZE // 2\nDOWNSAMPLING = 0.75  # Setting this to e.g. 0.5 means images will be loaded as 2x smaller. 1 does nothing.\nZ_DIM = 55   # Number of slices in the z direction. Max value is 65 - Z_START\nZ_START = 0  # Offset of slices in the z direction\nBATCH_SIZE = 8","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:46:11.919824Z","iopub.execute_input":"2023-03-24T16:46:11.920205Z","iopub.status.idle":"2023-03-24T16:46:11.926933Z","shell.execute_reply.started":"2023-03-24T16:46:11.920173Z","shell.execute_reply":"2023-03-24T16:46:11.925841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load the training data in CPU RAM as NumPy arrays","metadata":{}},{"cell_type":"code","source":"def resize(img):\n    if DOWNSAMPLING != 1.:\n        size = int(img.shape[1] * DOWNSAMPLING), int(img.shape[0] * DOWNSAMPLING)\n        img = cv2.resize(img, size)\n    return img\n\n\ndef load_mask(split, index):\n    img = cv2.imread(f\"{DATA_DIR}/{split}/{index}/mask.png\", 0)\n    img = resize(img)\n    return img.astype(\"bool\")\n\n\ndef load_labels(split, index):\n    img = cv2.imread(f\"{DATA_DIR}/{split}/{index}/inklabels.png\", 0)\n    img = resize(img)\n    return np.expand_dims(img, axis=-1)\n\n\ndef load_volume(split, index):\n    # A more memory-efficient volune loader\n    fnames = [f\"{DATA_DIR}/{split}/{index}/surface_volume/{i:02}.tif\"\n             for i in range(Z_START, Z_START + Z_DIM)]\n\n    batch_size = 8\n    fname_batches = [fnames[i :i + batch_size] for i in range(0, len(fnames), batch_size)]\n    volumes = []\n    for fname_batch in fname_batches:\n        z_slices = []\n        for fname in tqdm(fname_batch):\n            img = cv2.imread(fname, 0)\n            img = resize(img)\n            z_slices.append(img)\n        volumes.append(np.stack(z_slices, axis=-1))\n        del z_slices\n    return np.concatenate(volumes, axis=-1)\n\n\ndef load_sample(split, index):\n    print(f\"Loading '{split}/{index}'...\")\n    gc.collect()\n    if split == \"train\":\n        return load_volume(split, index), load_mask(split, index), load_labels(split, index)\n    return load_volume(split, index), load_mask(split, index), None","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:25:32.149409Z","iopub.execute_input":"2023-03-24T16:25:32.152013Z","iopub.status.idle":"2023-03-24T16:25:32.170086Z","shell.execute_reply.started":"2023-03-24T16:25:32.151973Z","shell.execute_reply":"2023-03-24T16:25:32.168708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"volume_1, mask_1, labels_1 = load_sample(split=\"train\", index=1)\nvolume_2, mask_2, labels_2 = load_sample(split=\"train\", index=2)\nvolume_3, mask_3, labels_3 = load_sample(split=\"train\", index=3)\ngc.collect()\nprint(\"Loading complete.\")","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:25:32.175136Z","iopub.execute_input":"2023-03-24T16:25:32.177813Z","iopub.status.idle":"2023-03-24T16:31:56.435613Z","shell.execute_reply.started":"2023-03-24T16:25:32.177754Z","shell.execute_reply":"2023-03-24T16:31:56.433115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define training / validation / production folds","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-03-24T16:46:18.730259Z","iopub.execute_input":"2023-03-24T16:46:18.730619Z","iopub.status.idle":"2023-03-24T16:46:18.753592Z","shell.execute_reply.started":"2023-03-24T16:46:18.730587Z","shell.execute_reply":"2023-03-24T16:46:18.752451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define data augmentation pipeline\n\nBe mindful that any geometric transformation on the input patches should be replicated on the label patches.","metadata":{}},{"cell_type":"code","source":"@tf.function\ndef train_augment_fn(patch, labels):\n    patch = tf.image.random_flip_left_right(patch, seed=1337)\n    labels = tf.image.random_flip_left_right(labels, seed=1337)\n    \n    patch = tf.image.random_flip_up_down(patch, seed=42)\n    labels = tf.image.random_flip_up_down(labels, seed=42)\n    return patch, labels","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:46:20.547629Z","iopub.execute_input":"2023-03-24T16:46:20.548343Z","iopub.status.idle":"2023-03-24T16:46:20.554734Z","shell.execute_reply.started":"2023-03-24T16:46:20.548304Z","shell.execute_reply":"2023-03-24T16:46:20.553648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create training & validation datasets that yield patches and their labels\n\nLet's define some utilities for training and validation datasets. It's all super simple, just very short functions composed together.\n\nWe have two ways to iterate over the data:\n\n- **Randomly**: We sample random patches within one of multiple volumes. This is what we use for our training data.\n- **Iteratively**: We sample patches in a deterministic grid-like pattern that provides us full coverage of the masked area of a given volume. This is what we use for our validation data and our test (production) data.","metadata":{}},{"cell_type":"code","source":"def sample_random_location(shape):\n    x = random.randint(PATCH_HALFSIZE, shape[0] - PATCH_HALFSIZE - 1)\n    y = random.randint(PATCH_HALFSIZE, shape[1] - PATCH_HALFSIZE - 1)\n    return (x, y)\n\n\ndef list_all_locations(mask, stride=PATCH_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_patch(location, volume):\n    x = location[0]\n    y = location[1]\n    patch = volume[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                   y - PATCH_HALFSIZE :y + PATCH_HALFSIZE, :]\n    return patch.astype(\"float32\") / 255.\n\n\ndef extract_labels(location, labels):\n    x = location[0]\n    y = location[1]\n    \n    label = labels[x - PATCH_HALFSIZE :x + PATCH_HALFSIZE,\n                   y - PATCH_HALFSIZE :y + PATCH_HALFSIZE, :]\n    return label.astype(\"float32\") / 255.\n\n\ndef make_random_data_generator(volume, mask, labels):\n    def data_generator():\n        while True:\n            loc = sample_random_location(mask.shape)\n            if mask[loc[0], loc[1]]:\n                patch = extract_patch(loc, volume)\n                label = extract_labels(loc, labels)\n                yield patch, label\n    return data_generator\n\n\ndef make_iterated_data_generator(volume, mask, labels=None):\n    locations = list_all_locations(mask)\n    def data_generator():\n        for loc in locations:\n            patch = extract_patch(loc, volume)\n            if labels is None:\n                yield patch\n            else:\n                label = extract_labels(loc, labels)\n                yield patch, 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.prefetch(tf.data.AUTOTUNE).batch(BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:46:22.133439Z","iopub.execute_input":"2023-03-24T16:46:22.134395Z","iopub.status.idle":"2023-03-24T16:46:22.149161Z","shell.execute_reply.started":"2023-03-24T16:46:22.134357Z","shell.execute_reply":"2023-03-24T16:46:22.148052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Next, let's prepare a utility that turns a fold's metadata into a training dataset and a validation dataset.","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-03-24T16:46:23.675711Z","iopub.execute_input":"2023-03-24T16:46:23.676842Z","iopub.status.idle":"2023-03-24T16:46:23.685686Z","shell.execute_reply.started":"2023-03-24T16:46:23.676761Z","shell.execute_reply":"2023-03-24T16:46:23.684571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check dataset output shapes","metadata":{}},{"cell_type":"code","source":"# We use the first fold datasets for this quick check\ntrain_ds, val_ds = make_datasets_for_fold(dev_folds[\"dev_1\"])\n\nfor patch, label in train_ds.take(1):\n    print(f\"Train patch shape: {patch.shape}\")\n    print(f\"Train label shape: {label.shape}\")\nprint(\"-\")\nfor patch, label in val_ds.take(1):\n    print(f\"Val patch shape: {patch.shape}\")\n    print(f\"Val label shape: {label.shape}\")","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:46:26.368500Z","iopub.execute_input":"2023-03-24T16:46:26.369225Z","iopub.status.idle":"2023-03-24T16:46:26.613282Z","shell.execute_reply.started":"2023-03-24T16:46:26.369188Z","shell.execute_reply":"2023-03-24T16:46:26.612164Z"},"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!\n\nIn this case, our datasets are *insanely* fast because they're compiled to natively parallel C++, rather than using multiprocessing in the Python runtime.","metadata":{}},{"cell_type":"code","source":"def check_throughout(ds):\n    n = 100\n    for i, _ in enumerate(ds.take(n + 1)):\n        if i == 1:  # Don't include dataset initialization time\n            t0 = time.time()\n    time_per_batch = (time.time() - t0) / n\n    print(f\"Time per batch: {time_per_batch:.4f}s\")\n    print(f\"Time per sample: {time_per_batch / BATCH_SIZE:.4f}s\")\n\n\ntrain_ds, val_ds = make_datasets_for_fold(dev_folds[\"dev_1\"])\n\ncheck_throughout(train_ds)","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:50:41.514614Z","iopub.execute_input":"2023-03-24T16:50:41.515328Z","iopub.status.idle":"2023-03-24T16:50:43.588934Z","shell.execute_reply.started":"2023-03-24T16:50:41.515288Z","shell.execute_reply":"2023-03-24T16:50:43.587846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"With data augmentation:","metadata":{}},{"cell_type":"code","source":"train_ds, val_ds = make_datasets_for_fold(dev_folds[\"dev_1\"], train_augment_fn=train_augment_fn)\n\ncheck_throughout(train_ds)","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:50:44.695439Z","iopub.execute_input":"2023-03-24T16:50:44.696300Z","iopub.status.idle":"2023-03-24T16:50:51.692926Z","shell.execute_reply.started":"2023-03-24T16:50:44.696240Z","shell.execute_reply":"2023-03-24T16:50:51.691901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Usage examples\n\nNote that the training datasets are infinitely streaming -- iterating over them will just yield new random patches forever.\nWhen using these datasets in `fit()`, make sure to provide the `steps_per_epoch=N` argument to draw N batches per epoch.","metadata":{}},{"cell_type":"code","source":"# This is how you create the datasets for a given fold:\ntrain_ds, val_ds = make_datasets_for_fold(dev_folds[\"dev_2\"], train_augment_fn=train_augment_fn)\n\n# This is how you create your production training dataset:\ntrain_ds = make_datasets_for_fold(prod_data, train_augment_fn=train_augment_fn)\n\n# This is how you create your production test dataset:\ntest_volume, test_mask, _ = load_sample(split=\"test\", index=\"a\")\ntest_ds = make_tf_dataset(\n    make_iterated_data_generator(test_volume, test_mask),\n    labeled=False,\n)\n\n# This is how you iterate over the production test dataset\n# and retrieve the locations (x, y) corresponding to the center pixel\n# of the current patch:\nlocations_ds = tf.data.Dataset.from_tensor_slices(\n    list_all_locations(test_mask, stride=PATCH_SIZE)\n).batch(BATCH_SIZE)\nfor loc_batch, patch_batch in tqdm(zip(locations_ds, test_ds)):\n    index = 3\n    x, y = loc_batch[index]\n    print(f\"Patch patch_batch[{index}] is for location ({x}, {y})\")\n    break","metadata":{"execution":{"iopub.status.busy":"2023-03-24T16:51:01.564182Z","iopub.execute_input":"2023-03-24T16:51:01.565121Z"},"trusted":true},"execution_count":null,"outputs":[]}]}