(batch_preds, batch_std_labels, sigma=None)
| 5 | import torch.nn.functional as F |
| 6 | |
| 7 | def get_pairwise_comp_probs(batch_preds, batch_std_labels, sigma=None): |
| 8 | batch_s_ij = torch.unsqueeze(batch_preds, dim=2) - torch.unsqueeze(batch_preds, dim=1) |
| 9 | batch_p_ij = torch.sigmoid(sigma * batch_s_ij) |
| 10 | |
| 11 | batch_std_diffs = torch.unsqueeze(batch_std_labels, dim=2) - torch.unsqueeze(batch_std_labels, dim=1) |
| 12 | batch_Sij = torch.clamp(batch_std_diffs, min=-1.0, max=1.0) |
| 13 | batch_std_p_ij = 0.5 * (1.0 + batch_Sij) |
| 14 | |
| 15 | return batch_p_ij, batch_std_p_ij |
| 16 | |
| 17 | def rankloss(score, metric,mask=None,sigma=None): |
| 18 | gtrank = (-metric).argsort().argsort().float() |
no outgoing calls
no test coverage detected