• yolov5 focal_loss源码解析


     一、以下为yolov5损失函数的源代码:

    1. import torch
    2. import torch.nn as nn
    3. import numpy as np
    4. class FocalLoss(nn.Module):
    5. # Wraps focal loss around existing loss_fcn(), i.e. criteria = FocalLoss(nn.BCEWithLogitsLoss(), gamma=1.5)
    6. def __init__(self, loss_fcn, gamma=1.5, alpha=0.25):
    7. super().__init__()
    8. self.loss_fcn = loss_fcn # must be nn.BCEWithLogitsLoss()
    9. self.gamma = gamma
    10. self.alpha = alpha
    11. self.reduction = loss_fcn.reduction
    12. self.loss_fcn.reduction = 'none' # required to apply FL to each element
    13. def forward(self, pred, true):
    14. print("pred:",pred.shape)
    15. print("true:", true.shape)
    16. loss = self.loss_fcn(pred, true)
    17. print("loss:",loss.shape)
    18. pred_prob = torch.sigmoid(pred) # prob from logits
    19. print("pred_prob:", pred_prob.shape)
    20. p_t = true * pred_prob + (1 - true) * (1 - pred_prob) # 计算概率
    21. print("p_t:",p_t.shape)
    22. alpha_factor = true * self.alpha + (1 - true) * (1 - self.alpha)
    23. print("alpha_factor:", alpha_factor.shape)
    24. modulating_factor = (1.0 - p_t) ** self.gamma
    25. print("modulating_factor:", modulating_factor.shape)
    26. """
    27. loss *= alpha_factor * modulating_factor
    28. 等同于:
    29. loss = loss * [true * self.alpha + (1 - true) * (1 - self.alpha)] * [(1.0 - p_t) ** self.gamma]
    30. loss = true * self.alpha * [(1.0 - p_t) ** self.gamma] * loss + (1 - true) * (1 - self.alpha) * [(1.0 - p_t) ** self.gamma] * loss
    31. """
    32. loss *= alpha_factor * modulating_factor
    33. print("focal_loss:", loss.shape)
    34. return loss.mean()
    35. if __name__ == "__main__":
    36. loss_fcn = nn.BCEWithLogitsLoss()
    37. focal_loss = FocalLoss(loss_fcn, gamma=1.5, alpha=0.25)
    38. pred = np.random.random((852,2))
    39. true = np.random.random((852, 2))
    40. pred = torch.tensor(pred)
    41. true = torch.tensor(true)
    42. loss = focal_loss(pred,true)
    43. print(loss)

    1、yolov5中如果开启focal_loss函数,则默认是分类损失、置信度损失都使用FocalLoss; 假设预测出852个目标框,则输出如下:

    1. pred: torch.Size([852, 2])
    2. true: torch.Size([852, 2])
    3. loss: torch.Size([852, 2])
    4. pred_prob: torch.Size([852, 2])
    5. p_t: torch.Size([852, 2])
    6. alpha_factor: torch.Size([852, 2])
    7. modulating_factor: torch.Size([852, 2])
    8. focal_loss: torch.Size([852, 2])
    9. tensor(0.1520, dtype=torch.float64)

    2、 当把yolov5修改为旋转目标检测时,会增加角度信息,会多一个角度loss,角度范围[0,180],则对应的输出为:

    1. pred: torch.Size([852, 180])
    2. true: torch.Size([852, 180])
    3. loss: torch.Size([852, 180])
    4. pred_prob: torch.Size([852, 180])
    5. p_t: torch.Size([852, 180])
    6. alpha_factor: torch.Size([852, 180])
    7. modulating_factor: torch.Size([852, 180])
    8. focal_loss: torch.Size([852, 180])
    9. tensor(0.1536, dtype=torch.float64)

    3、代码中的loss计算公式可以解析如下:

    loss *= alpha_factor * modulating_factor
    等同于:
    loss = loss * [true * self.alpha + (1 - true) * (1 - self.alpha)] * [(1.0 - p_t) ** self.gamma]
    等同于:
    loss = true * self.alpha * [(1.0 - p_t) ** self.gamma] * loss  +   (1 - true) * (1 - self.alpha) *  [(1.0 - p_t) ** self.gamma] * loss 


    二、FocalLoss数学公式解析:

    1、以下是focalloss的数学公式,可与代码对照着看,以下两者都对:

    2、另外,值得注意的是,大多数的讲解中,会有交叉熵和focalloss的公式对比,但是,在计算loss时,都会乘以真实标签,所以,在上述代码中会发现也乘了一个True,即真实的标签;

     

  • 相关阅读:
    十七、【渐变工具组】
    numpy常用创建
    电源自动测试系统NSAT-8000,精准高速可靠的电源测试设备
    Vue+python+django高校田径运动会成绩报名系统pycharm源码lw
    sql报错:sql injection violation, syntax error
    LeetCode 0144. 二叉树的前序遍历:二叉树必会题
    Mac 终端连接数据库
    UE5 C++ 使用TimeLine时间轴实现开关门
    el-tab 滚动条滚动到指定位置
    docker-compose概述与简单编排部署
  • 原文地址:https://blog.csdn.net/pangxing6491/article/details/127923914