MCPcopy Create free account
hub / github.com/DanielShalam/BPA / getProbas

Method getProbas

methods/pt_map/test_standard.py:97–116  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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))

Callers 2

performEpochMethod · 0.45
loopMethod · 0.45

Calls 2

_pairwise_distMethod · 0.95

Tested by

no test coverage detected