{
  "id": 372054,
  "title": "[PyTorch] Focal Loss using BCEWithLogitsLoss",
  "url": "/competitions/rsna-breast-cancer-detection/discussion/372054",
  "author_name": "moth",
  "post_date": "2022-12-13T21:09:24.526000",
  "votes": 16,
  "comment_count": 2,
  "views": 0,
  "content": "<h1><a href=\"https://amaarora.github.io/2020/06/29/FocalLoss.html\" target=\"_blank\">What is Focal Loss?</a></h1>\n<p><a href=\"https://arxiv.org/pdf/1708.02002.pdf\" target=\"_blank\">Paper: Focal Loss for Dense Object Detection</a></p>\n<p>The Focal Loss function addresses class imbalance during training in tasks like image classification. It applies a modulating term to the cross entropy loss in order to focus learning on hard misclassified examples. </p>\n<p>The Focal Loss is defined as:</p>\n<p>$$FL(p_{t}) = -\\alpha_{t}(1-p_{t})^{\\gamma}log(p_{t})$$</p>\n<p>Here is the PyTorch implementation of the Focal Loss for BCE:</p>\n<pre><code>criterion = nn.BCEWithLogitsLoss(reduction='none')\nGAMMA = 2 # default value\nALPHA = 0.25 # default value\n\n\nfor epoch in epochs:\n\n    ...\n\n    for images, labels in train_data_loader:\n\n        images = torch.tensor(images, device=device)\n        labels = torch.tensor(labels, device=device)\n        optimizer.zero_grad()\n        out = model(images)\n        # ===== FOCAL LOSS ======\n        bce_loss = criterion(out, labels.unsqueeze(1))\n        probas = torch.sigmoid(out)\n        loss = torch.where(labels &gt;= 0.5,\n                           ALPHA * (1-probas)**GAMMA * bce_loss,\n                           (1-ALPHA) * probas**GAMMA * bce_loss)\n        loss = loss.mean()\n        loss.backward()\n        # ======= END ========\n        optimizer.step()\n\n        ...\n</code></pre>",
  "messages": [
    {
      "id": 2064535,
      "postDate": "2022-12-13T21:09:24.527Z",
      "content": "<h1><a href=\"https://amaarora.github.io/2020/06/29/FocalLoss.html\" target=\"_blank\">What is Focal Loss?</a></h1>\n<p><a href=\"https://arxiv.org/pdf/1708.02002.pdf\" target=\"_blank\">Paper: Focal Loss for Dense Object Detection</a></p>\n<p>The Focal Loss function addresses class imbalance during training in tasks like image classification. It applies a modulating term to the cross entropy loss in order to focus learning on hard misclassified examples. </p>\n<p>The Focal Loss is defined as:</p>\n<p>$$FL(p_{t}) = -\\alpha_{t}(1-p_{t})^{\\gamma}log(p_{t})$$</p>\n<p>Here is the PyTorch implementation of the Focal Loss for BCE:</p>\n<pre><code>criterion = nn.BCEWithLogitsLoss(reduction='none')\nGAMMA = 2 # default value\nALPHA = 0.25 # default value\n\n\nfor epoch in epochs:\n\n    ...\n\n    for images, labels in train_data_loader:\n\n        images = torch.tensor(images, device=device)\n        labels = torch.tensor(labels, device=device)\n        optimizer.zero_grad()\n        out = model(images)\n        # ===== FOCAL LOSS ======\n        bce_loss = criterion(out, labels.unsqueeze(1))\n        probas = torch.sigmoid(out)\n        loss = torch.where(labels &gt;= 0.5,\n                           ALPHA * (1-probas)**GAMMA * bce_loss,\n                           (1-ALPHA) * probas**GAMMA * bce_loss)\n        loss = loss.mean()\n        loss.backward()\n        # ======= END ========\n        optimizer.step()\n\n        ...\n</code></pre>",
      "rawMarkdown": "# [What is Focal Loss?](https://amaarora.github.io/2020/06/29/FocalLoss.html)\n\n[Paper: Focal Loss for Dense Object Detection](https://arxiv.org/pdf/1708.02002.pdf)\n\nThe Focal Loss function addresses class imbalance during training in tasks like image classification. It applies a modulating term to the cross entropy loss in order to focus learning on hard misclassified examples. \n\nThe Focal Loss is defined as:\n\n$$FL(p_{t}) = -\\alpha_{t}(1-p_{t})^{\\gamma}log(p_{t})$$\n\nHere is the PyTorch implementation of the Focal Loss for BCE:\n\n```\ncriterion = nn.BCEWithLogitsLoss(reduction='none')\nGAMMA = 2 # default value\nALPHA = 0.25 # default value\n\n\nfor epoch in epochs:\n\t\n\t...\n\n\tfor images, labels in train_data_loader:\n\n        images = torch.tensor(images, device=device)\n        labels = torch.tensor(labels, device=device)\n        optimizer.zero_grad()\n        out = model(images)\n        # ===== FOCAL LOSS ======\n        bce_loss = criterion(out, labels.unsqueeze(1))\n        probas = torch.sigmoid(out)\n        loss = torch.where(labels >= 0.5,\n                           ALPHA * (1-probas)**GAMMA * bce_loss,\n                           (1-ALPHA) * probas**GAMMA * bce_loss)\n        loss = loss.mean()\n        loss.backward()\n        # ======= END ========\n        optimizer.step()\n\n        ...\n```",
      "votes": 16
    },
    {
      "id": 2067983,
      "postDate": "2022-12-17T11:47:18.360Z",
      "content": "<p>Hello !</p>\n<p>Have you seen this implementation of the focal loss by torchvision : <a href=\"https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html\" target=\"_blank\">https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html</a></p>\n<p>It also adds a reduction parameter and is made by torchvision so it is quite reliable.</p>",
      "rawMarkdown": "Hello !\n\nHave you seen this implementation of the focal loss by torchvision : https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html\n\nIt also adds a reduction parameter and is made by torchvision so it is quite reliable.",
      "votes": 2,
      "replies": [
        {
          "id": 2068131,
          "postDate": "2022-12-17T14:48:59.130Z",
          "content": "<p>I had not! Thanks for pointing it out.</p>",
          "rawMarkdown": "I had not! Thanks for pointing it out."
        }
      ]
    }
  ],
  "comments": [
    {
      "id": 2067983,
      "author_name": "Natyu",
      "author_url": "",
      "post_date": "2022-12-17T11:47:18.360000",
      "content": "<p>Hello !</p>\n<p>Have you seen this implementation of the focal loss by torchvision : <a href=\"https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html\" target=\"_blank\">https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html</a></p>\n<p>It also adds a reduction parameter and is made by torchvision so it is quite reliable.</p>",
      "votes": 2,
      "replies": [
        {
          "id": 2068131,
          "author_name": "moth",
          "author_url": "",
          "post_date": "2022-12-17T14:48:59.130000",
          "content": "<p>I had not! Thanks for pointing it out.</p>",
          "votes": 0,
          "replies": []
        }
      ]
    }
  ],
  "raw_markdown_by_id": {
    "2064535": "# [What is Focal Loss?](https://amaarora.github.io/2020/06/29/FocalLoss.html)\n\n[Paper: Focal Loss for Dense Object Detection](https://arxiv.org/pdf/1708.02002.pdf)\n\nThe Focal Loss function addresses class imbalance during training in tasks like image classification. It applies a modulating term to the cross entropy loss in order to focus learning on hard misclassified examples. \n\nThe Focal Loss is defined as:\n\n$$FL(p_{t}) = -\\alpha_{t}(1-p_{t})^{\\gamma}log(p_{t})$$\n\nHere is the PyTorch implementation of the Focal Loss for BCE:\n\n```\ncriterion = nn.BCEWithLogitsLoss(reduction='none')\nGAMMA = 2 # default value\nALPHA = 0.25 # default value\n\n\nfor epoch in epochs:\n\t\n\t...\n\n\tfor images, labels in train_data_loader:\n\n        images = torch.tensor(images, device=device)\n        labels = torch.tensor(labels, device=device)\n        optimizer.zero_grad()\n        out = model(images)\n        # ===== FOCAL LOSS ======\n        bce_loss = criterion(out, labels.unsqueeze(1))\n        probas = torch.sigmoid(out)\n        loss = torch.where(labels >= 0.5,\n                           ALPHA * (1-probas)**GAMMA * bce_loss,\n                           (1-ALPHA) * probas**GAMMA * bce_loss)\n        loss = loss.mean()\n        loss.backward()\n        # ======= END ========\n        optimizer.step()\n\n        ...\n```",
    "2067983": "Hello !\n\nHave you seen this implementation of the focal loss by torchvision : https://pytorch.org/vision/main/_modules/torchvision/ops/focal_loss.html\n\nIt also adds a reduction parameter and is made by torchvision so it is quite reliable."
  }
}