{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":47317,"databundleVersionId":5799376,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\nimport os\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:48:39.942659Z","iopub.execute_input":"2026-01-06T05:48:39.943221Z","iopub.status.idle":"2026-01-06T05:48:43.462577Z","shell.execute_reply.started":"2026-01-06T05:48:39.943194Z","shell.execute_reply":"2026-01-06T05:48:43.462026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PREFIX = '/kaggle/input/vesuvius-challenge-ink-detection/train/1'\nBUFFER = 30\nZ_START = 27\nZ_DIM = 10\nTRAINING_STEPS = 30000\nLEARNING_RATE = 0.03\nBATCH_SIZE = 32\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nVAL_EVERY = 3000","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:48:45.648589Z","iopub.execute_input":"2026-01-06T05:48:45.648996Z","iopub.status.idle":"2026-01-06T05:48:45.708045Z","shell.execute_reply.started":"2026-01-06T05:48:45.648969Z","shell.execute_reply":"2026-01-06T05:48:45.707338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(Image.open(PREFIX + \"/ir.png\"), cmap=\"gray\")\nplt.show()\nmask = np.array(Image.open(PREFIX + \"/mask.png\").convert('1'))\nlabel = np.array(Image.open(PREFIX + \"/inklabels.png\"))\nir = np.array(Image.open(PREFIX + \"/ir.png\"))\n\ntif_files = sorted(glob.glob(os.path.join(PREFIX, 'surface_volume', '*.tif')))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:48:59.301432Z","iopub.execute_input":"2026-01-06T05:48:59.301760Z","iopub.status.idle":"2026-01-06T05:49:08.453303Z","shell.execute_reply.started":"2026-01-06T05:48:59.301729Z","shell.execute_reply":"2026-01-06T05:49:08.452704Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images = [\n    np.array(Image.open(f), dtype=np.float32) / 65535.0\n    for f in tif_files[Z_START:Z_START + Z_DIM]\n]\n\nimage_stack = torch.stack([torch.from_numpy(img) for img in images], dim=0).to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:49:17.332068Z","iopub.execute_input":"2026-01-06T05:49:17.332357Z","iopub.status.idle":"2026-01-06T05:49:35.359028Z","shell.execute_reply.started":"2026-01-06T05:49:17.332332Z","shell.execute_reply":"2026-01-06T05:49:35.358410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label = torch.from_numpy(label).float()\n\nfig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(15, 5))\nax1.set_title(\"mask.png\")\nax1.imshow(mask, cmap='gray')\nax2.set_title(\"inklabels.png\")\nax2.imshow(label.cpu().numpy(), cmap='gray')\nax3.set_title(\"ir.png\")\nax3.imshow(ir, cmap='gray')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:49:39.470016Z","iopub.execute_input":"2026-01-06T05:49:39.470311Z","iopub.status.idle":"2026-01-06T05:50:02.319099Z","shell.execute_reply.started":"2026-01-06T05:49:39.470286Z","shell.execute_reply":"2026-01-06T05:50:02.318286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(1, len(images), figsize=(15, 3))\nfor i, (image, ax) in enumerate(zip(images, axes)):\n    img_uint8 = (image * 255).astype(np.uint8)\n    pil_img = Image.fromarray(img_uint8)\n    small_img = pil_img.resize((image.shape[1] // 16, image.shape[0] // 16))\n    ax.imshow(small_img, cmap='gray')\n    ax.set_title(f\"Layer {Z_START + i}\")\n    ax.axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:50:22.731920Z","iopub.execute_input":"2026-01-06T05:50:22.732616Z","iopub.status.idle":"2026-01-06T05:50:25.857581Z","shell.execute_reply.started":"2026-01-06T05:50:22.732588Z","shell.execute_reply":"2026-01-06T05:50:25.856753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rect = (1100, 3500, 700, 950)\n\nfig, ax = plt.subplots(figsize=(10, 8))\nax.imshow(label.cpu().numpy(), cmap='gray')\npatch = patches.Rectangle(\n    (rect[0], rect[1]), rect[2], rect[3],\n    linewidth=2, edgecolor='r', facecolor='none'\n)\nax.add_patch(patch)\nax.set_title(\"inklabels.png with selected rect\")\nax.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:50:37.730334Z","iopub.execute_input":"2026-01-06T05:50:37.730645Z","iopub.status.idle":"2026-01-06T05:50:45.868253Z","shell.execute_reply.started":"2026-01-06T05:50:37.730618Z","shell.execute_reply":"2026-01-06T05:50:45.867625Z"}},"outputs":[],"execution_count":null},{"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\n    def __len__(self):\n        return len(self.pixels)\n\n    def __getitem__(self, index):\n        y, x = self.pixels[index]\n\n        subvolume = self.image_stack[\n            :, y-BUFFER:y+BUFFER+1, x-BUFFER:x+BUFFER+1\n        ].view(1, Z_DIM, BUFFER*2+1, BUFFER*2+1)\n\n        inklabel = torch.tensor(\n            self.label[y, x],\n            dtype=torch.float32\n        ).view(1)\n\n        return subvolume, inklabel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:50:55.081298Z","iopub.execute_input":"2026-01-06T05:50:55.081597Z","iopub.status.idle":"2026-01-06T05:50:55.087085Z","shell.execute_reply.started":"2026-01-06T05:50:55.081571Z","shell.execute_reply":"2026-01-06T05:50:55.086480Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"not_border = np.zeros(mask.shape, dtype=bool)\nnot_border[BUFFER:-BUFFER, BUFFER:-BUFFER] = True\nvalid_area = mask & not_border\n\ninside_rect = np.zeros(mask.shape, dtype=bool)\ninside_rect[rect[1]:rect[1]+rect[3], rect[0]:rect[0]+rect[2]] = True\ninside_rect &= valid_area\n\noutside_rect = valid_area & (~inside_rect)\n\npixels_train = np.argwhere(outside_rect)\npixels_val = np.argwhere(inside_rect)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:51:12.187865Z","iopub.execute_input":"2026-01-06T05:51:12.188157Z","iopub.status.idle":"2026-01-06T05:51:12.937292Z","shell.execute_reply.started":"2026-01-06T05:51:12.188131Z","shell.execute_reply":"2026-01-06T05:51:12.936703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = data.DataLoader(\n    SubvolumeDataset(image_stack, label, pixels_train),\n    batch_size=BATCH_SIZE, shuffle=True\n)\n\nval_loader = data.DataLoader(\n    SubvolumeDataset(image_stack, label, pixels_val),\n    batch_size=BATCH_SIZE, shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:51:26.234966Z","iopub.execute_input":"2026-01-06T05:51:26.235554Z","iopub.status.idle":"2026-01-06T05:51:26.239600Z","shell.execute_reply.started":"2026-01-06T05:51:26.235526Z","shell.execute_reply":"2026-01-06T05:51:26.238835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = nn.Sequential(\n    nn.Conv3d(1, 16, 3, 1, 1), nn.MaxPool3d(2),\n    nn.Conv3d(16, 32, 3, 1, 1), nn.MaxPool3d(2),\n    nn.Conv3d(32, 64, 3, 1, 1), nn.MaxPool3d(2),\n    nn.Flatten(1),\n    nn.LazyLinear(128), nn.ReLU(),\n    nn.LazyLinear(1)\n).to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer, max_lr=LEARNING_RATE*10, total_steps=TRAINING_STEPS\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:51:37.959830Z","iopub.execute_input":"2026-01-06T05:51:37.960367Z","iopub.status.idle":"2026-01-06T05:51:40.585031Z","shell.execute_reply.started":"2026-01-06T05:51:37.960340Z","shell.execute_reply":"2026-01-06T05:51:40.584477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Training started...\")\n# =====================\ntrain_losses = []\nval_losses = []\n\nmodel.train()\nrunning_loss = 0.0\n\nfor step, (x, y) in enumerate(tqdm(train_loader, total=TRAINING_STEPS)):\n    if step >= TRAINING_STEPS:\n        break\n\n    x, y = x.to(DEVICE), y.to(DEVICE)\n\n    optimizer.zero_grad()\n    logits = model(x)\n    loss = criterion(logits, y)\n    loss.backward()\n    optimizer.step()\n    scheduler.step()\n\n    running_loss += loss.item()\n\n    if (step + 1) % VAL_EVERY == 0:\n        train_loss = running_loss / VAL_EVERY\n        train_losses.append(train_loss)\n        running_loss = 0.0\n\n        # Validation\n        model.eval()\n        val_loss = 0.0\n        with torch.no_grad():\n            for vx, vy in val_loader:\n                vx, vy = vx.to(DEVICE), vy.to(DEVICE)\n                v_logits = model(vx)\n                val_loss += criterion(v_logits, vy).item()\n\n        val_loss /= len(val_loader)\n        val_losses.append(val_loss)\n\n        print(f\"Step {step+1} | Train Loss: {train_loss:.5f} | Val Loss: {val_loss:.5f}\")\n        model.train()\nprint(\"Training completed!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-06T05:52:32.947491Z","iopub.execute_input":"2026-01-06T05:52:32.948335Z","iopub.status.idle":"2026-01-06T06:23:39.932494Z","shell.execute_reply.started":"2026-01-06T05:52:32.948306Z","shell.execute_reply":"2026-01-06T06:23:39.931739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}