(transp)
| 1039 | return W |
| 1040 | |
| 1041 | def vectorized(transp): |
| 1042 | labels_u, labels_idx = nx.unique(labels_a, return_inverse=True) |
| 1043 | n_labels = labels_u.shape[0] |
| 1044 | unroll_labels_idx = nx.eye(n_labels, type_as=transp)[labels_idx] |
| 1045 | W = ( |
| 1046 | nx.repeat(transp.T[:, :, None], n_labels, axis=2) |
| 1047 | * unroll_labels_idx[None, :, :] |
| 1048 | ) |
| 1049 | W = nx.sum(W, axis=1) |
| 1050 | W = p * ((W + epsilon) ** (p - 1)) |
| 1051 | W = nx.dot(W, unroll_labels_idx.T) |
| 1052 | return W.T |
| 1053 | |
| 1054 | assert np.allclose(unvectorized(T), vectorized(T)) |