(nx)
| 290 | |
| 291 | |
| 292 | def test_coot_log(nx): |
| 293 | n_samples = 90 # nb samples |
| 294 | |
| 295 | mu_s = np.array([-2, 0]) |
| 296 | cov_s = np.array([[1, 0], [0, 1]]) |
| 297 | |
| 298 | xs = ot.datasets.make_2D_samples_gauss(n_samples, mu_s, cov_s, random_state=43) |
| 299 | xt = xs[::-1].copy() |
| 300 | xs_nx = nx.from_numpy(xs) |
| 301 | xt_nx = nx.from_numpy(xt) |
| 302 | |
| 303 | pi_sample, pi_feature, log = coot(X=xs, Y=xt, log=True) |
| 304 | pi_sample_nx, pi_feature_nx, log_nx = coot(X=xs_nx, Y=xt_nx, log=True) |
| 305 | |
| 306 | duals_sample, duals_feature = log["duals_sample"], log["duals_feature"] |
| 307 | assert len(duals_sample) == 2 |
| 308 | assert len(duals_feature) == 2 |
| 309 | assert len(duals_sample[0]) == n_samples |
| 310 | assert len(duals_sample[1]) == n_samples |
| 311 | assert len(duals_feature[0]) == 2 |
| 312 | assert len(duals_feature[1]) == 2 |
| 313 | |
| 314 | duals_sample_nx = log_nx["duals_sample"] |
| 315 | assert len(duals_sample_nx) == 2 |
| 316 | assert len(duals_sample_nx[0]) == n_samples |
| 317 | assert len(duals_sample_nx[1]) == n_samples |
| 318 | |
| 319 | duals_feature_nx = log_nx["duals_feature"] |
| 320 | assert len(duals_feature_nx) == 2 |
| 321 | assert len(duals_feature_nx[0]) == 2 |
| 322 | assert len(duals_feature_nx[1]) == 2 |
| 323 | |
| 324 | list_coot = log["distances"] |
| 325 | assert len(list_coot) >= 1 |
| 326 | |
| 327 | list_coot_nx = log_nx["distances"] |
| 328 | assert len(list_coot_nx) >= 1 |
| 329 | |
| 330 | # test with coot distance |
| 331 | coot_np, log = coot2(X=xs, Y=xt, log=True) |
| 332 | coot_nx, log_nx = coot2(X=xs_nx, Y=xt_nx, log=True) |
| 333 | |
| 334 | duals_sample, duals_feature = log["duals_sample"], log["duals_feature"] |
| 335 | assert len(duals_sample) == 2 |
| 336 | assert len(duals_feature) == 2 |
| 337 | assert len(duals_sample[0]) == n_samples |
| 338 | assert len(duals_sample[1]) == n_samples |
| 339 | assert len(duals_feature[0]) == 2 |
| 340 | assert len(duals_feature[1]) == 2 |
| 341 | |
| 342 | duals_sample_nx = log_nx["duals_sample"] |
| 343 | assert len(duals_sample_nx) == 2 |
| 344 | assert len(duals_sample_nx[0]) == n_samples |
| 345 | assert len(duals_sample_nx[1]) == n_samples |
| 346 | |
| 347 | duals_feature_nx = log_nx["duals_feature"] |
| 348 | assert len(duals_feature_nx) == 2 |
| 349 | assert len(duals_feature_nx[0]) == 2 |
nothing calls this directly
no test coverage detected