MCPcopy Create free account
hub / github.com/PythonOT/POT / test_coot_log

Function test_coot_log

test/test_coot.py:292–356  ·  view source on GitHub ↗
(nx)

Source from the content-addressed store, hash-verified

290
291
292def 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

Callers

nothing calls this directly

Calls 2

from_numpyMethod · 0.80
copyMethod · 0.45

Tested by

no test coverage detected