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)
| 330 | return x |
| 331 | |
| 332 | def 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 | |
| 355 | def triplet_loss(f1, f2, margin=1.): |
| 356 | ''' |
nothing calls this directly
no outgoing calls
no test coverage detected