| 84 | |
| 85 | @torch.no_grad() |
| 86 | def sinkhorn_algorithm(distances, epsilon, sinkhorn_iterations): |
| 87 | Q = torch.exp(- distances / epsilon) |
| 88 | |
| 89 | B = Q.shape[0] # number of samples to assign |
| 90 | K = Q.shape[1] # how many centroids per block (usually set to 256) |
| 91 | |
| 92 | # make the matrix sums to 1 |
| 93 | sum_Q = Q.sum(-1, keepdim=True).sum(-2, keepdim=True) |
| 94 | Q /= sum_Q |
| 95 | # print(Q.sum()) |
| 96 | for it in range(sinkhorn_iterations): |
| 97 | |
| 98 | # normalize each column: total weight per sample must be 1/B |
| 99 | Q /= torch.sum(Q, dim=1, keepdim=True) |
| 100 | Q /= B |
| 101 | |
| 102 | # normalize each row: total weight per prototype must be 1/K |
| 103 | Q /= torch.sum(Q, dim=0, keepdim=True) |
| 104 | Q /= K |
| 105 | |
| 106 | |
| 107 | Q *= B # the colomns must sum to 1 so that Q is an assignment |
| 108 | return Q |