MCPcopy Create free account
hub / github.com/AkaliKong/MiniOneRec / sinkhorn_algorithm

Function sinkhorn_algorithm

rq/models/layers.py:86–108  ·  view source on GitHub ↗
(distances, epsilon, sinkhorn_iterations)

Source from the content-addressed store, hash-verified

84
85@torch.no_grad()
86def 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

Callers 1

forwardMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected