INNER CODE UNIT · Python
filler
nv-tlabs/GSCNN · loss.py:96
filler = torch.ones_like(target) * 255
return self.seg_loss(input,
torch.where(edge.max(1)[0] > 0.8, target, filler))
def forward(self, inputs, targets):
segin, edgein = inputs
segmask, edgemask = targets
losses = {}
losses['seg_loss'] = self.seg_weight * self.seg_loss(segin, segmask)
losses['edge_loss'] = self.edge_weight * 20 * self.bce2d(edgein, edgemask)
losses['att_loss'] = self.att_weight * self.edge_attention(segin, segmask, edgein)
losses['dual_loss'] = self.dual_weight * self.dual_task(segin, segmask)
return losses
#Img Weighted Loss