{"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":"This is a notebook explaining the [Ink Detection progress prize on Kaggle](https://www.kaggle.com/competitions/vesuvius-challenge), which is part of the larger [Vesuvius Challenge](https://scrollprize.org).\n\nFor more background on the process of ink detection, be sure to check out [Tutorial 4: Ink Detection](https://scrollprize.org/tutorial4) on the Vesuvius Challenge website.\n\nIn this notebook we'll see how to train a simple ML model to detect ink in a papyrus fragment from a 3d x-ray scan of the fragment.\n\n<img src=\"https://user-images.githubusercontent.com/177461/224853397-3cf86dc2-45b4-4e7c-9ec2-28a733791a75.jpg\" width=\"200\"/>\n\nFirst, initialize some variables, and let's look at a photo of the fragment. We won't use this for training, but it's useful to see.\n\nIt's an infrared photo, since the ink is better visible in infrared light.","metadata":{"_uuid":"748d3554-7759-4443-b1b8-916e19ca50ed","_cell_guid":"a18c7d0f-a17a-4603-956d-2b531016c536","trusted":true}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport glob\nimport PIL.Image as Image\nimport torch.utils.data as data\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom tqdm import tqdm\nfrom ipywidgets import interact, fixed\n\nPREFIX = '/kaggle/input/vesuvius-challenge/train/1/'\nBUFFER = 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\nTRAINING_STEPS = 30000\nLEARNING_RATE = 0.03\nBATCH_SIZE = 32\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nplt.imshow(Image.open(PREFIX+\"ir.png\"), cmap=\"gray\")","metadata":{"_uuid":"91db1348-a896-4607-8686-f6c6df6419ed","_cell_guid":"ff3c9fb9-0c86-4acf-9162-c741c46e53a4","collapsed":false,"jupyter":{"outputs_hidden":false},"_kg_hide-output":false,"execution":{"iopub.status.busy":"2023-04-01T16:10:47.066093Z","iopub.execute_input":"2023-04-01T16:10:47.066377Z","iopub.status.idle":"2023-04-01T16:10:54.384715Z","shell.execute_reply.started":"2023-04-01T16:10:47.066349Z","shell.execute_reply":"2023-04-01T16:10:54.383748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's load these binary images:\n* **mask.png**: a mask of which pixels contain data, and which pixels we should ignore.\n* **inklabels.png**: our label data: whether a pixel contains ink or no ink (which has been hand-labeled based on the infrared photo).","metadata":{"_uuid":"ee5f3f31-6cc0-4dbf-b4e1-e9afb18bc0ea","_cell_guid":"7607926b-5a4d-4107-8935-955288d53ebd","trusted":true}},{"cell_type":"code","source":"mask = np.array(Image.open(PREFIX+\"mask.png\").convert('1'))\nlabel = torch.from_numpy(np.array(Image.open(PREFIX+\"inklabels.png\"))).gt(0).float().to(DEVICE)\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.set_title(\"mask.png\")\nax1.imshow(mask, cmap='gray')\nax2.set_title(\"inklabels.png\")\nax2.imshow(label.cpu(), cmap='gray')\nplt.show()","metadata":{"_uuid":"351fefa9-30dd-4e8d-bf3b-1aaa0fb33905","_cell_guid":"43a9acbe-f2c5-4976-b9ee-027e62c27a83","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:10:54.389441Z","iopub.execute_input":"2023-04-01T16:10:54.391820Z","iopub.status.idle":"2023-04-01T16:11:00.462506Z","shell.execute_reply.started":"2023-04-01T16:10:54.391778Z","shell.execute_reply":"2023-04-01T16:11:00.461357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Next, we'll load the 3d x-ray of the fragment. This is represented as a .tif image stack. The image stack is an array of 16-bit grayscale images. Each image represents a \"slice\" in the z-direction, going from below the papyrus, to above the papyrus. We'll convert it to a 4D tensor of 32-bit floats. We'll also convert the pixel values to the range [0, 1].\n\nTo save memory, we'll only load the innermost slices (`Z_DIM` of them). Let's look at them when we're done.","metadata":{"_uuid":"74a6ae91-2aa2-4eb3-9f3c-956dc54cf017","_cell_guid":"09e95f98-439b-49c3-aae2-550fadd533df","trusted":true}},{"cell_type":"code","source":"# Load the 3d x-ray scan, one slice at a time\nimages = [np.array(Image.open(filename), dtype=np.float32)/65535.0 for filename in tqdm(sorted(glob.glob(PREFIX+\"surface_volume/*.tif\"))[Z_START:Z_START+Z_DIM])]\nimage_stack = torch.stack([torch.from_numpy(image) for image in images], dim=0).to(DEVICE)\n\nfig, axes = plt.subplots(1, len(images), figsize=(15, 3))\nfor image, ax in zip(images, axes):\n  ax.imshow(np.array(Image.fromarray(image).resize((image.shape[1]//20, image.shape[0]//20)), dtype=np.float32), cmap='gray')\n  ax.set_xticks([]); ax.set_yticks([])\nfig.tight_layout()\nplt.show()","metadata":{"_uuid":"929b9e7a-d30c-4462-a7b9-0765a76cdb4e","_cell_guid":"f75858f9-06ad-43cc-be5f-ab09738a58c1","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:11:00.464318Z","iopub.execute_input":"2023-04-01T16:11:00.464669Z","iopub.status.idle":"2023-04-01T16:11:18.037690Z","shell.execute_reply.started":"2023-04-01T16:11:00.464635Z","shell.execute_reply":"2023-04-01T16:11:18.036622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Can you see the ink in these slices of the 3d x-ray scan..? Neither can we.\n\nNow we'll create a dataset of subvolumes. We use a small rectangle around the letter \"P\" for our evaluation, and we'll exclude those pixels from the training set. (It's actually a Greek letter \"rho\", which looks similar to our \"P\".)","metadata":{"_uuid":"18a40c3b-4241-4f8d-b808-e1cab35b8f6e","_cell_guid":"173e0034-8ce4-4a12-9a0b-43cff0022249","trusted":true}},{"cell_type":"code","source":"rect = (1100, 3500, 700, 950)\nfig, ax = plt.subplots()\nax.imshow(label.cpu())\npatch = patches.Rectangle((rect[0], rect[1]), rect[2], rect[3], linewidth=2, edgecolor='r', facecolor='none')\nax.add_patch(patch)\nplt.show()","metadata":{"_uuid":"39a855d9-38f5-4f2d-ac7f-1f65529d1a3e","_cell_guid":"62d60bbe-85f1-4df4-b1e6-f4337268b11c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:11:18.040303Z","iopub.execute_input":"2023-04-01T16:11:18.041229Z","iopub.status.idle":"2023-04-01T16:11:19.931215Z","shell.execute_reply.started":"2023-04-01T16:11:18.041191Z","shell.execute_reply":"2023-04-01T16:11:19.930071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we'll define a PyTorch dataset and (super simple) model.","metadata":{"_uuid":"4061cd69-3dd7-495c-838f-17280c1041b2","_cell_guid":"456fbd67-b8ed-4622-925a-3efc9fb7eb4a","trusted":true}},{"cell_type":"code","source":"class SubvolumeDataset(data.Dataset):\n    def __init__(self, image_stack, label, pixels):\n        self.image_stack = image_stack\n        self.label = label\n        self.pixels = pixels\n    def __len__(self):\n        return len(self.pixels)\n    def __getitem__(self, index):\n        y, x = self.pixels[index]\n        subvolume = self.image_stack[:, y-BUFFER:y+BUFFER+1, x-BUFFER:x+BUFFER+1].view(1, Z_DIM, BUFFER*2+1, BUFFER*2+1)\n        inklabel = self.label[y, x].view(1)\n        return subvolume, inklabel\n\nmodel = nn.Sequential(\n    nn.Conv3d(1, 16, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Conv3d(16, 32, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Conv3d(32, 64, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Flatten(start_dim=1),\n    nn.LazyLinear(128), nn.ReLU(),\n    nn.LazyLinear(1), nn.Sigmoid()\n).to(DEVICE)","metadata":{"_uuid":"9cf6abaa-b06b-4d05-a6d8-b7f012cf2e2c","_cell_guid":"25930dc1-1883-4fcd-9d40-f4aaa2ca152b","collapsed":false,"jupyter":{"outputs_hidden":false},"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-01T16:11:19.932739Z","iopub.execute_input":"2023-04-01T16:11:19.934319Z","iopub.status.idle":"2023-04-01T16:11:19.957414Z","shell.execute_reply.started":"2023-04-01T16:11:19.934277Z","shell.execute_reply":"2023-04-01T16:11:19.956089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we'll train the model. Conceptually it looks like this:\n\n<a href=\"https://user-images.githubusercontent.com/22727759/224853655-3fad9edb-c798-452e-94d0-f74efe71c08e.mp4\"><img src=\"https://user-images.githubusercontent.com/22727759/224853385-ed190d89-f466-469c-82a9-499881759d57.gif\"/></a>\n\nThis typically takes about 10 minutes.","metadata":{"_uuid":"cfbc9413-b6ff-4a69-ae60-22dddef18f11","_cell_guid":"e7b84e1e-7903-4290-9e09-3bf2cb62b222","trusted":true}},{"cell_type":"code","source":"print(\"Generating pixel lists...\")\n# Split our dataset into train and val. The pixels inside the rect are the \n# val set, and the pixels outside the rect are the train set.\npixels_inside_rect = []\npixels_outside_rect = []\nfor pixel in zip(*np.where(mask == 1)):\n    if pixel[1] < BUFFER or pixel[1] >= mask.shape[1]-BUFFER or pixel[0] < BUFFER or pixel[0] >= mask.shape[0]-BUFFER:\n        continue # Too close to the edge\n    if pixel[1] >= rect[0] and pixel[1] <= rect[0]+rect[2] and pixel[0] >= rect[1] and pixel[0] <= rect[1]+rect[3]:\n        pixels_inside_rect.append(pixel)\n    else:\n        pixels_outside_rect.append(pixel)\n\nprint(\"Training...\")\ntrain_dataset = SubvolumeDataset(image_stack, label, pixels_outside_rect)\ntrain_loader = data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\ncriterion = nn.BCELoss()\noptimizer = optim.SGD(model.parameters(), lr=LEARNING_RATE)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=LEARNING_RATE, total_steps=TRAINING_STEPS)\nmodel.train()\n# running_loss = 0.0\nfor i, (subvolumes, inklabels) in tqdm(enumerate(train_loader), total=TRAINING_STEPS):\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#     running_loss += loss.item()\n#     if i % 3000 == 3000-1:\n#         print(\"Loss:\", running_loss / 3000)\n#         running_loss = 0.0","metadata":{"_uuid":"29fb8028-1ec5-42fb-b61a-658e1181c951","_cell_guid":"790108d2-5260-4276-9105-9da5a27c28f5","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:11:19.959148Z","iopub.execute_input":"2023-04-01T16:11:19.959549Z","iopub.status.idle":"2023-04-01T16:20:38.965609Z","shell.execute_reply.started":"2023-04-01T16:11:19.959512Z","shell.execute_reply":"2023-04-01T16:20:38.964476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Finally, we'll generate a prediction image. We'll use the model to predict the presence of ink for each pixel in our rectangle (the val set). Conceptually it looks like this:\n\n<a href=\"https://user-images.githubusercontent.com/22727759/224853653-7cffd0a4-c6fa-49a2-93c1-e3c820863a51.mp4\"><img src=\"https://user-images.githubusercontent.com/22727759/224853379-09ae991e-02be-4ecc-a652-313165b3005c.gif\"/></a>\n\n\nThis should take about a minute.\n\nRemember that the model has never seen the label data within the rectangle before!\n\nWe'll plot it side-by-side with the label image. Are you able to recognize the letter \"P\" in it?","metadata":{"_uuid":"8edfa121-b2ef-419e-acbe-6d430fb50133","_cell_guid":"3c6ad763-f47c-4aaa-8c12-cb95c2d28d74","trusted":true}},{"cell_type":"code","source":"eval_dataset = SubvolumeDataset(image_stack, label, pixels_inside_rect)\neval_loader = data.DataLoader(eval_dataset, batch_size=BATCH_SIZE, shuffle=False)\noutput = torch.zeros_like(label).float()\nmodel.eval()\nwith torch.no_grad():\n    for i, (subvolumes, _) in enumerate(tqdm(eval_loader)):\n        for j, value in enumerate(model(subvolumes.to(DEVICE))):\n            output[pixels_inside_rect[i*BATCH_SIZE+j]] = value\n\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.imshow(output.cpu(), cmap='gray')\nax2.imshow(label.cpu(), cmap='gray')\nplt.show()","metadata":{"_uuid":"68013338-2e52-4f0b-b15b-b18e071aa5da","_cell_guid":"56812c41-904d-4645-bddf-49b19fe2685d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:20:38.967340Z","iopub.execute_input":"2023-04-01T16:20:38.968000Z","iopub.status.idle":"2023-04-01T16:22:06.646062Z","shell.execute_reply.started":"2023-04-01T16:20:38.967959Z","shell.execute_reply":"2023-04-01T16:22:06.645088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Since our output has to be binary, we have to choose a threshold, say 40% confidence.","metadata":{"_uuid":"15c0a510-b4d1-4e14-974e-cbe5a7ac6b8e","_cell_guid":"49b12c62-7143-4d27-be79-e20e8cd9f5fe","trusted":true}},{"cell_type":"code","source":"THRESHOLD = 0.4\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.imshow(output.gt(THRESHOLD).cpu(), cmap='gray')\nax2.imshow(label.cpu(), cmap='gray')\nplt.show()","metadata":{"_uuid":"3eb989fa-2823-49f7-a55f-1e178884e344","_cell_guid":"4f1a3ed4-43c8-4a9c-9f49-dab4d3047048","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:22:06.647399Z","iopub.execute_input":"2023-04-01T16:22:06.648547Z","iopub.status.idle":"2023-04-01T16:22:09.431470Z","shell.execute_reply.started":"2023-04-01T16:22:06.648507Z","shell.execute_reply":"2023-04-01T16:22:09.430318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Finally, Kaggle expects a runlength-encoded submission.csv file, so let's output that.","metadata":{"_uuid":"04dc6e9a-5178-4ffa-95b3-a64783d4cf1a","_cell_guid":"df73edb8-b09d-4e84-b33a-b384dbe486fe","trusted":true}},{"cell_type":"code","source":"def rle(output):\n    flat_img = np.where(output.flatten().cpu() > THRESHOLD, 1, 0).astype(np.uint8)\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), ())))\nrle_output = rle(output)\n# This doesn't make too much sense, but let's just output in the required format\n# so notebook works as a submission. :-)\nprint(\"Id,Predicted\\na,\" + rle_output + \"\\nb,\" + rle_output, file=open('submission.csv', 'w'))","metadata":{"_uuid":"512cd6ab-2794-4ad3-87cc-fc240561f286","_cell_guid":"79b4b495-2e93-49ba-b9d8-434dccf49907","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-01T16:22:09.433203Z","iopub.execute_input":"2023-04-01T16:22:09.433588Z","iopub.status.idle":"2023-04-01T16:22:10.354779Z","shell.execute_reply.started":"2023-04-01T16:22:09.433547Z","shell.execute_reply":"2023-04-01T16:22:10.353745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Hurray! We've detected ink! Now, can you do better? :-) For example, you could start with this [example submission](https://www.kaggle.com/code/danielhavir/vesuvius-challenge-example-submission).","metadata":{"_uuid":"e84d0aa9-a297-4a90-b4f2-a8afb7c389c5","_cell_guid":"799ddbd9-6862-4a43-9ae4-3c2d2f01da85","trusted":true}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}