diff --git a/RELEASES.md b/RELEASES.md index 860f2babb..794fb8bdd 100644 --- a/RELEASES.md +++ b/RELEASES.md @@ -15,6 +15,7 @@ #### Closed issues +- Fix oversized `max_nz` budgets in `ot.utils.projection_sparse_simplex` so they return the documented unconstrained simplex projection for all axis modes (PR #876, Issue #875). - Allow `NumpyBackend.seed` to adopt an existing `np.random.RandomState` instance and remove NumPy-specific random sampling paths in sliced utilities (PR #849, Issue #848) - Remove a leftover debug `print` from `ot.utils.projection_sparse_simplex` with `axis=1`, and make the `ot.datasets.make_gauss_hd` docstring a raw string so importing `ot` no longer emits a `SyntaxWarning` (PR #860) - Fix `ot.dist` ignoring the weights `w` for `metric="cityblock"`, which returned the unweighted distance although the weights are documented for this metric (PR #859) diff --git a/ot/utils.py b/ot/utils.py index 8a40515b1..eeccbcc63 100644 --- a/ot/utils.py +++ b/ot/utils.py @@ -193,6 +193,8 @@ def projection_sparse_simplex(V, max_nz, z=1, axis=None, nx=None): raise ValueError("V.ndim must be <= 2") if axis == 1: + if max_nz >= V.shape[1]: + return proj_simplex(V.T, z).T # For each row of V, find top max_nz values; arrange the # corresponding column indices such that their values are # in a descending order. diff --git a/test/test_utils.py b/test/test_utils.py index 9c93d520a..c4593dd28 100644 --- a/test/test_utils.py +++ b/test/test_utils.py @@ -142,6 +142,45 @@ def double_sort_projection_sparse_simplex(X, max_nz, z=1, axis=None): np.testing.assert_allclose(slow_sparse_proj, fast_sparse_proj) +@pytest.mark.parametrize("axis", [None, 0, 1]) +@pytest.mark.parametrize("extra_budget", [0, 2]) +def test_projection_sparse_simplex_unconstrained_budget(nx, axis, extra_budget): + values = np.array([[0.8, 0.1, -1.0], [0.4, 0.3, 0.2]]) + z = 0.5 + if axis is None: + dimension = values.size + expected = np.array([0.45, 0.0, 0.0, 0.05, 0.0, 0.0]) + elif axis == 0: + dimension = values.shape[0] + expected = np.array([[0.45, 0.15, 0.0], [0.05, 0.35, 0.5]]) + else: + dimension = values.shape[1] + expected = np.array([[0.5, 0.0, 0.0], [4 / 15, 1 / 6, 1 / 15]]) + result = ot.utils.projection_sparse_simplex( + nx.from_numpy(values), dimension + extra_budget, z=z, axis=axis + ) + np.testing.assert_allclose(nx.to_numpy(result), expected, atol=1e-12) + + +@pytest.mark.parametrize("axis", [0, 1]) +@pytest.mark.parametrize("extra_budget", [0, 2]) +def test_projection_sparse_simplex_unconstrained_vector_mass(nx, axis, extra_budget): + values = np.array([[0.8, 0.1, -1.0], [0.4, 0.3, 0.2]]) + if axis == 0: + mass = np.array([0.5, 1.0, 0.2]) + expected = np.array([[0.45, 0.4, 0.0], [0.05, 0.6, 0.2]]) + else: + mass = np.array([0.5, 1.0]) + expected = np.array([[0.5, 0.0, 0.0], [13 / 30, 1 / 3, 7 / 30]]) + result = ot.utils.projection_sparse_simplex( + nx.from_numpy(values), + values.shape[axis] + extra_budget, + z=nx.from_numpy(mass), + axis=axis, + ) + np.testing.assert_allclose(nx.to_numpy(result), expected, atol=1e-12) + + def test_parmap(): n = 10