diff --git a/CHANGELOG.md b/CHANGELOG.md index 7b3482e39..81e6f8f3a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,6 +26,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Allow JAX and JAXlib 0.11 in downstream environments by removing the `<0.11` dependency bounds ([#801](https://github.com/QuantClimate/GPJax/issues/801)). +### Fixed + +- Read the observation count from `y` for unsupervised `Dataset` objects, so + `n` and `full_size` work when `X` is absent. + ## [1.0.0] — 2026-09-28 ### Added diff --git a/gpjax/dataset.py b/gpjax/dataset.py index e98cad435..e18097f34 100644 --- a/gpjax/dataset.py +++ b/gpjax/dataset.py @@ -75,7 +75,7 @@ def __add__(self, other: "Dataset") -> "Dataset": @property def n(self) -> int: r"""Number of observations.""" - return self.X.shape[0] + return self.X.shape[0] if self.X is not None else self.y.shape[0] @property def full_size(self) -> int: diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 00488398b..a5471b1f7 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -62,6 +62,17 @@ def test_dataset_init(n: int, in_dim: int) -> None: assert jtu.tree_leaves(D) == [x, y] +@pytest.mark.parametrize("n", [0, 1, 3]) +@pytest.mark.parametrize("n_outputs", [1, 2]) +def test_unsupervised_dataset_observation_count(n, n_outputs): + data = Dataset(y=jnp.ones((n, n_outputs))) + assert data.is_unsupervised() + assert data.n == n + assert data.full_size == n + leaves, treedef = jtu.tree_flatten(data) + assert jtu.tree_unflatten(treedef, leaves).n == n + + @pytest.mark.parametrize("n1", [1, 2, 10]) @pytest.mark.parametrize("n2", [1, 2, 10]) @pytest.mark.parametrize("in_dim", [1, 2, 10])