{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import warnings\nimport numpy as np\n\nwarnings.simplefilter('ignore')\n\nSEED = 333\nnp.random.seed(SEED)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:02.428837Z","iopub.execute_input":"2025-03-24T18:36:02.429116Z","iopub.status.idle":"2025-03-24T18:36:02.433335Z","shell.execute_reply.started":"2025-03-24T18:36:02.429094Z","shell.execute_reply":"2025-03-24T18:36:02.432460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\nvesuvius_path = '/kaggle/input/vesuvius-challenge-ink-detection/'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:02.434499Z","iopub.execute_input":"2025-03-24T18:36:02.434795Z","iopub.status.idle":"2025-03-24T18:36:05.627331Z","shell.execute_reply.started":"2025-03-24T18:36:02.434766Z","shell.execute_reply":"2025-03-24T18:36:05.626649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom torch import Tensor\nimport cv2\n\nTIF_STARTS = 27\nTIF_RANGE = 10\n\ndef load_image_stack(ab_123,\n                     train_rectangle=None,\n                     test_train='train'):\n\n  surface_volume_filepath = vesuvius_path + \\\n    f'{test_train}/{ab_123}/surface_volume/'\n  tif_filepaths = \\\n    [surface_volume_filepath + tif_filepath \\\n     for tif_filepath in sorted(os.listdir(surface_volume_filepath))[:-4]]\n  tif_filepaths_stack = tif_filepaths[TIF_STARTS:TIF_STARTS + TIF_RANGE]\n\n  image_stack = []\n  for tif_filepath_stack in tif_filepaths_stack:\n    loaded_img = Tensor(cv2.imread(tif_filepath_stack,\n                                         0) / 65535.0).float() \\\n      .to(device)\n    if train_rectangle is None:\n      image_stack.append(loaded_img)\n    else:\n      image_stack.append(loaded_img[train_rectangle[1]:train_rectangle[1] + \\\n                                    train_rectangle[3],\n                                    train_rectangle[0]:train_rectangle[0] + \\\n                                    train_rectangle[2]])\n\n  return torch.stack(image_stack, dim=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.628901Z","iopub.execute_input":"2025-03-24T18:36:05.629382Z","iopub.status.idle":"2025-03-24T18:36:05.936648Z","shell.execute_reply.started":"2025-03-24T18:36:05.629348Z","shell.execute_reply":"2025-03-24T18:36:05.935774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_train_mask_label(train_123, train_rectangle=None):\n\n  mask_filepath = vesuvius_path + f'train/{train_123}/mask.png'\n  label_filepath = vesuvius_path + f'train/{train_123}/inklabels.png'\n\n  mask = cv2.imread(mask_filepath, 0) / 255.\n  label = cv2.imread(label_filepath, 0) / 255.\n  mask = mask[train_rectangle[1]:train_rectangle[1] + train_rectangle[3],\n              train_rectangle[0]:train_rectangle[0] + train_rectangle[2]]\n  label = label[train_rectangle[1]:train_rectangle[1] + train_rectangle[3],\n                train_rectangle[0]:train_rectangle[0] + train_rectangle[2]]\n\n  return mask, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.937848Z","iopub.execute_input":"2025-03-24T18:36:05.938065Z","iopub.status.idle":"2025-03-24T18:36:05.942621Z","shell.execute_reply.started":"2025-03-24T18:36:05.938047Z","shell.execute_reply":"2025-03-24T18:36:05.941827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMAGE_SIZE = 64\nRANGE = IMAGE_SIZE // 2\nSTRIDE = 17\n\ndef get_non_zero_indices(mask):\n  trim_margins = np.zeros(mask.shape, dtype=bool)\n  trim_margins[RANGE:mask.shape[0] - RANGE,\n             RANGE:mask.shape[1] - RANGE] = True\n\n  trim_margins_mask = np.array(mask) * trim_margins\n  del trim_margins\n\n  sparse_mask = np.zeros(mask.shape, dtype=bool)\n  sparse_mask[::STRIDE, ::STRIDE] = True\n\n  return np.argwhere(sparse_mask * trim_margins_mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.943441Z","iopub.execute_input":"2025-03-24T18:36:05.943735Z","iopub.status.idle":"2025-03-24T18:36:05.958031Z","shell.execute_reply.started":"2025-03-24T18:36:05.943707Z","shell.execute_reply":"2025-03-24T18:36:05.957329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nfrom torch.utils.data import Dataset\n\nclass SubvolumeDataset(Dataset):\n\n  def __init__(self, image_stack, label, non_zero_indices):\n    self.image_stack = image_stack\n    self.label = Tensor(label).float()\n    self.non_zero_indices = non_zero_indices\n\n  def __len__(self):\n    return len(self.non_zero_indices)\n\n  def __getitem__(self, idx):\n    y, x = self.non_zero_indices[idx]\n    subvolume = self.image_stack[:,\n                                 y - RANGE:y + RANGE,\n                                 x - RANGE:x + RANGE]\n    ink_label = self.label[y, x]\n\n    return subvolume, ink_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.958742Z","iopub.execute_input":"2025-03-24T18:36:05.958958Z","iopub.status.idle":"2025-03-24T18:36:05.975809Z","shell.execute_reply.started":"2025-03-24T18:36:05.958940Z","shell.execute_reply":"2025-03-24T18:36:05.975187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.nn import Module, Sequential, Conv3d, ReLU, BatchNorm3d\n\nBATCH_NORM_MOMENTUM = 0.1\nFILTERS = [16, 32, 64, 128]\nFILTER_SIZES = [1] + FILTERS\nFILTER_PAIRS = list(zip(FILTER_SIZES[:-1], FILTER_SIZES[1:]))\nSTRIDES = [1, 2, 2, 2]\nKERNEL_SIZE = 3\nPADDING = 1\n\nclass Subvolume3DcnnEncoder(Module):\n\n  def __init__(self):\n\n    super().__init__()\n    self.conv_layers = Sequential(\n        *[Sequential(\n            Conv3d(chan_in,\n                   chan_out,\n                   kernel_size=KERNEL_SIZE,\n                   stride=stride, padding=PADDING),\n            ReLU(),\n            BatchNorm3d(num_features=filter_,\n                        momentum=BATCH_NORM_MOMENTUM)) \\\n          for (chan_in, chan_out), stride, filter_ in zip(FILTER_PAIRS,\n                                                          STRIDES,\n                                                          FILTERS)])\n    self.apply(self.init_weight)\n\n  @staticmethod\n  def init_weight(w):\n    if isinstance(w, Conv3d):\n      nn.init.xavier_uniform_(w.weight)\n      nn.init.zeros_(w.bias)\n\n  def forward(self, x):\n    return self.conv_layers(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.977785Z","iopub.execute_input":"2025-03-24T18:36:05.978010Z","iopub.status.idle":"2025-03-24T18:36:05.989676Z","shell.execute_reply.started":"2025-03-24T18:36:05.977991Z","shell.execute_reply":"2025-03-24T18:36:05.988906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.nn import Linear, Flatten, Sigmoid\n\nclass LinearInkDecoder(nn.Module):\n\n  def __init__(self, input_shape):\n\n    super().__init__()\n    self.linear = Linear(int(np.prod(input_shape)), 1)\n    self.flatten = Flatten()\n    self.sigmoid = Sigmoid()\n\n  def forward(self, x):\n    return self.sigmoid(self.linear(self.flatten(x)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:05.990913Z","iopub.execute_input":"2025-03-24T18:36:05.991191Z","iopub.status.idle":"2025-03-24T18:36:06.008098Z","shell.execute_reply.started":"2025-03-24T18:36:05.991163Z","shell.execute_reply":"2025-03-24T18:36:06.007354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SUBVOLUME_SHAPE = [IMAGE_SIZE, IMAGE_SIZE, 10]\n\nclass InkClassifier3DCNN(nn.Module):\n\n  def __init__(self):\n    super().__init__()\n    self.encoder = Subvolume3DcnnEncoder()\n    self.decoder = LinearInkDecoder(\n        self.encoder(torch.zeros((1,\n                                  1,\n                                  *SUBVOLUME_SHAPE))).shape[1:])\n  def forward(self, x):\n    return self.decoder(self.encoder(x))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:06.008842Z","iopub.execute_input":"2025-03-24T18:36:06.009089Z","iopub.status.idle":"2025-03-24T18:36:06.023373Z","shell.execute_reply.started":"2025-03-24T18:36:06.009062Z","shell.execute_reply":"2025-03-24T18:36:06.022607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.optim import Adam, lr_scheduler\n\nLR = 3e-3\n\nmodel = InkClassifier3DCNN().to(device)\ntorch_rand = torch.rand((5, 1, IMAGE_SIZE, IMAGE_SIZE, 10)).to(device)\n\nTOTAL_TRAINING_STEPS = 14000\n\nSUBVOLUME_TRAINING_STEPS = 4000\n\nloss_fn = nn.BCELoss()\noptimizer = Adam(model.parameters(), lr=LR)\nscheduler = lr_scheduler.OneCycleLR(optimizer,\n                                                max_lr=LR,\n                                                total_steps=TOTAL_TRAINING_STEPS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:06.024071Z","iopub.execute_input":"2025-03-24T18:36:06.024342Z","iopub.status.idle":"2025-03-24T18:36:08.254472Z","shell.execute_reply.started":"2025-03-24T18:36:06.024312Z","shell.execute_reply":"2025-03-24T18:36:08.253611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.train()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:08.255356Z","iopub.execute_input":"2025-03-24T18:36:08.255662Z","iopub.status.idle":"2025-03-24T18:36:08.261676Z","shell.execute_reply.started":"2025-03-24T18:36:08.255642Z","shell.execute_reply":"2025-03-24T18:36:08.260865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# v4\nTRAIN_RECTANGLES = [\n     [1, [500, 1500, 3700, 6400]],\n     [1, [2000, 400, 2500, 1000]],\n     [1, [2200, 400, 2300, 900]],\n     [2, [700, 300, 5000, 4300]],\n     [2, [500, 4800, 8900, 10000]],\n     [2, [1200, 4800, 5400, 9700]],\n     [2, [1400, 10600, 7700, 3700]],\n     [2, [600, 6100, 8700, 5700]],\n     [3, [1950, 800, 1950, 900]],\n     [3, [1450, 2350, 3750, 1300]],\n     [3, [100, 4150, 4350, 1650]]\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:08.262602Z","iopub.execute_input":"2025-03-24T18:36:08.262925Z","iopub.status.idle":"2025-03-24T18:36:08.280151Z","shell.execute_reply.started":"2025-03-24T18:36:08.262894Z","shell.execute_reply":"2025-03-24T18:36:08.279481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nfrom torch.utils.data import DataLoader\nimport gc\n\ntraining_step = 0\nBATCH_SIZE = 32\n\nwhile training_step < TOTAL_TRAINING_STEPS:\n\n  random.shuffle(TRAIN_RECTANGLES)\n  for train_123, train_rectangle in TRAIN_RECTANGLES:\n\n    if training_step >= TOTAL_TRAINING_STEPS:\n      break\n\n    mask, label = load_train_mask_label(train_123,\n                                        train_rectangle=train_rectangle)\n    image_stack = load_image_stack(train_123, train_rectangle=train_rectangle)\n    non_zero_indices = get_non_zero_indices(mask)\n    del mask\n\n    dataset = SubvolumeDataset(image_stack, label, non_zero_indices)\n    dataloader = DataLoader(dataset,\n                            batch_size=BATCH_SIZE,\n                            shuffle=True)\n\n    dataloader_loss = []\n    dataloader_step = 0\n\n    for subvolumes, ink_labels in dataloader:\n\n      if dataloader_step >= SUBVOLUME_TRAINING_STEPS \\\n        or training_step >= TOTAL_TRAINING_STEPS:\n        break\n\n      ink_labels = ink_labels.to(device)\n      logits = model(subvolumes.permute(0, 2, 3, 1).unsqueeze(dim=1))\n      loss = loss_fn(logits, ink_labels.unsqueeze(dim=1))\n      optimizer.zero_grad()\n      loss.backward()\n      optimizer.step()\n      scheduler.step()\n      dataloader_loss.append(loss.item())\n      training_step += 1\n      dataloader_step += 1\n\n    del image_stack, label, dataset, dataloader\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:36:08.280843Z","iopub.execute_input":"2025-03-24T18:36:08.281072Z","iopub.status.idle":"2025-03-24T18:48:23.089627Z","shell.execute_reply.started":"2025-03-24T18:36:08.281053Z","shell.execute_reply":"2025-03-24T18:48:23.088686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model, 'vc_train_s_333_tts_14000_sts_4000_v4.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T18:48:23.090518Z","iopub.execute_input":"2025-03-24T18:48:23.090785Z","iopub.status.idle":"2025-03-24T18:48:23.103152Z","shell.execute_reply.started":"2025-03-24T18:48:23.090757Z","shell.execute_reply":"2025-03-24T18:48:23.102407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}