(X1, X2, w1, w2)
| 117 | w = nx.ones(r, type_as=Xa) / r |
| 118 | |
| 119 | def solve_ot(X1, X2, w1, w2): |
| 120 | M = dist(X1, X2) |
| 121 | if reg > 0: |
| 122 | G, log = sinkhorn(w1, w2, M, reg, log=True, **kwargs) |
| 123 | log["cost"] = nx.sum(G * M) |
| 124 | return G, log |
| 125 | else: |
| 126 | return emd(w1, w2, M, log=True, **kwargs) |
| 127 | |
| 128 | norm_delta = [] |
| 129 |
no test coverage detected