setu.sliced_wasserstein#

setu.sliced_wasserstein(samples_a, samples_b, n_projections=100, seed=42)[source]#

Sliced Wasserstein distance.

Projects samples onto random 1D directions and computes the average 1D Wasserstein distance across projections.

Parameters:
  • samples_a (Array) – First set of samples (n, dim)

  • samples_b (Array) – Second set of samples (m, dim)

  • n_projections (int) – Number of random projections

  • seed (int) – Random seed for reproducibility

Return type:

float

Returns:

Sliced Wasserstein distance (lower = more similar)