{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\nIndividual masks are valuable to offer nuanced insights. For instance, if 4 out of 5 labelers detect a contrail on a pixel, the model could ideally assign it a value of 0.8, or 0.4 if only 2 of 5 labelers detect it.\n\nHowever, when I integrated the Dice loss in the training process, I noticed a degradation in performance. After analyzing the scores of ideal predictions for each case, I found a problem when using Dice loss with this kind of soft labeling.\n\nThis document is structured into three main parts:\n\n1. Visualization of adjusted dice score vs dice score on one example.\n2. The issue associated with dice score and soft labels.\n3. The proposed solution.\n4. The Adjusted Dice loss function, designed for PyTorch models.\n\n#### Update: \nThe reason Dice loss did not work well for me was that I left the threshold at 0.5 for the model outputs,\nthe ideal threshold when using soft labels in Dice loss training is around 0.999 for me. The problem described here still remains but it does not worsen performance in DL model training by as much as I thought when starting this avenue.\n\n#### Update 2:\nThe adjusted dice loss works well for model training!\n* Data: Train on full trainset and validate on given validation set, Ash Color Scheme, no upscaling, soft labels.\n* Model: Segmentation models pytorch, UNet decoder, ResNeSt26d encoder.\n* Regularization: No augmentation, only weight decay with AdamsW at 0.8.\n* Other settings: Epochs 10, Cosine scheduler, starting at 5e-4 learning rate.\n\nI get 0.629 for regular Dice Loss and 0.633 for Adjusted Dice Loss!\n\n#### Update 3:\nFixed the pytorch implementation. I had it correct at my computer, but here I missed adding \"torch.minimum(target, output)\".\n\n## Update 4:\nVerified that pytorch implementation works!","metadata":{}},{"cell_type":"markdown","source":"## Definitions","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.480773Z","iopub.execute_input":"2023-09-14T07:16:59.482095Z","iopub.status.idle":"2023-09-14T07:16:59.517769Z","shell.execute_reply.started":"2023-09-14T07:16:59.482036Z","shell.execute_reply":"2023-09-14T07:16:59.516433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize(values, dices, adjusted_dices):\n    plt.plot(values, dices, label='Orginal Dice')\n    plt.plot(values, adjusted_dices, label='Adjusted Dice')\n    plt.legend()\n    plt.xlabel('Value, right side of prediction')\n    plt.ylabel('Score')\n\n    # Add a title\n    plt.title('Adjusted vs Original Dice Score, perfect prediction at 0.2')\n\n    # Display the plot\n    plt.show()   ","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.520197Z","iopub.execute_input":"2023-09-14T07:16:59.520671Z","iopub.status.idle":"2023-09-14T07:16:59.527626Z","shell.execute_reply.started":"2023-09-14T07:16:59.52064Z","shell.execute_reply":"2023-09-14T07:16:59.526208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice(inputs, targets):\n    intersection = (inputs * targets).sum()                            \n    dice = 2.*intersection/(inputs.sum() + targets.sum())  \n    return dice, intersection, inputs.sum(), targets.sum()","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.529541Z","iopub.execute_input":"2023-09-14T07:16:59.530007Z","iopub.status.idle":"2023-09-14T07:16:59.539299Z","shell.execute_reply.started":"2023-09-14T07:16:59.529964Z","shell.execute_reply":"2023-09-14T07:16:59.538325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def adjusted_dice(inputs, targets):\n    # Conceptually this 'absolute' leads to the best intersection and score being achieved at target==input  \n    intersection = (np.minimum(inputs, targets) * (1-abs(targets - inputs))).sum()   \n    dice = 2.*intersection/(inputs.sum() + targets.sum())  \n    return dice, intersection, inputs.sum(), targets.sum()","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.541595Z","iopub.execute_input":"2023-09-14T07:16:59.54193Z","iopub.status.idle":"2023-09-14T07:16:59.551885Z","shell.execute_reply.started":"2023-09-14T07:16:59.541893Z","shell.execute_reply":"2023-09-14T07:16:59.550798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization","metadata":{}},{"cell_type":"code","source":"groundtruth = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\ngroundtruth","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.553518Z","iopub.execute_input":"2023-09-14T07:16:59.553912Z","iopub.status.idle":"2023-09-14T07:16:59.567935Z","shell.execute_reply.started":"2023-09-14T07:16:59.553872Z","shell.execute_reply":"2023-09-14T07:16:59.566883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dices = []\nadjusted_dices = [] \nvalues = [m/100 for m in range(100)]\nfor m in values:\n    prediction = np.array([[1,1,m,m],[1,1,m,m]])\n    dices.append(dice(prediction, groundtruth)[0])\n    adjusted_dices.append(adjusted_dice(prediction, groundtruth)[0])\nvisualize(values, dices, adjusted_dices)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:16:59.569411Z","iopub.execute_input":"2023-09-14T07:16:59.570292Z","iopub.status.idle":"2023-09-14T07:17:00.008175Z","shell.execute_reply.started":"2023-09-14T07:16:59.570251Z","shell.execute_reply":"2023-09-14T07:17:00.007229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The Problem with Dice and Soft Labels","metadata":{}},{"cell_type":"code","source":"groundtruth = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\ngroundtruth","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.009575Z","iopub.execute_input":"2023-09-14T07:17:00.010678Z","iopub.status.idle":"2023-09-14T07:17:00.020319Z","shell.execute_reply.started":"2023-09-14T07:17:00.010633Z","shell.execute_reply":"2023-09-14T07:17:00.018978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_perfect = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\nprediction_perfect","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.024827Z","iopub.execute_input":"2023-09-14T07:17:00.025258Z","iopub.status.idle":"2023-09-14T07:17:00.033456Z","shell.execute_reply.started":"2023-09-14T07:17:00.025226Z","shell.execute_reply":"2023-09-14T07:17:00.032478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_less_good = np.array([[1,1,0,0],[1,1,0,0]])\nprediction_less_good","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.035082Z","iopub.execute_input":"2023-09-14T07:17:00.035799Z","iopub.status.idle":"2023-09-14T07:17:00.047545Z","shell.execute_reply.started":"2023-09-14T07:17:00.035752Z","shell.execute_reply":"2023-09-14T07:17:00.04617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Undershooting is encouraged in some circumstance, the best dice score can be achieved with setting the 0.2 pixels from the gt to 0 in the prediction.\nprint('Outputs are: Dice Score, Intersection, sum(inp), sum(gt)')\nprint(dice(prediction_perfect, groundtruth))\nprint(dice(prediction_less_good, groundtruth))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.049239Z","iopub.execute_input":"2023-09-14T07:17:00.049666Z","iopub.status.idle":"2023-09-14T07:17:00.058792Z","shell.execute_reply.started":"2023-09-14T07:17:00.049632Z","shell.execute_reply":"2023-09-14T07:17:00.057579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Overshooting is encouraged in other circumstances, the best dice score here is achieved by setting the 0.2 pixel from the gt to 1 in the prediction.\ngroundtruth = np.array([0.2])\nprediction_perfect = np.array([0.2])\nprediction_lesser = np.array([1])\n\nprint(dice(prediction_perfect, groundtruth))\nprint(dice(prediction_lesser, groundtruth))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.060347Z","iopub.execute_input":"2023-09-14T07:17:00.061027Z","iopub.status.idle":"2023-09-14T07:17:00.072657Z","shell.execute_reply.started":"2023-09-14T07:17:00.060978Z","shell.execute_reply":"2023-09-14T07:17:00.07123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Solution","metadata":{}},{"cell_type":"code","source":"# Perfect solutions gain adjusted dice score 1, overshooting is discouraged\ngroundtruth = np.array([0.2])\nprediction_perfect = np.array([0.2])\nprediction_lesser = np.array([1])\n\nprint(adjusted_dice(prediction_perfect, groundtruth))\nprint(adjusted_dice(prediction_lesser, groundtruth))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.074007Z","iopub.execute_input":"2023-09-14T07:17:00.074941Z","iopub.status.idle":"2023-09-14T07:17:00.091868Z","shell.execute_reply.started":"2023-09-14T07:17:00.074891Z","shell.execute_reply":"2023-09-14T07:17:00.090842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Perfect solution gains adjusted dice score = 1 again.\ngroundtruth = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\nprediction_perfect = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\nprediction_lesser = np.array([[1,1,0,0],[1,1,0,0]])\n\nprint(adjusted_dice(prediction_perfect, groundtruth))\nprint(adjusted_dice(prediction_lesser, groundtruth))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.093605Z","iopub.execute_input":"2023-09-14T07:17:00.094275Z","iopub.status.idle":"2023-09-14T07:17:00.113067Z","shell.execute_reply.started":"2023-09-14T07:17:00.094241Z","shell.execute_reply":"2023-09-14T07:17:00.111727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## One additional example","metadata":{}},{"cell_type":"code","source":"groundtruth = np.array([[1,1,0.2,0.2],[1,1,0.2,0.2]])\ngroundtruth","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.115321Z","iopub.execute_input":"2023-09-14T07:17:00.116068Z","iopub.status.idle":"2023-09-14T07:17:00.127406Z","shell.execute_reply.started":"2023-09-14T07:17:00.116032Z","shell.execute_reply":"2023-09-14T07:17:00.126044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_better = np.array([[1,1,0,0],[1,1,0,0]])\nprediction_better","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.129015Z","iopub.execute_input":"2023-09-14T07:17:00.129411Z","iopub.status.idle":"2023-09-14T07:17:00.143507Z","shell.execute_reply.started":"2023-09-14T07:17:00.129378Z","shell.execute_reply":"2023-09-14T07:17:00.14209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_lesser = np.array([[0,0,0.2,0.2],[0,0,0.2,0.2]])\nprediction_lesser","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.147454Z","iopub.execute_input":"2023-09-14T07:17:00.148406Z","iopub.status.idle":"2023-09-14T07:17:00.157443Z","shell.execute_reply.started":"2023-09-14T07:17:00.148346Z","shell.execute_reply":"2023-09-14T07:17:00.156369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(adjusted_dice(prediction_better, groundtruth))\nprint(adjusted_dice(prediction_lesser, groundtruth))","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.162686Z","iopub.execute_input":"2023-09-14T07:17:00.163368Z","iopub.status.idle":"2023-09-14T07:17:00.170581Z","shell.execute_reply.started":"2023-09-14T07:17:00.163323Z","shell.execute_reply":"2023-09-14T07:17:00.169328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dice Loss Replacement for Pytorch \n#### + Small example to show that it works","metadata":{}},{"cell_type":"markdown","source":"## Imports and Modified Dice Loss in style of segmentation_models_pytorch","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\nimport time","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:00.172516Z","iopub.execute_input":"2023-09-14T07:17:00.173253Z","iopub.status.idle":"2023-09-14T07:17:28.500668Z","shell.execute_reply.started":"2023-09-14T07:17:00.173211Z","shell.execute_reply":"2023-09-14T07:17:28.499129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Credits to segmentation models pytorch\nfrom typing import Optional, List\nimport torch\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\n\nclass ModifiedDiceLoss(_Loss):\n    def __init__(\n        self,\n        mode: str,\n        classes: Optional[List[int]] = None,\n        log_loss: bool = False,\n        from_logits: bool = True,\n        smooth: float = 0.0,\n        ignore_index: Optional[int] = None,\n        eps: float = 1e-7,):\n            \n        super(ModifiedDiceLoss, self).__init__()\n        self.classes = classes\n        self.from_logits = from_logits\n        self.smooth = smooth\n        self.eps = eps\n        self.log_loss = log_loss\n        self.ignore_index = ignore_index\n        \n\n    def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:\n        assert y_true.size(0) == y_pred.size(0)\n        if self.from_logits:\n            # Apply activations to get [0..1] class probabilities\n            # Using Log-Exp as this gives more numerically stable result and does not cause vanishing gradient on\n            # extreme values 0 and 1\n            y_pred = F.logsigmoid(y_pred).exp()\n        bs = y_true.size(0)\n        num_classes = y_pred.size(1)\n        dims = (0, 2)\n        y_true = y_true.view(bs, 1, -1)\n        y_pred = y_pred.view(bs, 1, -1)\n        if self.ignore_index is not None:\n            mask = y_true != self.ignore_index\n            y_pred = y_pred * mask\n            y_true = y_true * mask\n        scores = self.compute_score(y_pred, y_true.type_as(y_pred), smooth=self.smooth, eps=self.eps, dims=dims)\n        if self.log_loss:\n            loss = -torch.log(scores.clamp_min(self.eps))\n        else:\n            loss = 1.0 - scores\n        mask = y_true.sum(dims) > 0\n        loss *= mask.to(loss.dtype)\n        if self.classes is not None:\n            loss = loss[self.classes]\n        return self.aggregate_loss(loss)\n\n    def aggregate_loss(self, loss):\n        return loss.mean()\n\n    def compute_score(self, output, target, smooth=0.0, eps=1e-7, dims=None) -> torch.Tensor:\n        return soft_dice_score(output, target, smooth, eps, dims)\n\ndef soft_dice_score(\n    output: torch.Tensor,\n    target: torch.Tensor,\n    smooth: float = 0.0,\n    eps: float = 1e-7,\n    dims=None,\n) -> torch.Tensor:\n    assert output.size() == target.size()\n    if dims is not None:\n        intersection = torch.sum((1-torch.abs(target - output)) * torch.minimum(target, output), dim=dims) # CHANGED FROM ORIGINAL\n        cardinality = torch.sum(output + target, dim=dims)\n    else:\n        intersection = torch.sum((1-torch.abs(target - output)) * torch.minimum(target, output))           # CHANGED FROM ORIGINAL\n        cardinality = torch.sum(output + target)\n    dice_score = (2.0 * intersection + smooth) / (cardinality + smooth).clamp_min(eps)\n    return dice_score","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:28.502421Z","iopub.execute_input":"2023-09-14T07:17:28.503439Z","iopub.status.idle":"2023-09-14T07:17:28.526448Z","shell.execute_reply.started":"2023-09-14T07:17:28.5034Z","shell.execute_reply":"2023-09-14T07:17:28.524543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation","metadata":{}},{"cell_type":"code","source":"def evaluate(original_dice=True, random_target=True):\n    t_start = time.time()\n    \n    # Step 1: Create a 2x2 target matrix with random values between 0 and 1\n    if random_target:\n        target = torch.rand((3, 3))   \n    else:   \n        target = torch.tensor([[1,0,1],[0,1,0],[1,0,1]])\n    \n    # Step 2: Create a randomized 2x2 prediction.\n    logit_prediction = torch.randn((3, 3))    \n    logit_prediction.requires_grad = True\n    optimizer = torch.optim.SGD([logit_prediction], lr=0.3)\n    \n    # Step 3: Set up a loss function (Dice Loss)\n    if original_dice:\n        loss_fn = smp.losses.DiceLoss(mode='binary') \n    else:\n        loss_fn = ModifiedDiceLoss(mode='binary') \n    \n    # Before Training\n    print('Original Prediction:')\n    print(torch.sigmoid(logit_prediction))\n    print('Ground Truth:')\n    print(target)\n    \n    # Step 5: Training loop\n    num_epochs = 10000\n    for epoch in range(num_epochs):\n        # Zero the parameter gradients\n        optimizer.zero_grad()\n\n        # Forward pass\n        loss = loss_fn(logit_prediction, target)\n\n        # Backward pass and optimization\n        loss.backward()\n        optimizer.step()\n\n        # Print loss\n        if epoch % 1000 == 0:\n            print(f'Epoch {epoch}, Loss: {loss.item()}')\n    \n    # After Training\n    print('Prediction after Training:')\n    print(torch.sigmoid(logit_prediction))\n    print(f'The process took {time.time()-t_start} seconds')\n    print()","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:28.528908Z","iopub.execute_input":"2023-09-14T07:17:28.529913Z","iopub.status.idle":"2023-09-14T07:17:28.548738Z","shell.execute_reply.started":"2023-09-14T07:17:28.529866Z","shell.execute_reply":"2023-09-14T07:17:28.547004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Evaluate simple target with 1s and 0s on both functions to show that both work for simple settings\nevaluate(original_dice=True, random_target=False)\nevaluate(original_dice=False, random_target=False)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:28.550125Z","iopub.execute_input":"2023-09-14T07:17:28.550523Z","iopub.status.idle":"2023-09-14T07:17:40.513115Z","shell.execute_reply.started":"2023-09-14T07:17:28.55049Z","shell.execute_reply":"2023-09-14T07:17:40.511617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show that original dice does not work for soft labels, as it almost always either goes to 1 or 0 in the optimization (see figure at the top)\nevaluate(original_dice=True, random_target=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:40.51578Z","iopub.execute_input":"2023-09-14T07:17:40.516203Z","iopub.status.idle":"2023-09-14T07:17:45.692738Z","shell.execute_reply.started":"2023-09-14T07:17:40.516168Z","shell.execute_reply":"2023-09-14T07:17:45.691384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Modified dice works with arbitrary target!\nevaluate(original_dice=False, random_target=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-14T07:17:45.694566Z","iopub.execute_input":"2023-09-14T07:17:45.695071Z","iopub.status.idle":"2023-09-14T07:17:52.312524Z","shell.execute_reply.started":"2023-09-14T07:17:45.695027Z","shell.execute_reply":"2023-09-14T07:17:52.311232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}