{"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":"# Why implement weight maps?\nPixels of segmentation borders have greater weights in the original implementation of [U-Net](https://arxiv.org/pdf/1505.04597.pdf), in order to address the challenge of separating touching objects of the same class.\n> The separation border is computed using morphological operations. The\nweight map is then computed as\n>\n>$w(\\mathbf{x})=w_c(\\mathbf{x})+w_0 \\cdot \\exp \\left(-\\frac{\\left(d_1(\\mathbf{x})+d_2(\\mathbf{x})\\right)^2}{2 \\sigma^2}\\right)$ \n>\n>where wc : Ω → R is the weight map to balance the class frequencies, d1 : Ω → R\ndenotes the distance to the border of the nearest cell and d2 : Ω → R the distance\nto the border of the second nearest cell. In our experiments we set w0 = 10 and\nσ ≈ 5 pixels.\n>\n> (d) map with a pixel-wise loss weight to force the network to learn the border pixels.\n>\n> ![](https://i.postimg.cc/90RrVJXF/image-20230520174104573.png)\n\nA similar challenge also exists in our contrail identification task. As you can see, boundaries of contrails aren't localized very well.\n\n![](https://i.postimg.cc/ZYWSNWCv/QQ-20230615011445.png)\n\nThus, we can also assign greater weights on boundaries of contrails.\n\n![](https://i.postimg.cc/kMj7vkN0/QQ-20230615012341.png)","metadata":{}},{"cell_type":"markdown","source":"# Visualizing Weight Maps","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom IPython import display\nimport matplotlib.pyplot as plt\nfrom scipy.ndimage import distance_transform_edt","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:51.557371Z","iopub.execute_input":"2023-06-14T17:57:51.557688Z","iopub.status.idle":"2023-06-14T17:57:51.665521Z","shell.execute_reply.started":"2023-06-14T17:57:51.557666Z","shell.execute_reply":"2023-06-14T17:57:51.664634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ids(tar_path):\n    ids = []\n    for img_id in os.listdir(tar_path):\n        ids.append(img_id)\n    print(f\"{len(ids)} samples in {tar_path}\")\n    return ids\n\ntar_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train\"\nids = get_ids(tar_path)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:51.67048Z","iopub.execute_input":"2023-06-14T17:57:51.672508Z","iopub.status.idle":"2023-06-14T17:57:51.910267Z","shell.execute_reply.started":"2023-06-14T17:57:51.672475Z","shell.execute_reply":"2023-06-14T17:57:51.909529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_weight_map(image, sigma=0.2):\n    \"\"\"\n    Generates a weight map based on the ground truth.\n    \n    Args:\n        image (numpy.ndarray): Ground truth, binary image of shape (height, width).\n        sigma (float): Controls the range of boundaries.\n        \n    Returns:\n        weight_map (numpy.ndarray): Weight map of the same shape as the ground truth.\n    \"\"\"\n    distance = distance_transform_edt(1 - image)\n    distance = distance / np.max(distance)\n    weight_map = np.exp(-0.5 * (distance / sigma) ** 2)\n    weight_map[image == 1] = 0\n    return weight_map","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:51.911198Z","iopub.execute_input":"2023-06-14T17:57:51.911954Z","iopub.status.idle":"2023-06-14T17:57:51.916405Z","shell.execute_reply.started":"2023-06-14T17:57:51.911929Z","shell.execute_reply":"2023-06-14T17:57:51.915753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img_id in ids[:20]:  # visualize 20 images\n    sample_path = f\"{tar_path}/{img_id}\"\n    gt = np.load(f\"{sample_path}/human_pixel_masks.npy\")\n    if np.all(gt == 0):  # skip images without contrails\n        continue\n    weight_map = gen_weight_map(gt)  # generate the weight map\n    plt.figure(figsize=(6, 3))\n    ax = plt.subplot(1, 2, 1)\n    ax.imshow(gt, interpolation='none')\n    ax.set_title('GroundTruth')\n    ax = plt.subplot(1, 2, 2)\n    ax.imshow(weight_map, interpolation='none')\n    ax.set_title('WeightMap')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:51.919155Z","iopub.execute_input":"2023-06-14T17:57:51.919579Z","iopub.status.idle":"2023-06-14T17:57:54.749237Z","shell.execute_reply.started":"2023-06-14T17:57:51.919551Z","shell.execute_reply":"2023-06-14T17:57:54.748246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  Example for Loss Calculation","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import nn","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:54.750483Z","iopub.execute_input":"2023-06-14T17:57:54.751232Z","iopub.status.idle":"2023-06-14T17:57:57.355622Z","shell.execute_reply.started":"2023-06-14T17:57:54.751203Z","shell.execute_reply":"2023-06-14T17:57:57.354342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:57.356976Z","iopub.execute_input":"2023-06-14T17:57:57.357433Z","iopub.status.idle":"2023-06-14T17:57:57.360889Z","shell.execute_reply.started":"2023-06-14T17:57:57.35741Z","shell.execute_reply":"2023-06-14T17:57:57.360169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def bce_loss(pred, gt):\n    class_weight=torch.Tensor([73.2]).to(device)  # balance the class frequencies\n    criterion = nn.BCELoss(class_weight, reduction='none')\n    loss = criterion(pred, gt)\n    gt.detach().cpu()\n    weight_map = gen_weight_map(gt)\n    weight_map = torch.from_numpy(weight_map).to(device)\n    loss = loss + loss * weight_map\n    return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:57.36191Z","iopub.execute_input":"2023-06-14T17:57:57.362487Z","iopub.status.idle":"2023-06-14T17:57:57.374335Z","shell.execute_reply.started":"2023-06-14T17:57:57.36246Z","shell.execute_reply":"2023-06-14T17:57:57.37309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test the loss function\ngt = np.load(f\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/1000934780200790693/human_pixel_masks.npy\")\ngt = torch.from_numpy(gt).permute(2, 0, 1).reshape(1, 1, 256, 256)\nprint(gt.size())  # [batch_size, n_channels, height, width]\nloss = bce_loss(torch.randn(1, 1, 256, 256).clamp(0, 1), gt.float())\nprint(loss)","metadata":{"execution":{"iopub.status.busy":"2023-06-14T17:57:57.375664Z","iopub.execute_input":"2023-06-14T17:57:57.375956Z","iopub.status.idle":"2023-06-14T17:57:57.522159Z","shell.execute_reply.started":"2023-06-14T17:57:57.375932Z","shell.execute_reply":"2023-06-14T17:57:57.521162Z"},"trusted":true},"execution_count":null,"outputs":[]}]}