target_pos = target.bool() target_neg = ~target_pos # loss_pos and loss_neg below contain non-zero values only for those elements # that are positive pairs and negative pairs respectively. loss_pos = torch.zeros(x.size(0), x.size(0)).masked_scatter(target_pos, loss[target_pos]) loss_neg = torch.zeros(x.size(0), x.size(0)).masked_scatter(target_neg, loss[target_neg])