mmocr DBLoss
·
目录
DBLoss由三种loss组成。
1. balance bce loss
balance_bce_loss(pred=pred_prob, gt=gt_shrink, mask=gt_shrink_mask)


def balance_bce_loss(self, pred, gt, mask):
"""
pred: (b, w,h) 预测分数image, (0, 1)之间;
gt: (b, w,h) label image, 值0或者1;
mask:(b, w,h) 屏蔽背景位置,其中屏蔽位置为0,保留位置为1
"""
positive = (gt * mask) # positive前景位置,前景位置值为1,其他位置为0(利用mask屏蔽不需要训练的位置,比如字符检测时,模糊不清的字符)
negative = ((1 - gt) * mask) # negative背景位置,(1-gt是背景位置,再利用mask屏蔽不训练的背景位置。)
positive_count = int(positive.float().sum()) # 前景面积
negative_count = min(
int(negative.float().sum()),
int(positive_count * self.negative_ratio)) # 背景面积。(利用negative_ratio,背景区域不能大于前景的3倍)
assert gt.max() <= 1 and gt.min() >= 0
assert pred.max() <= 1 and pred.min() >= 0
loss = F.binary_cross_entropy(pred, gt, reduction='none') # 所有位置的loss: (b, w, h)
positive_loss = loss * positive.float() # 前景位置的loss
negative_loss = loss * negative.float() # 背景位置的loss
negative_loss, _ = torch.topk(negative_loss.view(-1), negative_count) # 只取前negative_count个背景位置的loss
balance_loss = (positive_loss.sum() + negative_loss.sum()) / (
positive_count + negative_count + self.eps) # 最后把所有位置的loss相加,求平均。
return balance_loss
2. dice loss
和前面bce loss输入的label是一样的。 前面的bce loss同时关注前后景,dice_loss只关注前景。
dice_loss(pred=pred_db, gt=gt_shrink, mask=gt_shrink_mask)


def forward(self, pred, target, mask=None):
"""
pred: (b, w,h) 预测分数image, (0, 1)之间;
target: (b, w,h) label image, 值0或者1;
mask:(b, w,h) 屏蔽背景位置,其中屏蔽位置为0,保留位置为1
"""
pred = pred.contiguous().view(pred.size()[0], -1) # (b, w*h)
target = target.contiguous().view(target.size()[0], -1) # (b, w*h)
if mask is not None:
mask = mask.contiguous().view(mask.size()[0], -1) # (b, w*h)
pred = pred * mask # 利用mask,将不需要参与训练的预测值置零
target = target * mask # 利用mask,将不需要参与训练的label值置零
a = torch.sum(pred * target) # 这里往后就是dice loss公式了: 1 - 2|A*B| / ( |A| + |B|)
b = torch.sum(pred)
c = torch.sum(target)
d = (2 * a) / (b + c + self.eps)
return 1 - d
3. l1 loss
求绝对值距离。
l1_th_loss(pred=pred_thr, gt=gt_thr, mask=gt_thr_mask)


def l1_thr_loss(self, pred, gt, mask):
"""
l1: pred与gt差值绝对值的平均值。
Args:
pred: torch.Size([8, 640, 640]).
gt: torch.Size([8, 640, 640]).
mask: torch.Size([8, 640, 640]). 只计算mask位置的loss:mask位置的预测值与真实值之间差值绝对值的平均值。
Returns:
"""
thr_loss = torch.abs((pred - gt) * mask).sum() / (
mask.sum() + self.eps)
return thr_loss
更多推荐



所有评论(0)