{"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":"# Import Libraries and set paths","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.utils.data as data\nimport numpy as np\nimport glob\nimport PIL.Image as Image\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom tqdm import tqdm\nfrom io import StringIO\nfrom sklearn.metrics import fbeta_score\nimport torchvision\nfrom skimage.util import view_as_windows\nfrom scipy.ndimage import distance_transform_edt\n\n# Constants\nPREFIX = '/kaggle/input/vesuvius-challenge-ink-detection/train/3/'\nZ_START = 28\nZ_DIM = 7\nXY_BORDER = 15 # depends on model\nXY_WINDOW = 112\nLARGE_XY_WINDOW = 224\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Load mask image\nmask = np.array(Image.open(PREFIX+\"mask.png\").convert('1'))\nnot_border = np.zeros(mask.shape, dtype=bool)\nnot_border[XY_BORDER:-XY_BORDER, XY_BORDER:-XY_BORDER] = True\nmask *= not_border\nmaskT = torch.from_numpy(mask).to(DEVICE)\n\n# Load label image\ninklabels = (np.array(Image.open(PREFIX+\"inklabels.png\")) > 0).astype(np.float32)\n# Soften labels so that model does not depend on boundary artifacts between ink and non ink\nlabels = np.mean(view_as_windows(np.pad(inklabels,5),11), axis = (2,3))\nlabel = torch.from_numpy(labels).to(DEVICE)\n# Build matrix of weights for fitness function based on square distance to nearest ink\ndistance_to_ink = distance_transform_edt(inklabels == 0.)\nfitness_weight = torch.from_numpy(2. / (1. + distance_to_ink**2) - 1.).float().to(DEVICE)\n\n# 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])]\n\n# 3d CNNs apply same weights to different layers, so they expect homogeneous distribution of input data\n#normed = [(image - image.mean()) / image.std() for image in tqdm(images, total=Z_DIM)]\nimage_stack = torch.stack([torch.from_numpy(image) for image in images], dim=0).to(DEVICE)\n\nprint(image_stack.shape, label.shape, maskT.shape)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:16.403079Z","iopub.execute_input":"2023-05-15T18:21:16.403469Z","iopub.status.idle":"2023-05-15T18:21:52.458065Z","shell.execute_reply.started":"2023-05-15T18:21:16.403436Z","shell.execute_reply":"2023-05-15T18:21:52.456968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fitness_weight[2703:2803, 1583:1683]","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:52.460172Z","iopub.execute_input":"2023-05-15T18:21:52.461222Z","iopub.status.idle":"2023-05-15T18:21:52.536208Z","shell.execute_reply.started":"2023-05-15T18:21:52.461187Z","shell.execute_reply":"2023-05-15T18:21:52.535234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_boundary = np.array([6, 11, 9, 5])\nR = eval_boundary * LARGE_XY_WINDOW\nR[:2] += XY_BORDER\n\ndef in_R(y, x):\n    return x >= R[0] and y >= R[1] and x < R[0]+R[2] and y < R[1]+R[3]\n\ndef out_R(y, x):\n    return not in_R(y, x)\n\ndef plot_batch_grid():\n    fig = plt.figure(figsize=(label.shape[1] // 500,label.shape[0] // 500))\n    ax = plt.gca()\n\n    plt.imshow(fitness_weight.cpu(),cmap=\"gray\")\n    plt.xticks(np.arange(XY_BORDER,label.shape[1],LARGE_XY_WINDOW), rotation=60)\n    plt.yticks(np.arange(XY_BORDER,label.shape[0],LARGE_XY_WINDOW))\n\n    rect = plt.Rectangle((R[0], R[1]), R[2], R[3], linewidth=2, edgecolor=\"r\", facecolor=\"None\")\n    ax.add_patch(rect)\n\n    ax.grid(color='g')\n    plt.show()\n    \nplot_batch_grid()","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:52.541710Z","iopub.execute_input":"2023-05-15T18:21:52.542277Z","iopub.status.idle":"2023-05-15T18:21:55.418492Z","shell.execute_reply.started":"2023-05-15T18:21:52.542235Z","shell.execute_reply":"2023-05-15T18:21:55.416857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SubvolumeDataset(data.Dataset):\n    def __init__(self, image_stack, fitness_weight, mask, window = XY_WINDOW, border = XY_BORDER, training = False, constraint = None):\n        self.image_stack = image_stack\n        self.fitness_weight = fitness_weight\n        self.mask = mask\n        self.window = window\n        self.border = border\n        self.Y_SIZE = (self.fitness_weight.shape[0]-2*self.border) // self.window\n        self.X_SIZE = (self.fitness_weight.shape[1]-2*self.border) // self.window\n        # at least one valid point\n        valid = nn.functional.max_pool2d((self.mask[self.border:-self.border,self.border:-self.border]).unsqueeze(0).float(), self.window) > 0\n        if training:\n            #has_data = nn.functional.max_pool2d((self.fitness_weight[self.border:-self.border,self.border:-self.border]).unsqueeze(0).float(), self.window) > 0.9\n            # all points are valid\n            valid = nn.functional.max_pool2d((~self.mask[self.border:-self.border,self.border:-self.border]).unsqueeze(0).float(), self.window) == 0\n            #valid = valid * has_data\n        self.valid_indices = (valid).nonzero()[:,1:]\n        if constraint is not None:\n            self.valid_indices = list(filter(lambda yx: constraint(yx[0]*self.window + self.border, yx[1]*self.window + self.border), self.valid_indices))\n\n    def __len__(self):\n        return len(self.valid_indices)\n\n    def __getitem__(self, index):\n        y, x = self.valid_indices[index]\n        y, x = y*self.window + self.border, x*self.window + self.border\n        subvolume = self.image_stack[:, y-self.border:y+self.window+self.border, x-self.border:x+self.window+self.border]\n        weight = self.fitness_weight[y:y+self.window, x:x+self.window]\n        truth = max(1, (weight > 0).sum().item())\n        #validmask = self.mask[y:y+self.window, x:x+self.window]\n        return subvolume.unsqueeze(0), weight, truth, y, x\n    \nclass InkDetectionModel(nn.Module):\n    def __init__(self):\n        super(InkDetectionModel, self).__init__()\n        # Since GA does not care about back propagation, we can use activation functions that can better solve the problem\n        # Rather then worrying about vanishing gradients, dying neurons, covariance shifts, etc.\n        # Interesting options are square, or more generally polinomial activations, squared ReLU, tanh and many more\n        self.conv1 = nn.Conv3d(1, 16, kernel_size=[3,3,3])\n        self.act1 = nn.Tanh()\n        self.conv2 = nn.Conv3d(16, 32, kernel_size=[3,3,3], dilation=[1,2,2])\n        self.act2 = nn.Tanh()\n        self.conv3 = nn.Conv3d(32, 64, kernel_size=[3,3,3], dilation=[1,4,4])\n        self.act3 = nn.Tanh()\n        # Same as having Fully connected layer, but allows to make simultaneous predictions on window of any size,\n        # Which is especially useful during evaluation\n        self.conv4 = nn.Conv3d(64, 128, kernel_size=[1,3,3], dilation=[1,8,8])\n        self.act4 = nn.Tanh()\n        self.fc2 = nn.Conv3d(128, 1, kernel_size=1)\n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self, x):\n        b,c,d,w,h = x.shape\n        x = self.act1(self.conv1(x))\n        x = self.act2(self.conv2(x))\n        x = self.act3(self.conv3(x))\n        x = self.act4(self.conv4(x))\n        x = self.sigmoid(self.fc2(x))\n        return x.squeeze(1,2)\n\n# Instantiate the model\nmodel = InkDetectionModel().to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:55.419588Z","iopub.execute_input":"2023-05-15T18:21:55.420020Z","iopub.status.idle":"2023-05-15T18:21:55.450994Z","shell.execute_reply.started":"2023-05-15T18:21:55.419983Z","shell.execute_reply":"2023-05-15T18:21:55.450355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = SubvolumeDataset(image_stack, fitness_weight, maskT, training = True, constraint = out_R)\ntrain_loader = data.DataLoader(train_dataset, batch_size=8, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:55.452166Z","iopub.execute_input":"2023-05-15T18:21:55.453105Z","iopub.status.idle":"2023-05-15T18:21:55.749951Z","shell.execute_reply.started":"2023-05-15T18:21:55.453072Z","shell.execute_reply":"2023-05-15T18:21:55.749089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"!pip install pygad>=2.10.0","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:21:55.751642Z","iopub.execute_input":"2023-05-15T18:21:55.752217Z","iopub.status.idle":"2023-05-15T18:22:08.697934Z","shell.execute_reply.started":"2023-05-15T18:21:55.752184Z","shell.execute_reply":"2023-05-15T18:22:08.696668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pygad.torchga\n\nnum_solutions = 10\n\ntorch_ga = pygad.torchga.TorchGA(model=model, num_solutions=num_solutions)\n\ndata_outputs = labels * mask\n\nall_samples = list(train_loader)\nscore_cache = {}\n\nprint(len(train_loader))\n\ndef fitness_func(ga_instance, solution, sol_idx):\n    index = ga_instance.generations_completed % len(train_loader)\n    key = f'{index}_{sol_idx}'\n    if key not in score_cache:\n        model_weights_dict = pygad.torchga.model_weights_as_dict(model=model, weights_vector=solution)\n        model.load_state_dict(model_weights_dict)\n        subvolumes, weights, truth, y, x = all_samples[index]\n        predictions = model(subvolumes.to(DEVICE))\n        score_cache[key] = (torch.dot(predictions.ravel(), weights.ravel()) / truth.to(DEVICE).sum()).item()\n    #print(key, score_cache[key])\n    return score_cache[key]","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:22:08.700386Z","iopub.execute_input":"2023-05-15T18:22:08.700799Z","iopub.status.idle":"2023-05-15T18:22:09.468633Z","shell.execute_reply.started":"2023-05-15T18:22:08.700759Z","shell.execute_reply":"2023-05-15T18:22:09.467492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 10\nwith tqdm(total=EPOCHS * len(all_samples)) as pbar:\n    def callback_generation(ga_instance):\n        pbar.set_postfix({\"Fitness\": ga_instance.best_solution()[1]})\n        pbar.update(1)\n\n    ga_instance = pygad.GA(num_generations=EPOCHS * len(all_samples),\n                           num_parents_mating=5,\n                           initial_population=torch_ga.population_weights,\n                           fitness_func=fitness_func,\n                           mutation_probability = 0.05,\n                           crossover_probability = 0.25,\n                           random_seed=17,\n                           on_generation=callback_generation)\n    #with torch.no_grad():\n    #model.eval()\n    ga_instance.run()\n    ga_instance.plot_result(title=\"PyGAD & PyTorch - Iteration vs. Fitness\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:38:46.916136Z","iopub.execute_input":"2023-05-15T18:38:46.916493Z","iopub.status.idle":"2023-05-15T18:39:37.725287Z","shell.execute_reply.started":"2023-05-15T18:38:46.916465Z","shell.execute_reply":"2023-05-15T18:39:37.723870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution, solution_fitness, solution_idx = ga_instance.best_solution()\nprint(\"Fitness value of the best solution = {solution_fitness}\".format(solution_fitness=solution_fitness))\nprint(\"Index of the best solution : {solution_idx}\".format(solution_idx=solution_idx))","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:32:34.211044Z","iopub.execute_input":"2023-05-15T18:32:34.211464Z","iopub.status.idle":"2023-05-15T18:32:34.612838Z","shell.execute_reply.started":"2023-05-15T18:32:34.211435Z","shell.execute_reply":"2023-05-15T18:32:34.611471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluate","metadata":{}},{"cell_type":"code","source":"# Create the evaluation dataset and data loader\neval_dataset = SubvolumeDataset(image_stack, label, maskT, window = LARGE_XY_WINDOW)\neval_loader = data.DataLoader(eval_dataset, batch_size=1, shuffle=False)\n\n# Initialize an output tensor to store predictions\noutput = torch.zeros_like(label).float()\n\n# Evaluation loop\nmodel.eval()\nwith torch.no_grad():\n    for subvolumes, _, _, y, x in tqdm(eval_loader, mininterval=2):\n        #if x >= R[0] and y >= R[1] and x < R[0]+R[2] and y < R[1]+R[3]:\n        prediction = pygad.torchga.predict(model=model, solution=solution, data=subvolumes.to(DEVICE))\n        output[y:y+LARGE_XY_WINDOW, x:x+LARGE_XY_WINDOW] = prediction\noutput *= maskT\n\ntorch.save(model.state_dict(), \"/kaggle/working/model.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:32:36.990360Z","iopub.execute_input":"2023-05-15T18:32:36.990710Z","iopub.status.idle":"2023-05-15T18:32:43.798888Z","shell.execute_reply.started":"2023-05-15T18:32:36.990681Z","shell.execute_reply":"2023-05-15T18:32:43.797816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize the prediction results\nfig, ax = plt.subplots(1, 2, figsize=(14,8))\nax[0].imshow(label.cpu(), cmap='gray')\nax[1].imshow(output.cpu(), cmap='gray', vmin = 0, vmax = 1)\nfor i in range(2):\n    rect = plt.Rectangle((R[0], R[1]), R[2], R[3], linewidth=2, edgecolor=\"r\", facecolor=\"None\")\n    ax[i].add_patch(rect)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:32:46.203843Z","iopub.execute_input":"2023-05-15T18:32:46.204568Z","iopub.status.idle":"2023-05-15T18:32:51.170702Z","shell.execute_reply.started":"2023-05-15T18:32:46.204537Z","shell.execute_reply":"2023-05-15T18:32:51.169844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target = torch.masked_select(label[R[1]:R[1]+R[3],R[0]:R[0]+R[2]],maskT[R[1]:R[1]+R[3],R[0]:R[0]+R[2]]).cpu()\nsource = torch.masked_select(output[R[1]:R[1]+R[3],R[0]:R[0]+R[2]],maskT[R[1]:R[1]+R[3],R[0]:R[0]+R[2]]).cpu()\nscores = np.array([\n    fbeta_score(target > 0.5, source >= th / 20., beta=0.5)\n    for th in tqdm(range(1,20), total=19)\n])\nprint(scores)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:33:05.454855Z","iopub.execute_input":"2023-05-15T18:33:05.455726Z","iopub.status.idle":"2023-05-15T18:33:16.085840Z","shell.execute_reply.started":"2023-05-15T18:33:05.455690Z","shell.execute_reply":"2023-05-15T18:33:16.084740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visTensor(tensor, ch=0, allkernels=True, padding=1): \n    n,c,d,w,h = tensor.shape\n    nrow = 8 * d\n\n    if allkernels: tensor = tensor.view(n*c*d, -1, w, h)\n    elif c != 3: tensor = tensor[:,ch,:,:].unsqueeze(dim=1)\n\n    rows = np.min((tensor.shape[0] // nrow + 1, 32))    \n    grid = torchvision.utils.make_grid(tensor, nrow=nrow, normalize=True, padding=padding)\n    plt.figure( figsize=(nrow,rows) )\n    plt.imshow(grid.cpu().numpy().transpose((1, 2, 0)))\n    plt.axis('off')\n    plt.ioff()\n    plt.show()\n\nimage0, label0, _, _, _ = eval_dataset[101]\nprint(image0.shape, label0.nonzero().shape)\nvisTensor(image0.unsqueeze(2), allkernels=True)\nfig, ax = plt.subplots(1, 1, figsize=(2,2))\nax.imshow(label0.cpu(), cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:35:02.654706Z","iopub.execute_input":"2023-05-15T18:35:02.655736Z","iopub.status.idle":"2023-05-15T18:35:03.003682Z","shell.execute_reply.started":"2023-05-15T18:35:02.655693Z","shell.execute_reply":"2023-05-15T18:35:03.002772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visTensor(model.conv1.weight.data.clone(), allkernels=True)\nimage1 = model.act1(model.conv1(image0))\nvisTensor(image1.unsqueeze(2), allkernels=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:35:15.356971Z","iopub.execute_input":"2023-05-15T18:35:15.357337Z","iopub.status.idle":"2023-05-15T18:35:16.607316Z","shell.execute_reply.started":"2023-05-15T18:35:15.357309Z","shell.execute_reply":"2023-05-15T18:35:16.606504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visTensor(model.conv2.weight.data.clone(), allkernels=True)\nimage2 = model.act2(model.conv2(image1))\nvisTensor(image2.unsqueeze(2), allkernels=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:35:46.637791Z","iopub.execute_input":"2023-05-15T18:35:46.638182Z","iopub.status.idle":"2023-05-15T18:35:47.894462Z","shell.execute_reply.started":"2023-05-15T18:35:46.638151Z","shell.execute_reply":"2023-05-15T18:35:47.893484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visTensor(model.conv3.weight.data.clone(), allkernels=True)\nimage3 = model.act3(model.conv3(image2))\nvisTensor(image3.unsqueeze(2), allkernels=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:36:19.337176Z","iopub.execute_input":"2023-05-15T18:36:19.337741Z","iopub.status.idle":"2023-05-15T18:36:20.134774Z","shell.execute_reply.started":"2023-05-15T18:36:19.337708Z","shell.execute_reply":"2023-05-15T18:36:20.133902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visTensor(model.conv4.weight.data.clone(), allkernels=True)\nimage4 = model.act4(model.conv4(image3))\nvisTensor(image4.unsqueeze(2), allkernels=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:36:45.959357Z","iopub.execute_input":"2023-05-15T18:36:45.959735Z","iopub.status.idle":"2023-05-15T18:36:47.466433Z","shell.execute_reply.started":"2023-05-15T18:36:45.959707Z","shell.execute_reply":"2023-05-15T18:36:47.462193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visTensor(model.fc2.weight.data.clone(), allkernels=True)\nimage5 = model.sigmoid(model.fc2(image4))\nvisTensor(image5.unsqueeze(2), allkernels=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:38:16.966321Z","iopub.execute_input":"2023-05-15T18:38:16.966676Z","iopub.status.idle":"2023-05-15T18:38:17.025177Z","shell.execute_reply.started":"2023-05-15T18:38:16.966651Z","shell.execute_reply":"2023-05-15T18:38:17.024093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Apply a threshold to the predictions to obtain binary output\nTHRESHOLD = (np.argmax(scores)+1) * 0.05  # Adjust the threshold value for better results\nprint(THRESHOLD)\nbinary_output = output.gt(THRESHOLD).cpu()\n\n# Visualize the binary prediction results\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.imshow(binary_output, cmap='gray')\nax2.imshow(label.cpu(), cmap='gray')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:38:23.712365Z","iopub.execute_input":"2023-05-15T18:38:23.712726Z","iopub.status.idle":"2023-05-15T18:38:26.412922Z","shell.execute_reply.started":"2023-05-15T18:38:23.712698Z","shell.execute_reply":"2023-05-15T18:38:26.411843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to generate run-length encoding (RLE) for the binary mask\ndef rle(img):\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    f = StringIO()\n    np.savetxt(f, runs.reshape(1, -1), delimiter=\" \", fmt=\"%d\")\n    predicted = f.getvalue().strip()\n    return predicted\n\n# Generate RLE for the binary output\nrle_output = rle(binary_output)\n\n# Save the RLE to a CSV file for submission\nwith open('submission.csv', 'w') as f:\n    f.write(\"Id,Predicted\\n\")\n    f.write(\"a,\" + rle_output + \"\\n\")\n    f.write(\"b,\" + rle_output + \"\\n\")\n\nprint(\"Submission file 'submission.csv' has been generated.\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T18:31:09.985340Z","iopub.status.idle":"2023-05-15T18:31:09.986035Z","shell.execute_reply.started":"2023-05-15T18:31:09.985774Z","shell.execute_reply":"2023-05-15T18:31:09.985797Z"},"trusted":true},"execution_count":null,"outputs":[]}]}