{
  "id": 213075,
  "title": "BCE Focal Loss for Pytorch",
  "url": "/competitions/rfcx-species-audio-detection/discussion/213075",
  "author_name": "",
  "post_date": "2021-01-21T11:55:01.462398500Z",
  "votes": 77,
  "comment_count": 9,
  "views": 0,
  "content": "<p>I see many focal loss implementations in Pytorch but they look complex to me.  Here is one that leverages built in BCE  loss. <code>preds</code> is that output of your model, and <code>targets</code> is the ground truth.  For BCE loss you would then use (I moved the average reduction outside the built-in function to make it easier to code focal loss)::</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nloss = bce_loss.mean()\n</code></pre>\n<p>For focal loss the idea is to weight bce loss by the complement of the predicted probability raised to some power gamma.  The code assumes batch dimension is the first dimension, as is usual with Pytorch.</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets &gt;= 0.5, (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n</code></pre>\n<p>The original focal loss paper also introduced an alpha weight for positive class.  It is a simple addition:</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n</code></pre>\n<p>The above code assumes that targets are binary. </p>\n<p>Edit. It is important to use the logits instead of taking the log or probabilities after a sigmoid layer.  Indeed, the latter has very bad numerical behavior when the probabilities are near zero.</p>\n<p>Edit2.  If targets are not binary one can use this code.  I haven't used it (yet) hence I have no clue about its usefulness.</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = targets * alpha * (1. - probas)**gamma * bce_loss + (1. -  targets) *  probas**gamma * bce_loss\nloss = loss.mean()\n</code></pre>\n<p>Edit3.  You can replace the last call to <code>mean</code> by any other reduction, for instance <code>sum</code> if the loss is too small.</p>",
  "messages": [
    {
      "id": "1162904",
      "postDate": "01/21/2021 11:55:01",
      "content": "<p>I see many focal loss implementations in Pytorch but they look complex to me.  Here is one that leverages built in BCE  loss. <code>preds</code> is that output of your model, and <code>targets</code> is the ground truth.  For BCE loss you would then use (I moved the average reduction outside the built-in function to make it easier to code focal loss)::</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nloss = bce_loss.mean()\n</code></pre>\n<p>For focal loss the idea is to weight bce loss by the complement of the predicted probability raised to some power gamma.  The code assumes batch dimension is the first dimension, as is usual with Pytorch.</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets &gt;= 0.5, (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n</code></pre>\n<p>The original focal loss paper also introduced an alpha weight for positive class.  It is a simple addition:</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n</code></pre>\n<p>The above code assumes that targets are binary. </p>\n<p>Edit. It is important to use the logits instead of taking the log or probabilities after a sigmoid layer.  Indeed, the latter has very bad numerical behavior when the probabilities are near zero.</p>\n<p>Edit2.  If targets are not binary one can use this code.  I haven't used it (yet) hence I have no clue about its usefulness.</p>\n<pre><code>loss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = targets * alpha * (1. - probas)**gamma * bce_loss + (1. -  targets) *  probas**gamma * bce_loss\nloss = loss.mean()\n</code></pre>\n<p>Edit3.  You can replace the last call to <code>mean</code> by any other reduction, for instance <code>sum</code> if the loss is too small.</p>",
      "rawMarkdown": "I see many focal loss implementations in Pytorch but they look complex to me.  Here is one that leverages built in BCE  loss. ` preds` is that output of your model, and `targets` is the ground truth.  For BCE loss you would then use (I moved the average reduction outside the built-in function to make it easier to code focal loss)::\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nloss = bce_loss.mean()\n```\nFor focal loss the idea is to weight bce loss by the complement of the predicted probability raised to some power gamma.  The code assumes batch dimension is the first dimension, as is usual with Pytorch.\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets >= 0.5, (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n```\nThe original focal loss paper also introduced an alpha weight for positive class.  It is a simple addition:\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n```\nThe above code assumes that targets are binary. \n\nEdit. It is important to use the logits instead of taking the log or probabilities after a sigmoid layer.  Indeed, the latter has very bad numerical behavior when the probabilities are near zero.\n\nEdit2.  If targets are not binary one can use this code.  I haven't used it (yet) hence I have no clue about its usefulness.\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = targets * alpha * (1. - probas)**gamma * bce_loss + (1. -  targets) *  probas**gamma * bce_loss\nloss = loss.mean()\n```\n\nEdit3.  You can replace the last call to `mean` by any other reduction, for instance `sum` if the loss is too small.",
      "votes": null
    },
    {
      "id": "1163017",
      "postDate": "01/21/2021 13:05:10",
      "content": "<p>Thank you!</p>\n<p>This is instead my Focal Loss and Binary Focal Loss, in an object.</p>\n<pre><code>##################################################\n# Focal Loss\n##################################################\n\nclass FocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2, preds_with_logits=True, smooth_factor=0.0, labels_are_oh=False):\n        self.gamma = gamma\n        self.preds_with_logits = preds_with_logits\n        self.smooth_factor = smooth_factor\n        self.labels_are_oh = labels_are_oh\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        if not self.labels_are_oh:\n            num_classes = preds.shape[1]\n            labels = F.one_hot(labels, num_classes).to(preds.dtype)\n            labels = labels.transpose(1, -1)\n\n        if self.smooth_factor &gt; 0:\n            num_classes = preds.shape[1]\n            labels = (1.0 - self.smooth_factor) * labels + self.smooth_factor / num_classes\n\n        if self.preds_with_logits:\n            preds = F.softmax(preds, 1)\n            preds = torch.where(preds &gt; 1e-8, preds, 1e-8 * torch.ones_like(preds))\n\n        fl = (- labels * (1 - preds).pow(self.gamma) * torch.log(preds)).sum(1)\n        if reduce_bach == 'mean':\n            fl = fl.mean()\n        elif reduce_beatch == 'sum':\n            fl = fl.sum()\n        return fl\n</code></pre>\n<pre><code>##################################################\n# Binary Focall Loss\n##################################################\n\nclass BFocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2):\n        self.gamma = gamma\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        \"\"\"\n        Args:\n            preds: tensor of shape [batch_size, num_classes], sigmoid is already applied.\n            labels: tensor of shape [batch_size, num_classes].\n        \"\"\"\n        preds = torch.where(preds &lt; 1e-8, 1e-8 * torch.ones_like(preds), preds)\n        bfl = - labels * (1 - preds).pow(self.gamma) * torch.log(preds)\n        bfl = bfl - (1 - labels) * preds.pow(self.gamma) * torch.log(1 - preds)\n        if reduce_bach == 'mean':\n            bfl = bfl.mean()\n        elif reduce_beatch == 'sum':\n            bfl = bfl.sum()\n        return bfl\n</code></pre>",
      "rawMarkdown": "Thank you!\n\nThis is instead my Focal Loss and Binary Focal Loss, in an object.\n\n```python\n##################################################\n# Focal Loss\n##################################################\n\nclass FocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2, preds_with_logits=True, smooth_factor=0.0, labels_are_oh=False):\n        self.gamma = gamma\n        self.preds_with_logits = preds_with_logits\n        self.smooth_factor = smooth_factor\n        self.labels_are_oh = labels_are_oh\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        if not self.labels_are_oh:\n            num_classes = preds.shape[1]\n            labels = F.one_hot(labels, num_classes).to(preds.dtype)\n            labels = labels.transpose(1, -1)\n\n        if self.smooth_factor > 0:\n            num_classes = preds.shape[1]\n            labels = (1.0 - self.smooth_factor) * labels + self.smooth_factor / num_classes\n\n        if self.preds_with_logits:\n            preds = F.softmax(preds, 1)\n            preds = torch.where(preds > 1e-8, preds, 1e-8 * torch.ones_like(preds))\n\n        fl = (- labels * (1 - preds).pow(self.gamma) * torch.log(preds)).sum(1)\n        if reduce_bach == 'mean':\n            fl = fl.mean()\n        elif reduce_beatch == 'sum':\n            fl = fl.sum()\n        return fl\n```\n\n```python\n##################################################\n# Binary Focall Loss\n##################################################\n\nclass BFocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2):\n        self.gamma = gamma\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        \"\"\"\n        Args:\n            preds: tensor of shape [batch_size, num_classes], sigmoid is already applied.\n            labels: tensor of shape [batch_size, num_classes].\n        \"\"\"\n        preds = torch.where(preds < 1e-8, 1e-8 * torch.ones_like(preds), preds)\n        bfl = - labels * (1 - preds).pow(self.gamma) * torch.log(preds)\n        bfl = bfl - (1 - labels) * preds.pow(self.gamma) * torch.log(1 - preds)\n        if reduce_bach == 'mean':\n            bfl = bfl.mean()\n        elif reduce_beatch == 'sum':\n            bfl = bfl.sum()\n        return bfl\n```",
      "votes": null
    },
    {
      "id": "1163020",
      "postDate": "01/21/2021 13:08:10",
      "content": "<p>You take the log of predictions in the second version, which is unstable numerically.</p>\n<p>Maybe it works fine though, but I am worried by the numerical instability.</p>\n<p>Edited per next comment.</p>",
      "rawMarkdown": "You take the log of predictions in the second version, which is unstable numerically.\n\n Maybe it works fine though, but I am worried by the numerical instability.\n\nEdited per next comment.",
      "votes": null
    },
    {
      "id": "1163030",
      "postDate": "01/21/2021 13:16:32",
      "content": "<ul>\n<li>for the numerical instability I simpy clip before taking the log:</li>\n</ul>\n<pre><code>preds = torch.where(preds &lt; 1e-8, 1e-8 * torch.ones_like(preds), preds)\n</code></pre>\n<ul>\n<li>for computing binary focal loss you need <code>p</code> in [0, 1] like the output of the sigmoid, I'm not inverting the sigmoid.</li>\n</ul>\n<p>I think my implementation is correct.</p>",
      "rawMarkdown": "for the numerical instability I simpy clip before taking the log:\n```\npreds = torch.where(preds < 1e-8, 1e-8 * torch.ones_like(preds), preds)\n```\n\n- for computing binary focal loss you need `p` in [0, 1] like the output of the sigmoid, I'm not inverting the sigmoid.\n\nI think my implementation is correct.",
      "votes": null
    },
    {
      "id": "1163045",
      "postDate": "01/21/2021 13:33:11",
      "content": "<p>Yes, it is mathematically correct, I've edited my previous comment.  But it is still numerically unstable.  That's why I use the native bce with logits which implements the log of sigmoid in a numerically stable way.</p>",
      "rawMarkdown": "Yes, it is mathematically correct, I've edited my previous comment.  But it is still numerically unstable.  That's why I use the native bce with logits which implements the log of sigmoid in a numerically stable way.",
      "votes": null
    },
    {
      "id": "1163048",
      "postDate": "01/21/2021 13:38:30",
      "content": "<p>Thank you for saying that is correct :)</p>\n<p>I'm checking either the gradients and the prediction values. For now, I haven't any numerical instability.<br>\nIn general, if you clip something before the log, in the region close to 0, you avoid in toto the instability. </p>\n<p>I just introduce some noise in the prediciton (very very small, the numbers are less than 1e-8).</p>\n<p>Best,</p>\n<p>Guglielmo</p>",
      "rawMarkdown": "Thank you for saying that is correct :)\n\nI'm checking either the gradients and the prediction values. For now, I haven't any numerical instability.\nIn general, if you clip something before the log, in the region close to 0, you avoid in toto the instability. \n\nI just introduce some noise in the prediciton (very very small, the numbers are less than 1e-8).\n\nBest,\n\nGuglielmo",
      "votes": null
    },
    {
      "id": "1163128",
      "postDate": "01/21/2021 14:56:24",
      "content": "<p>You must also clip probabilities near 1.</p>\n<p>Anyway, I don't get what is the upside to use clipping when there is a safe built in function.  Maybe I am missing something.</p>",
      "rawMarkdown": "You must also clip probabilities near 1.\n\nAnyway, I don't get what is the upside to use clipping when there is a safe built in function.  Maybe I am missing something.",
      "votes": null
    },
    {
      "id": "1164064",
      "postDate": "01/22/2021 05:50:43",
      "content": "<blockquote>\n  <p>loss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)<strong>gamma * bce_loss, probas</strong>gamma * bce_loss)</p>\n</blockquote>\n<p>Should Alpha only be applied on positive class? In my understanding Alpha should be also applied on negative class as:</p>\n<p><code>loss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)**gamma * bce_loss, (1-alpha) * probas**gamma * bce_loss)</code></p>",
      "rawMarkdown": "> loss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\n\nShould Alpha only be applied on positive class? In my understanding Alpha should be also applied on negative class as:\n\n`loss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, (1-alpha) * probas**gamma * bce_loss)`",
      "votes": null
    },
    {
      "id": "1164293",
      "postDate": "01/22/2021 09:23:00",
      "content": "<p>It is almost equivalent.  In the paper there is one alpha per class, i.e.</p>\n<p><code>loss = torch.where(targets &gt;= 0.5, alpha1 * (1. - probas)**gamma * bce_loss, alpha0 * probas**gamma * bce_loss)</code></p>\n<p>If you divide loss by <code>alpha0</code> then you get my version.  If you divide by <code>alpha0 + alpha1</code> then you get your version.</p>",
      "rawMarkdown": "It is almost equivalent.  In the paper there is one alpha per class, i.e.\n\n`loss = torch.where(targets >= 0.5, alpha1 * (1. - probas)**gamma * bce_loss, alpha0 * probas**gamma * bce_loss)`\n\nIf you divide loss by `alpha0` then you get my version.  If you divide by `alpha0 + alpha1` then you get your version.",
      "votes": null
    },
    {
      "id": "1181110",
      "postDate": "02/01/2021 16:47:31",
      "content": "<p>I agree using clipping there may become problematic, specially if afterwards you want to use schedule target clipping as described in the spec augment paper. You would be mixing purpose.</p>",
      "rawMarkdown": "I agree using clipping there may become problematic, specially if afterwards you want to use schedule target clipping as described in the spec augment paper. You would be mixing purpose.",
      "votes": null
    }
  ],
  "comments": [
    {
      "id": 1163017,
      "author_name": "guglielmocamporese",
      "author_url": "",
      "post_date": "01/21/2021 13:05:10",
      "content": "<p>Thank you!</p>\n<p>This is instead my Focal Loss and Binary Focal Loss, in an object.</p>\n<pre><code>##################################################\n# Focal Loss\n##################################################\n\nclass FocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2, preds_with_logits=True, smooth_factor=0.0, labels_are_oh=False):\n        self.gamma = gamma\n        self.preds_with_logits = preds_with_logits\n        self.smooth_factor = smooth_factor\n        self.labels_are_oh = labels_are_oh\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        if not self.labels_are_oh:\n            num_classes = preds.shape[1]\n            labels = F.one_hot(labels, num_classes).to(preds.dtype)\n            labels = labels.transpose(1, -1)\n\n        if self.smooth_factor &gt; 0:\n            num_classes = preds.shape[1]\n            labels = (1.0 - self.smooth_factor) * labels + self.smooth_factor / num_classes\n\n        if self.preds_with_logits:\n            preds = F.softmax(preds, 1)\n            preds = torch.where(preds &gt; 1e-8, preds, 1e-8 * torch.ones_like(preds))\n\n        fl = (- labels * (1 - preds).pow(self.gamma) * torch.log(preds)).sum(1)\n        if reduce_bach == 'mean':\n            fl = fl.mean()\n        elif reduce_beatch == 'sum':\n            fl = fl.sum()\n        return fl\n</code></pre>\n<pre><code>##################################################\n# Binary Focall Loss\n##################################################\n\nclass BFocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2):\n        self.gamma = gamma\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        \"\"\"\n        Args:\n            preds: tensor of shape [batch_size, num_classes], sigmoid is already applied.\n            labels: tensor of shape [batch_size, num_classes].\n        \"\"\"\n        preds = torch.where(preds &lt; 1e-8, 1e-8 * torch.ones_like(preds), preds)\n        bfl = - labels * (1 - preds).pow(self.gamma) * torch.log(preds)\n        bfl = bfl - (1 - labels) * preds.pow(self.gamma) * torch.log(1 - preds)\n        if reduce_bach == 'mean':\n            bfl = bfl.mean()\n        elif reduce_beatch == 'sum':\n            bfl = bfl.sum()\n        return bfl\n</code></pre>",
      "votes": null,
      "replies": [
        {
          "id": 1163020,
          "author_name": "cpmpml",
          "author_url": "",
          "post_date": "01/21/2021 13:08:10",
          "content": "<p>You take the log of predictions in the second version, which is unstable numerically.</p>\n<p>Maybe it works fine though, but I am worried by the numerical instability.</p>\n<p>Edited per next comment.</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1163030,
          "author_name": "guglielmocamporese",
          "author_url": "",
          "post_date": "01/21/2021 13:16:32",
          "content": "<ul>\n<li>for the numerical instability I simpy clip before taking the log:</li>\n</ul>\n<pre><code>preds = torch.where(preds &lt; 1e-8, 1e-8 * torch.ones_like(preds), preds)\n</code></pre>\n<ul>\n<li>for computing binary focal loss you need <code>p</code> in [0, 1] like the output of the sigmoid, I'm not inverting the sigmoid.</li>\n</ul>\n<p>I think my implementation is correct.</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1163045,
          "author_name": "cpmpml",
          "author_url": "",
          "post_date": "01/21/2021 13:33:11",
          "content": "<p>Yes, it is mathematically correct, I've edited my previous comment.  But it is still numerically unstable.  That's why I use the native bce with logits which implements the log of sigmoid in a numerically stable way.</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1163048,
          "author_name": "guglielmocamporese",
          "author_url": "",
          "post_date": "01/21/2021 13:38:30",
          "content": "<p>Thank you for saying that is correct :)</p>\n<p>I'm checking either the gradients and the prediction values. For now, I haven't any numerical instability.<br>\nIn general, if you clip something before the log, in the region close to 0, you avoid in toto the instability. </p>\n<p>I just introduce some noise in the prediciton (very very small, the numbers are less than 1e-8).</p>\n<p>Best,</p>\n<p>Guglielmo</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1163128,
          "author_name": "cpmpml",
          "author_url": "",
          "post_date": "01/21/2021 14:56:24",
          "content": "<p>You must also clip probabilities near 1.</p>\n<p>Anyway, I don't get what is the upside to use clipping when there is a safe built in function.  Maybe I am missing something.</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1181110,
          "author_name": "felipebihaiek",
          "author_url": "",
          "post_date": "02/01/2021 16:47:31",
          "content": "<p>I agree using clipping there may become problematic, specially if afterwards you want to use schedule target clipping as described in the spec augment paper. You would be mixing purpose.</p>",
          "votes": null,
          "replies": []
        }
      ]
    },
    {
      "id": 1164064,
      "author_name": "superchenhao",
      "author_url": "",
      "post_date": "01/22/2021 05:50:43",
      "content": "<blockquote>\n  <p>loss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)<strong>gamma * bce_loss, probas</strong>gamma * bce_loss)</p>\n</blockquote>\n<p>Should Alpha only be applied on positive class? In my understanding Alpha should be also applied on negative class as:</p>\n<p><code>loss = torch.where(targets &gt;= 0.5, alpha * (1. - probas)**gamma * bce_loss, (1-alpha) * probas**gamma * bce_loss)</code></p>",
      "votes": null,
      "replies": [
        {
          "id": 1164293,
          "author_name": "cpmpml",
          "author_url": "",
          "post_date": "01/22/2021 09:23:00",
          "content": "<p>It is almost equivalent.  In the paper there is one alpha per class, i.e.</p>\n<p><code>loss = torch.where(targets &gt;= 0.5, alpha1 * (1. - probas)**gamma * bce_loss, alpha0 * probas**gamma * bce_loss)</code></p>\n<p>If you divide loss by <code>alpha0</code> then you get my version.  If you divide by <code>alpha0 + alpha1</code> then you get your version.</p>",
          "votes": null,
          "replies": []
        }
      ]
    }
  ],
  "raw_markdown_by_id": {
    "1162904": "I see many focal loss implementations in Pytorch but they look complex to me.  Here is one that leverages built in BCE  loss. ` preds` is that output of your model, and `targets` is the ground truth.  For BCE loss you would then use (I moved the average reduction outside the built-in function to make it easier to code focal loss)::\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nloss = bce_loss.mean()\n```\nFor focal loss the idea is to weight bce loss by the complement of the predicted probability raised to some power gamma.  The code assumes batch dimension is the first dimension, as is usual with Pytorch.\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets >= 0.5, (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n```\nThe original focal loss paper also introduced an alpha weight for positive class.  It is a simple addition:\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\nloss = loss.mean()\n```\nThe above code assumes that targets are binary. \n\nEdit. It is important to use the logits instead of taking the log or probabilities after a sigmoid layer.  Indeed, the latter has very bad numerical behavior when the probabilities are near zero.\n\nEdit2.  If targets are not binary one can use this code.  I haven't used it (yet) hence I have no clue about its usefulness.\n\n```\nloss_fct = nn.BCEWithLogitsLoss(reduction='none')\nbce_loss = loss_fct(preds, targets)\nprobas = torch.sigmoid(preds)\nloss = targets * alpha * (1. - probas)**gamma * bce_loss + (1. -  targets) *  probas**gamma * bce_loss\nloss = loss.mean()\n```\n\nEdit3.  You can replace the last call to `mean` by any other reduction, for instance `sum` if the loss is too small.",
    "1163017": "Thank you!\n\nThis is instead my Focal Loss and Binary Focal Loss, in an object.\n\n```python\n##################################################\n# Focal Loss\n##################################################\n\nclass FocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2, preds_with_logits=True, smooth_factor=0.0, labels_are_oh=False):\n        self.gamma = gamma\n        self.preds_with_logits = preds_with_logits\n        self.smooth_factor = smooth_factor\n        self.labels_are_oh = labels_are_oh\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        if not self.labels_are_oh:\n            num_classes = preds.shape[1]\n            labels = F.one_hot(labels, num_classes).to(preds.dtype)\n            labels = labels.transpose(1, -1)\n\n        if self.smooth_factor > 0:\n            num_classes = preds.shape[1]\n            labels = (1.0 - self.smooth_factor) * labels + self.smooth_factor / num_classes\n\n        if self.preds_with_logits:\n            preds = F.softmax(preds, 1)\n            preds = torch.where(preds > 1e-8, preds, 1e-8 * torch.ones_like(preds))\n\n        fl = (- labels * (1 - preds).pow(self.gamma) * torch.log(preds)).sum(1)\n        if reduce_bach == 'mean':\n            fl = fl.mean()\n        elif reduce_beatch == 'sum':\n            fl = fl.sum()\n        return fl\n```\n\n```python\n##################################################\n# Binary Focall Loss\n##################################################\n\nclass BFocalLoss(object):\n    \"\"\"\n    From https://arxiv.org/pdf/1708.02002.pdf\n    \"\"\"\n    def __init__(self, gamma=2):\n        self.gamma = gamma\n\n    def __call__(self, preds, labels, reduce_bach='mean'):\n        \"\"\"\n        Args:\n            preds: tensor of shape [batch_size, num_classes], sigmoid is already applied.\n            labels: tensor of shape [batch_size, num_classes].\n        \"\"\"\n        preds = torch.where(preds < 1e-8, 1e-8 * torch.ones_like(preds), preds)\n        bfl = - labels * (1 - preds).pow(self.gamma) * torch.log(preds)\n        bfl = bfl - (1 - labels) * preds.pow(self.gamma) * torch.log(1 - preds)\n        if reduce_bach == 'mean':\n            bfl = bfl.mean()\n        elif reduce_beatch == 'sum':\n            bfl = bfl.sum()\n        return bfl\n```",
    "1163020": "You take the log of predictions in the second version, which is unstable numerically.\n\n Maybe it works fine though, but I am worried by the numerical instability.\n\nEdited per next comment.",
    "1163030": "for the numerical instability I simpy clip before taking the log:\n```\npreds = torch.where(preds < 1e-8, 1e-8 * torch.ones_like(preds), preds)\n```\n\n- for computing binary focal loss you need `p` in [0, 1] like the output of the sigmoid, I'm not inverting the sigmoid.\n\nI think my implementation is correct.",
    "1163045": "Yes, it is mathematically correct, I've edited my previous comment.  But it is still numerically unstable.  That's why I use the native bce with logits which implements the log of sigmoid in a numerically stable way.",
    "1163048": "Thank you for saying that is correct :)\n\nI'm checking either the gradients and the prediction values. For now, I haven't any numerical instability.\nIn general, if you clip something before the log, in the region close to 0, you avoid in toto the instability. \n\nI just introduce some noise in the prediciton (very very small, the numbers are less than 1e-8).\n\nBest,\n\nGuglielmo",
    "1163128": "You must also clip probabilities near 1.\n\nAnyway, I don't get what is the upside to use clipping when there is a safe built in function.  Maybe I am missing something.",
    "1164064": "> loss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, probas**gamma * bce_loss)\n\nShould Alpha only be applied on positive class? In my understanding Alpha should be also applied on negative class as:\n\n`loss = torch.where(targets >= 0.5, alpha * (1. - probas)**gamma * bce_loss, (1-alpha) * probas**gamma * bce_loss)`",
    "1164293": "It is almost equivalent.  In the paper there is one alpha per class, i.e.\n\n`loss = torch.where(targets >= 0.5, alpha1 * (1. - probas)**gamma * bce_loss, alpha0 * probas**gamma * bce_loss)`\n\nIf you divide loss by `alpha0` then you get my version.  If you divide by `alpha0 + alpha1` then you get your version.",
    "1181110": "I agree using clipping there may become problematic, specially if afterwards you want to use schedule target clipping as described in the spec augment paper. You would be mixing purpose."
  },
  "source": "meta"
}