{"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":"## Batch loading of image data\n### Brett Olsen, April 2023\n\nThis is an extension of earlier posted work using zarr to store chunked fragment volume data.\nThis notebook shows how to use this chunked data to train with all the training data by swapping between different slices of the data that are small enough to hold in memory.","metadata":{}},{"cell_type":"code","source":"GITHUB_REPO_URL = \"https://github.com/caethan/vesuvius_image.git\"\n!git clone {GITHUB_REPO_URL}\n!git -C /kaggle/working/vesuvius_image/ pull","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:10:22.743389Z","iopub.execute_input":"2023-04-11T01:10:22.743672Z","iopub.status.idle":"2023-04-11T01:10:25.660741Z","shell.execute_reply.started":"2023-04-11T01:10:22.743644Z","shell.execute_reply":"2023-04-11T01:10:25.659337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os\nclass HiddenPrints:\n    def __enter__(self):\n        self._original_stdout = sys.stdout\n        sys.stdout = open(os.devnull, 'w')\n\n    def __exit__(self, exc_type, exc_val, exc_tb):\n        sys.stdout.close()\n        sys.stdout = self._original_stdout","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:10:25.667443Z","iopub.execute_input":"2023-04-11T01:10:25.670237Z","iopub.status.idle":"2023-04-11T01:10:25.679984Z","shell.execute_reply.started":"2023-04-11T01:10:25.670183Z","shell.execute_reply":"2023-04-11T01:10:25.678541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We're using this context manager to do our conda env update without massively cluttering up the output\nwith HiddenPrints():\n    !mamba env update --name base --file vesuvius_image/conda_env.yaml","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:10:25.685454Z","iopub.execute_input":"2023-04-11T01:10:25.688229Z","iopub.status.idle":"2023-04-11T01:15:02.176098Z","shell.execute_reply.started":"2023-04-11T01:10:25.688179Z","shell.execute_reply":"2023-04-11T01:15:02.174772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.utils.data as data\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:15:02.179036Z","iopub.execute_input":"2023-04-11T01:15:02.179381Z","iopub.status.idle":"2023-04-11T01:15:03.324600Z","shell.execute_reply.started":"2023-04-11T01:15:02.179346Z","shell.execute_reply":"2023-04-11T01:15:03.323461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if \"/kaggle/working/vesuvius_image/\" not in sys.path:\n    sys.path.append(\"/kaggle/working/vesuvius_image/\")","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:15:04.310905Z","iopub.execute_input":"2023-04-11T01:15:04.311336Z","iopub.status.idle":"2023-04-11T01:15:04.318695Z","shell.execute_reply.started":"2023-04-11T01:15:04.311298Z","shell.execute_reply":"2023-04-11T01:15:04.316646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from vesuvius.image import PapyrusImage\nfrom vesuvius.volume import PapyrusVolume\nfrom vesuvius.utils import Timer\n\nfrom tqdm import tqdm\nimport numpy as np\nfrom typing import Union\nimport scipy\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:15:18.121860Z","iopub.execute_input":"2023-04-11T01:15:18.122581Z","iopub.status.idle":"2023-04-11T01:15:18.295677Z","shell.execute_reply.started":"2023-04-11T01:15:18.122541Z","shell.execute_reply":"2023-04-11T01:15:18.293551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_FOLDER = \"/kaggle/input/vesuvius-challenge-ink-detection/\"\nWORKING_FOLDER = \"/kaggle/working/\"","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:16:09.116686Z","iopub.execute_input":"2023-04-11T01:16:09.117205Z","iopub.status.idle":"2023-04-11T01:16:09.122424Z","shell.execute_reply.started":"2023-04-11T01:16:09.117166Z","shell.execute_reply":"2023-04-11T01:16:09.121251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This cell will generate the zarr files if they don't already exist.\n# Note we're running with \"multiscale=False\" to avoid using more space than we need for the\n# downsampled copies.\nfor subdir in [\"1\", \"2\", \"3\"]:\n    train_input_dir = os.path.join(INPUT_FOLDER, \"train\", subdir)\n    train_output_dir = os.path.join(WORKING_FOLDER, f\"train_{subdir}\")\n    os.makedirs(train_output_dir, exist_ok=True)\n    print(train_output_dir)\n    # First the surface volume stack\n    frag_output = os.path.join(train_output_dir, \"frag.volzarr\")\n    if not os.path.exists(frag_output):\n        # TODO: experiment with different chunk sizes here\n        PapyrusVolume.build_from_tiffdir(\n            os.path.join(train_input_dir, \"surface_volume\"), \n            frag_output, \n            multiscale=False,\n        )\n    # The fragment mask\n    mask_output = os.path.join(train_output_dir, \"mask.imagezarr\")\n    if not os.path.exists(mask_output):\n        PapyrusImage.build_from_image(\n            os.path.join(train_input_dir, \"mask.png\"),\n            mask_output,\n            multiscale=False,\n        )\n    # The truth set\n    ink_output = os.path.join(train_output_dir, \"inklabels.imagezarr\")\n    if not os.path.exists(ink_output):\n        PapyrusImage.build_from_image(\n            os.path.join(train_input_dir, \"inklabels.png\"),\n            ink_output,\n            multiscale=False\n        )","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:16:10.012093Z","iopub.execute_input":"2023-04-11T01:16:10.012469Z","iopub.status.idle":"2023-04-11T01:25:29.085466Z","shell.execute_reply.started":"2023-04-11T01:16:10.012435Z","shell.execute_reply":"2023-04-11T01:25:29.084451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BUFFER = 30  # Buffer size in x and y direction\nZ_START = 27 # First slice in the z direction to use\nZ_DIM = 10   # Number of slices in the z direction\nZ_SLICE = slice(Z_START, Z_START + Z_DIM, None)\n\n\nSUBFOLDERS = [\"train_1\", \"train_2\", \"train_3\"]\n\ndef load_data_slice(query_index, max_mem=1e9, count_only=False):\n    \"\"\"Loads a subset of the training data into GPU memory.  Slicing is done so as to\n    produce roughly square regions of each dataset so that the slices do not exceed\n    the required maximum memory size.\n    \n    Note that the mask data is only loaded into CPU, not onto GPU.\n    \"\"\"\n    # one float stored for the labels and one for each level of Z used\n    bytes_per_pixel = (Z_DIM + 1) * np.dtype(\"float32\").itemsize\n    \n    slice_index = 0\n    for subfolder in SUBFOLDERS:\n        if count_only:\n            print(slice_index)\n        with HiddenPrints():\n            labelzarr = PapyrusImage(f\"{WORKING_FOLDER}/{subfolder}/inklabels.imagezarr\")\n            fragzarr = PapyrusVolume(f\"{WORKING_FOLDER}/{subfolder}/frag.volzarr\")\n            maskzarr = PapyrusImage(f\"{WORKING_FOLDER}/{subfolder}/mask.imagezarr\")\n        \n        # Keep cutting the largest dimensions in half until we're small enough to fit\n        # in our memory limits.\n        xsize, ysize = labelzarr.shape\n        xcount, ycount = 1, 1\n        while (xsize + BUFFER) * (ysize + BUFFER) * bytes_per_pixel > max_mem:\n            if xsize > ysize:\n                xsize = xsize // 2\n                xcount += 1\n            else:\n                ysize = ysize // 2\n                ycount += 1\n        for x in range(xcount):\n            for y in range(ycount):\n                if (slice_index == query_index) and not count_only:\n                    xslice = slice(max(0, xsize * x - BUFFER),\n                                   min(labelzarr.shape[0], xsize * (x + 1) + BUFFER),\n                                   None)\n                    yslice = slice(max(0, ysize * y - BUFFER),\n                                   min(labelzarr.shape[1], ysize * (y + 1) + BUFFER),\n                                   None)\n                    labels = torch.from_numpy(\n                        labelzarr[xslice,yslice,0].astype(np.float32)\n                    ).to(DEVICE)\n                    surface_volume = torch.from_numpy(\n                        fragzarr[xslice,yslice,Z_SLICE,0].astype(np.float32) / np.iinfo(np.uint16).max\n                    ).to(DEVICE)\n                    mask = maskzarr[xslice,yslice,0]\n                    return labels, surface_volume, mask\n                slice_index += 1\n    \n    if count_only:\n        return slice_index\n    \n    raise ValueError(f\"Query index {query_index} too large; only {slice_index} slices of data.\")\n      \n# Here we're running it with \"count_only\" so that we can see how many slices of data there\n# are in our zarrs.\ntotal_slices = load_data_slice(0, max_mem=1e9, count_only=True)\nprint(total_slices)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:28:32.566561Z","iopub.execute_input":"2023-04-11T01:28:32.566987Z","iopub.status.idle":"2023-04-11T01:28:32.729145Z","shell.execute_reply.started":"2023-04-11T01:28:32.566924Z","shell.execute_reply":"2023-04-11T01:28:32.728011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Great.  So there's 12 total slices, slices 0-3 are from fragment 1, slices 4-9 are from fragment 2, and slices 10-11 are from fragment 3.\nWe'll select one slice from each fragment to be validation data and the remaining slices will provide training data.","metadata":{}},{"cell_type":"code","source":"validation_slices = [2, 6, 11]\ntraining_slices = [i for i in range(total_slices) if i not in validation_slices]\nprint(training_slices)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:28:34.594769Z","iopub.execute_input":"2023-04-11T01:28:34.595957Z","iopub.status.idle":"2023-04-11T01:28:34.602085Z","shell.execute_reply.started":"2023-04-11T01:28:34.595893Z","shell.execute_reply":"2023-04-11T01:28:34.600918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we will generate a custom PyTorch dataset:  if provided with a given slice index, it will load\nthe data from that slice into GPU memory.\nThen every pixel (within the mask) inside that slice is available to access a small region around that pixel.","metadata":{}},{"cell_type":"code","source":"class SlicedDataset(data.Dataset):\n    def __init__(self, query_slice):\n        self.query_slice = query_slice\n        labels, surface_volume, mask = load_data_slice(query_slice)\n        self.labels = labels\n        self.surface_volume = surface_volume\n        self.mask = mask\n        self.xvals, self.yvals = np.where(self.mask)\n        xmask = ((self.xvals - BUFFER) > 0) & ((self.xvals + BUFFER) < self.labels.shape[0])\n        ymask = ((self.yvals - BUFFER) > 0) & ((self.yvals + BUFFER) < self.labels.shape[1])\n        self.xvals = self.xvals[xmask & ymask]\n        self.yvals = self.yvals[xmask & ymask]\n        \n    def __len__(self):\n        return len(self.xvals)\n    \n    def __getitem__(self, index):\n        x = self.xvals[index]\n        y = self.yvals[index]\n        subvolume = self.surface_volume[x-BUFFER:x+BUFFER+1, y-BUFFER:y+BUFFER+1, :].view(1, BUFFER*2+1, BUFFER*2+1, Z_DIM)\n        inklabel = self.labels[x, y].view(1)\n        return x, y, subvolume, inklabel","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:28:40.534277Z","iopub.execute_input":"2023-04-11T01:28:40.534654Z","iopub.status.idle":"2023-04-11T01:28:40.545248Z","shell.execute_reply.started":"2023-04-11T01:28:40.534622Z","shell.execute_reply":"2023-04-11T01:28:40.543750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And here's a dead simple neural network model (the same one used for the demo from the organizers).\nBut now I'll show how to train and validate this model with the entire data from all three fragments.","metadata":{}},{"cell_type":"code","source":"model = nn.Sequential(\n    nn.Conv3d(in_channels=1, out_channels=16, kernel_size=3, stride=1, padding=1),\n    nn.MaxPool3d(kernel_size=2, stride=2),\n    nn.Conv3d(in_channels=16, out_channels=32, kernel_size=3, stride=1, padding=1),\n    nn.MaxPool3d(kernel_size=2, stride=2),\n    nn.Conv3d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=1),\n    nn.MaxPool3d(kernel_size=2, stride=2),\n    #Flattens dimensions 1+ into a single dimension, so it is a dim=2 matrix afterwards\n    nn.Flatten(start_dim=1),\n    nn.LazyLinear(out_features=128),\n    nn.ReLU(),\n    nn.LazyLinear(out_features=1),\n    # Convert from (-inf, inf) to (0, 1) probability with a sigmoid\n    nn.Sigmoid()\n).to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:28:53.066186Z","iopub.execute_input":"2023-04-11T01:28:53.067197Z","iopub.status.idle":"2023-04-11T01:28:55.998775Z","shell.execute_reply.started":"2023-04-11T01:28:53.067142Z","shell.execute_reply":"2023-04-11T01:28:55.997712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we're going to train (badly) our model to show how we load data from different slices in and out of GPU memory. ","metadata":{}},{"cell_type":"code","source":"training_losses = []\nvalidation_losses = []","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:29:48.299883Z","iopub.execute_input":"2023-04-11T01:29:48.300604Z","iopub.status.idle":"2023-04-11T01:29:48.305724Z","shell.execute_reply.started":"2023-04-11T01:29:48.300565Z","shell.execute_reply":"2023-04-11T01:29:48.304649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Training...\")\nTRAINING_STEPS_PER_SLICE = 64\nLEARNING_RATE = 0.03\nBATCH_SIZE = 16\nEPOCH_COUNT = 5\n\noptimizer = optim.SGD(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCELoss()\n\nfor i in range(EPOCH_COUNT):\n    print(f\"Processing epoch {i+1}\")\n\n    epoch_training_loss = [0.0, 0]\n    for tr_item in training_slices:\n        train_dataset = SlicedDataset(tr_item)\n        if len(train_dataset) == 0:\n            # Just in case one of the slices is entirely empty (because of the mask)\n            continue\n        print(f\"Training on data slice {tr_item}\")\n        training_steps = min(len(train_dataset), TRAINING_STEPS_PER_SLICE)\n        train_loader = data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=LEARNING_RATE, total_steps=training_steps\n        )\n        model.train()\n        for i, (_, _, subvolumes, inklabels) in enumerate(train_loader):\n            if i >= training_steps:\n                break\n            optimizer.zero_grad()\n            outputs = model(subvolumes.to(DEVICE))\n            loss = criterion(outputs, inklabels.to(DEVICE))\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            epoch_training_loss[0] += loss.item()\n            epoch_training_loss[1] += 1\n        del train_dataset, train_loader, scheduler, _, loss, subvolumes, inklabels\n        torch.cuda.empty_cache()\n    training_losses.append(epoch_training_loss[0] / epoch_training_loss[1])\n        \n    val_training_loss = [0.0, 0]\n    for val_item in validation_slices:\n        val_dataset = SlicedDataset(val_item)\n        if len(val_dataset) == 0:\n            continue\n        print(f\"Validating on data slice {val_item}\")\n        val_steps = min(len(val_dataset), TRAINING_STEPS_PER_SLICE)\n        val_loader = data.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=True)\n        model.eval()\n        with torch.no_grad():\n            for i, (_, _, subvolumes, inklabels) in enumerate(val_loader):\n                if i > val_steps:\n                    break\n                outputs = model(subvolumes.to(DEVICE))\n                loss = criterion(outputs, inklabels.to(DEVICE))\n                val_training_loss[0] += loss.item()\n                val_training_loss[1] += 1\n        del val_dataset, val_loader, _, loss, subvolumes, inklabels\n        torch.cuda.empty_cache()\n    validation_losses.append(val_training_loss[0] / val_training_loss[1])\n        \ndel optimizer\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:29:49.397167Z","iopub.execute_input":"2023-04-11T01:29:49.397603Z","iopub.status.idle":"2023-04-11T01:33:28.304901Z","shell.execute_reply.started":"2023-04-11T01:29:49.397567Z","shell.execute_reply":"2023-04-11T01:33:28.303600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure()\nplt.grid()\n\nplt.plot(training_losses, label=\"Training\")\nplt.plot(validation_losses, label=\"Validation\")\n\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:33:43.002896Z","iopub.execute_input":"2023-04-11T01:33:43.003624Z","iopub.status.idle":"2023-04-11T01:33:43.275988Z","shell.execute_reply.started":"2023-04-11T01:33:43.003585Z","shell.execute_reply":"2023-04-11T01:33:43.274979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"So you can see that the model and training performance is pretty bad, but the actual data loading works fine:  we just work with each bit of memory separately for either training, validation, or inference.","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}