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

Function triplet_loss_hard_negative_mining_plus

script/feature/misc.py:399–435  ·  view source on GitHub ↗

triplet loss with hard negative mining, four cases. inspired by http://www.bmva.org/bmvc/2016/papers/paper119/paper119.pdf section3.3 :param criterion: loss function :param f1: [lvl, B, C, H, W] :param f2: [lvl, B, C, H, W] :return: loss

(f1, f2, margin=1.)

Source from the content-addressed store, hash-verified

397 return loss
398
399def triplet_loss_hard_negative_mining_plus(f1, f2, margin=1.):
400 '''
401 triplet loss with hard negative mining, four cases. inspired by http://www.bmva.org/bmvc/2016/papers/paper119/paper119.pdf section3.3
402 :param criterion: loss function
403 :param f1: [lvl, B, C, H, W]
404 :param f2: [lvl, B, C, H, W]
405 :return:
406 loss
407 '''
408 criterion = nn.TripletMarginLoss(margin=margin, reduction='mean')
409 anchor = f1
410 anchor_negative = torch.roll(f1, shifts=1, dims=1)
411 positive = f2
412 negative = torch.roll(f2, shifts=1, dims=1)
413
414 # select in-triplet hard negative, reference: section3.3
415 mse = nn.MSELoss(reduction='mean')
416 with torch.no_grad():
417 case1 = mse(anchor, negative)
418 case2 = mse(positive, anchor_negative)
419 case3 = mse(anchor, anchor_negative)
420 case4 = mse(positive, negative)
421 distance_list = torch.stack([case1,case2,case3,case4])
422 loss_case = torch.argmin(distance_list)
423
424 # perform anchor swap if necessary
425 if loss_case == 0:
426 loss = criterion(anchor, positive, negative)
427 elif loss_case == 1:
428 loss = criterion(positive, anchor, anchor_negative)
429 elif loss_case == 2:
430 loss = criterion(anchor, positive, anchor_negative)
431 elif loss_case == 3:
432 loss = criterion(positive, anchor, negative)
433 else:
434 raise NotImplementedError
435 return loss
436
437def perturb_rotation(c2w, theta, phi, psi=0):
438 last_row = np.array([[0, 0, 0, 1]])# np.tile(np.array([0, 0, 0, 1]), (1, 1)) # (N_images, 1, 4)

Callers 2

train_on_batchFunction · 0.85

Calls 1

mseFunction · 0.85

Tested by

no test coverage detected