MCPcopy Create free account
hub / github.com/ActiveVisionLab/DFNet / masked_loss

Function masked_loss

script/feature/misc.py:332–353  ·  view source on GitHub ↗

compute loss only in masked region :param criterion: loss function :param f1: [3, batch_size, H, W] :param f2: [3, batch_size, H, W] :param valid_mask: [batch_size, H, W] :return: loss

(criterion, f1, f2, valid_mask)

Source from the content-addressed store, hash-verified

330 return x
331
332def masked_loss(criterion, f1, f2, valid_mask):
333 '''
334 compute loss only in masked region
335 :param criterion: loss function
336 :param f1: [3, batch_size, H, W]
337 :param f2: [3, batch_size, H, W]
338
339 :param valid_mask: [batch_size, H, W]
340 :return:
341 loss
342 '''
343 # apply mask on f1
344 f1 = f1 * (valid_mask)
345 f2 = f2
346
347 # # print masked img features f1, and unwarped features f2
348 # plot_features(f1, 'f1.png', False)
349 # plot_features(f2, 'f2.png', False)
350 loss = criterion(f1, f2)
351 loss = (loss * valid_mask).sum() # take sum of loss
352 loss = loss / (valid_mask.sum()) # take mean of loss
353 return loss
354
355def triplet_loss(f1, f2, margin=1.):
356 '''

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected