(x, y)
| 2 | |
| 3 | |
| 4 | def pairwise_dist(x, y): |
| 5 | xx, yy, zz = torch.mm(x, x.t()), torch.mm(y, y.t()), torch.mm(x, y.t()) |
| 6 | rx = xx.diag().unsqueeze(0).expand_as(xx) |
| 7 | ry = yy.diag().unsqueeze(0).expand_as(yy) |
| 8 | P = rx.t() + ry - 2 * zz |
| 9 | return P |
| 10 | |
| 11 | |
| 12 | def NN_loss(x, y, dim=0): |