(self)
| 95 | return (a.unsqueeze(2) - b.unsqueeze(1)).norm(dim=3).pow(2) |
| 96 | |
| 97 | def getProbas(self): |
| 98 | global ndatas, n_nfeat |
| 99 | # compute squared dist to centroids [n_runs][n_samples][n_ways] |
| 100 | if self.distance_metric == 'cosine': |
| 101 | dist = 1-torch.bmm(F.normalize(ndatas), F.normalize(self.mus.transpose(1, 2))) |
| 102 | elif self.distance_metric == 'ce': |
| 103 | dist = -torch.bmm(torch.log(ndatas + 1e-5), self.mus.transpose(1, 2)) |
| 104 | else: |
| 105 | dist = self._pairwise_dist(ndatas, self.mus) |
| 106 | |
| 107 | p_xj = torch.zeros_like(dist) |
| 108 | r = torch.ones(n_runs, n_usamples, device='cuda') |
| 109 | c = torch.ones(n_runs, n_ways, device='cuda') * n_queries |
| 110 | p_xj_test = self.compute_optimal_transport(dist[:, n_lsamples:], r, c, epsilon=1e-4) |
| 111 | p_xj[:, n_lsamples:] = p_xj_test |
| 112 | |
| 113 | p_xj[:, :n_lsamples].fill_(0) |
| 114 | p_xj[:, :n_lsamples].scatter_(2, labels[:, :n_lsamples].unsqueeze(2), 1) |
| 115 | |
| 116 | return p_xj |
| 117 | |
| 118 | def estimateFromMask(self, mask): |
| 119 | emus = mask.permute(0, 2, 1).matmul(ndatas).div(mask.sum(dim=1).unsqueeze(2)) |
no test coverage detected