From 93cb275f9dccc55a4f630fd3d0ca3dadd650a2ee Mon Sep 17 00:00:00 2001 From: Aaron Smull Date: Thu, 20 Aug 2026 13:31:13 -0700 Subject: [PATCH 1/3] Fixing exporting of dataset to xarray with permuted data axis. Fixing ordering of data when using permuted data axes Revert "Updating documentation" This reverts commit 60a47357c4f48075eafbc3ac660e969a35134cba. Fixing linting --- .../dataset/exporters/export_to_xarray.py | 136 +++++++++++++----- .../dataset/test_inferred_multiple_parents.py | 16 +-- ...st_parameter_with_setpoints_has_control.py | 4 +- 3 files changed, 111 insertions(+), 45 deletions(-) diff --git a/src/qcodes/dataset/exporters/export_to_xarray.py b/src/qcodes/dataset/exporters/export_to_xarray.py index 1734831dd820..041fd1c21947 100644 --- a/src/qcodes/dataset/exporters/export_to_xarray.py +++ b/src/qcodes/dataset/exporters/export_to_xarray.py @@ -67,6 +67,7 @@ def _add_inferred_data_vars( name: str, sub_dict: Mapping[str, npt.NDArray], xr_dataset: xr.Dataset, + index: pd.Index | pd.MultiIndex | None, ) -> xr.Dataset: """Add inferred parameters as data variables to an xarray dataset. @@ -74,6 +75,7 @@ def _add_inferred_data_vars( and present in sub_dict but not yet in the dataset are added as data variables along the existing dimensions. """ + from pandas import Series interdeps = dataset.description.interdeps meas_paramspec = interdeps.graph.nodes[name]["value"] @@ -106,7 +108,14 @@ def _add_inferred_data_vars( expected_shape = tuple(xr_dataset.sizes[d] for d in dims) expected_size = prod(expected_shape) if flat.shape[0] == expected_size: - xr_dataset[inf.name] = (dims, flat.reshape(expected_shape)) + if index is not None: + # If an index is provided, we should align the inferred data with the index. + # This is necessary because data may be reordered when transforming from a pandas DataFrame to an xarray Dataset. + # Passing an index allows the original data ordering to be preserved on reconstruction. + indexed_data = Series(flat, index=index, name=inf.name).to_xarray() + xr_dataset[inf.name] = indexed_data + else: + xr_dataset[inf.name] = (dims, flat.reshape(expected_shape)) else: _LOG.warning( "Cannot add inferred parameter '%s' to xarray dataset for '%s' " @@ -164,7 +173,7 @@ def _load_to_xarray_dataset_dict_no_metadata( dependent_parameter=name, ).to_xarray() xr_dataset_dict[name] = _add_inferred_data_vars( - dataset, name, sub_dict, xr_dataset + dataset, name, sub_dict, xr_dataset, index ) elif index_is_unique: df = _data_to_dataframe( @@ -177,7 +186,7 @@ def _load_to_xarray_dataset_dict_no_metadata( dataset, use_multi_index, name, df, index ) xr_dataset_dict[name] = _add_inferred_data_vars( - dataset, name, sub_dict, xr_dataset + dataset, name, sub_dict, xr_dataset, index ) else: df = _data_to_dataframe( @@ -188,7 +197,7 @@ def _load_to_xarray_dataset_dict_no_metadata( ) xr_dataset = df.reset_index().to_xarray() xr_dataset_dict[name] = _add_inferred_data_vars( - dataset, name, sub_dict, xr_dataset + dataset, name, sub_dict, xr_dataset, index ) return xr_dataset_dict @@ -234,7 +243,9 @@ def _xarray_data_set_from_pandas_multi_index( def _xarray_data_set_direct( - dataset: DataSetProtocol, name: str, sub_dict: Mapping[str, npt.NDArray] + dataset: DataSetProtocol, + name: str, + sub_dict: Mapping[str, npt.NDArray], ) -> xr.Dataset: import xarray as xr @@ -242,56 +253,109 @@ def _xarray_data_set_direct( _, deps, inferred = dataset.description.interdeps.all_parameters_in_tree_by_group( meas_paramspec ) - # Build coordinate axes from direct dependencies preserving their order + + shape = sub_dict[name].shape + expected_size = prod(shape) + + if len(deps) != len(shape): + raise ValueError( + f"Parameter {name!r} has shape {shape}, but has {len(deps)} dependencies" + ) + dep_axis: dict[str, npt.NDArray] = {} - for axis, dep in enumerate(deps): - dep_array = sub_dict[dep.name] - dep_axis[dep.name] = dep_array[ - tuple(slice(None) if i == axis else 0 for i in range(dep_array.ndim)) - ] + destination = np.zeros(expected_size, dtype=np.intp) + + for dimension, dep in enumerate(deps): + dep_data = sub_dict[dep.name].ravel() + + if dep_data.size != expected_size: + raise ValueError( + f"Dependency {dep.name!r} contains {dep_data.size} values, " + f"but {expected_size} were expected" + ) + + values, first_indices, sorted_codes = np.unique( + dep_data, + return_index=True, + return_inverse=True, + equal_nan=True, + ) + + if values.size != shape[dimension]: + raise ValueError( + f"Dependency {dep.name!r} does not define an axis of length " + f"{shape[dimension]}: found {values.size} unique values" + ) + + # Restore first-seen order because np.unique sorts its output. + first_seen_order = np.argsort(first_indices) + dep_axis[dep.name] = values[first_seen_order] + + code_remapping = np.empty(values.size, dtype=np.intp) + code_remapping[first_seen_order] = np.arange(values.size) + codes = code_remapping[sorted_codes] + + # Encode the multidimensional coordinate using C-order mixed radix. + destination = destination * shape[dimension] + codes + + if not np.array_equal( + np.sort(destination), + np.arange(expected_size, dtype=np.intp), + ): + raise ValueError( + f"Dependencies for {name!r} do not form a complete Cartesian " + "grid without duplicate points" + ) + + permutation = np.argsort(destination) + + def reorder(data: npt.NDArray) -> npt.NDArray: + if data.size != expected_size: + raise ValueError( + f"Parameter contains {data.size} values, " + f"but {expected_size} were expected" + ) + return data.ravel()[permutation].reshape(shape) + + reordered_data = { + parameter_name: reorder(parameter_data) + for parameter_name, parameter_data in sub_dict.items() + } extra_coords: dict[str, tuple[tuple[str, ...], npt.NDArray]] = {} extra_data_vars: dict[str, tuple[tuple[str, ...], npt.NDArray]] = {} + for inf in inferred: - # skip parameters already used as primary coordinate axes - if inf.name in dep_axis: - continue - # add only if data for this parameter is available - if inf.name not in sub_dict: + if inf.name in dep_axis or inf.name not in reordered_data: continue inf_related = dataset.description.interdeps.find_all_parameters_in_tree(inf) - related_deps = inf_related.intersection(set(deps)) related_top_level = inf_related.intersection({meas_paramspec}) - if len(related_top_level) > 0: - # If inferred param is related to the top-level measurement parameter, - # add it as a data variable with the full dependency dimensions - inf_data_full = sub_dict[inf.name] - inf_dims_full = tuple(dep_axis.keys()) - extra_data_vars[inf.name] = (inf_dims_full, inf_data_full) + if related_top_level: + extra_data_vars[inf.name] = ( + tuple(dep_axis), + reordered_data[inf.name], + ) else: - # Otherwise, add as a coordinate along the related dependency axes only - inf_data = sub_dict[inf.name][ + inf_data = reordered_data[inf.name][ tuple(slice(None) if dep in related_deps else 0 for dep in deps) ] - inf_coords = [dep.name for dep in deps if dep in related_deps] - - extra_coords[inf.name] = (tuple(inf_coords), inf_data) + inf_dims = tuple(dep.name for dep in deps if dep in related_deps) + extra_coords[inf.name] = (inf_dims, inf_data) - # Compose coordinates dict including dependency axes and extra inferred coords - coords: dict[str, tuple[tuple[str, ...], npt.NDArray] | npt.NDArray] - coords = {**dep_axis, **extra_coords} + coords: dict[ + str, + tuple[tuple[str, ...], npt.NDArray] | npt.NDArray, + ] = {**dep_axis, **extra_coords} - # Compose data variables dict including measured var and any inferred data vars data_vars: dict[str, tuple[tuple[str, ...], npt.NDArray]] = { - name: (tuple(dep_axis.keys()), sub_dict[name]) + name: (tuple(dep_axis), reordered_data[name]), + **extra_data_vars, } - data_vars.update(extra_data_vars) - ds = xr.Dataset(data_vars, coords=coords) - return ds + return xr.Dataset(data_vars, coords=coords) def load_to_xarray_dataset_dict( diff --git a/tests/dataset/test_inferred_multiple_parents.py b/tests/dataset/test_inferred_multiple_parents.py index 770044ab3547..e13e4e08a7bf 100644 --- a/tests/dataset/test_inferred_multiple_parents.py +++ b/tests/dataset/test_inferred_multiple_parents.py @@ -70,7 +70,7 @@ def test_single_parent_matching_size_is_included(self) -> None: coords={"sp": sub_dict["sp"]}, ) - result = _add_inferred_data_vars(ds, "meas", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "meas", sub_dict, xr_ds, index=None) assert "inf_param" in result.data_vars npt.assert_array_almost_equal(result["inf_param"].values, sub_dict["inf_param"]) @@ -110,7 +110,7 @@ def test_inferred_matches_all_parents_is_included(self) -> None: coords={"sp": sub_dict["sp"]}, ) - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" in result.data_vars npt.assert_array_almost_equal(result["inf_param"].values, sub_dict["inf_param"]) @@ -155,7 +155,7 @@ def test_inferred_matches_first_parent_only(self) -> None: coords={"sp1": sub_dict["sp1"]}, ) - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) # Current behavior: included because it matches parent1 assert "inf_param" in result.data_vars @@ -206,7 +206,7 @@ def test_inferred_matches_second_parent_not_dataset_dims_warns( with caplog.at_level( logging.WARNING, logger="qcodes.dataset.exporters.export_to_xarray" ): - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" not in result.data_vars assert any( @@ -252,7 +252,7 @@ def test_inferred_matches_no_parent_warns( with caplog.at_level( logging.WARNING, logger="qcodes.dataset.exporters.export_to_xarray" ): - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" not in result.data_vars assert any( @@ -292,7 +292,7 @@ def test_matches_available_parent_ignores_missing(self) -> None: coords={"sp": sub_dict["sp"]}, ) - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" in result.data_vars npt.assert_array_almost_equal(result["inf_param"].values, sub_dict["inf_param"]) @@ -328,7 +328,7 @@ def test_warns_when_only_available_parent_mismatches( with caplog.at_level( logging.WARNING, logger="qcodes.dataset.exporters.export_to_xarray" ): - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" not in result.data_vars assert any( @@ -368,7 +368,7 @@ def test_both_parents_same_size_included(self) -> None: coords={"sp": sub_dict["sp"]}, ) - result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds) + result = _add_inferred_data_vars(ds, "parent1", sub_dict, xr_ds, index=None) assert "inf_param" in result.data_vars npt.assert_array_almost_equal(result["inf_param"].values, sub_dict["inf_param"]) diff --git a/tests/dataset/test_parameter_with_setpoints_has_control.py b/tests/dataset/test_parameter_with_setpoints_has_control.py index 08e641a89bdb..d2aae82e39f5 100644 --- a/tests/dataset/test_parameter_with_setpoints_has_control.py +++ b/tests/dataset/test_parameter_with_setpoints_has_control.py @@ -187,7 +187,9 @@ def test_parameter_with_setpoints_has_control_size_mismatch_warns( with caplog.at_level( logging.WARNING, logger="qcodes.dataset.exporters.export_to_xarray" ): - result = _add_inferred_data_vars(ds.dataset, "p2", sub_dict, xr_dataset) + result = _add_inferred_data_vars( + ds.dataset, "p2", sub_dict, xr_dataset, index=None + ) assert "p1" not in result.data_vars assert any( From b4ac0b470197dbc0040a64a063e982b6294b57ae Mon Sep 17 00:00:00 2001 From: Aaron Smull Date: Tue, 25 Aug 2026 15:42:08 -0700 Subject: [PATCH 2/3] Test coverage for direct-export xarray changes --- .../dataset/exporters/export_to_xarray.py | 2 +- tests/dataset/test_dataset_export.py | 116 +++++++++++++++++- 2 files changed, 116 insertions(+), 2 deletions(-) diff --git a/src/qcodes/dataset/exporters/export_to_xarray.py b/src/qcodes/dataset/exporters/export_to_xarray.py index 041fd1c21947..3e564dc08f38 100644 --- a/src/qcodes/dataset/exporters/export_to_xarray.py +++ b/src/qcodes/dataset/exporters/export_to_xarray.py @@ -326,7 +326,7 @@ def reorder(data: npt.NDArray) -> npt.NDArray: extra_data_vars: dict[str, tuple[tuple[str, ...], npt.NDArray]] = {} for inf in inferred: - if inf.name in dep_axis or inf.name not in reordered_data: + if inf.name not in reordered_data: continue inf_related = dataset.description.interdeps.find_all_parameters_in_tree(inf) diff --git a/tests/dataset/test_dataset_export.py b/tests/dataset/test_dataset_export.py index bf7e21911689..16dfc2155a37 100644 --- a/tests/dataset/test_dataset_export.py +++ b/tests/dataset/test_dataset_export.py @@ -35,7 +35,10 @@ from qcodes.dataset.descriptions.versioning import serialization as serial from qcodes.dataset.export_config import DataExportType from qcodes.dataset.exporters.export_to_pandas import _generate_pandas_index -from qcodes.dataset.exporters.export_to_xarray import _calculate_index_shape +from qcodes.dataset.exporters.export_to_xarray import ( + _calculate_index_shape, + _xarray_data_set_direct, +) from qcodes.dataset.linked_datasets.links import links_to_str from qcodes.parameters import ManualParameter, Parameter, ParamSpecBase @@ -176,6 +179,21 @@ def _make_mock_dataset_grid_with_shapes(experiment: Experiment) -> DataSet: return dataset +@pytest.fixture(name="direct_export_dataset") +def _make_direct_export_dataset(experiment: Experiment) -> DataSet: + dataset = new_data_set("direct_export_dataset") + xparam = ParamSpecBase("x", "numeric") + yparam = ParamSpecBase("y", "numeric") + signalparam = ParamSpecBase("signal", "numeric") + inferredparam = ParamSpecBase("inferred", "numeric") + idps = InterDependencies_( + dependencies={signalparam: (xparam, yparam)}, + inferences={inferredparam: (xparam,)}, + ) + dataset.set_interdependencies(idps, shapes={"signal": (2, 2)}) + return dataset + + @pytest.fixture(name="mock_dataset_grid_incomplete") def _make_mock_dataset_grid_incomplete(experiment: Experiment) -> DataSet: dataset = new_data_set("dataset") @@ -1586,6 +1604,102 @@ def test_multi_index_options_grid_with_shape( assert xds_always.sizes == {"multi_index": 50} +def test_export_to_xarray_dataset_permuted_grid(experiment: Experiment) -> None: + dataset = new_data_set("permuted_grid") + xparam = ParamSpecBase("x", "numeric") + yparam = ParamSpecBase("y", "numeric") + signalparam = ParamSpecBase("signal", "numeric") + idps = InterDependencies_(dependencies={signalparam: (xparam, yparam)}) + dataset.set_interdependencies(idps, shapes={"signal": (3, 3)}) + + x_values = np.array([[1, 0, 2], [1, 1, 2], [0, 2, 0]]) + y_values = np.array([[0, 0, 0], [1, 2, 1], [1, 2, 2]]) + + dataset.mark_started() + for x, y in zip(x_values.ravel(), y_values.ravel()): + dataset.add_results([{"x": x, "y": y, "signal": 10 * x + y}]) + dataset.mark_completed() + + xarray_dataset = dataset.to_xarray_dataset() + + assert_array_equal(xarray_dataset.coords["x"], [1, 0, 2]) + assert_array_equal(xarray_dataset.coords["y"], [0, 1, 2]) + assert_array_equal( + xarray_dataset["signal"], + np.array([[10, 11, 12], [0, 1, 2], [20, 21, 22]]), + ) + + +@pytest.mark.parametrize( + ("data", "error"), + [ + ( + { + "signal": np.arange(4), + "x": np.array([0, 0, 1, 1]), + "y": np.array([0, 1, 0, 1]), + }, + "has shape .* but has 2 dependencies", + ), + ( + { + "signal": np.arange(4).reshape(2, 2), + "x": np.array([0, 0, 1]), + "y": np.array([[0, 1], [0, 1]]), + }, + "Dependency 'x' contains 3 values, but 4 were expected", + ), + ( + { + "signal": np.arange(4).reshape(2, 2), + "x": np.array([[0, 1], [2, 3]]), + "y": np.array([[0, 1], [0, 1]]), + }, + "Dependency 'x' does not define an axis of length 2", + ), + ( + { + "signal": np.arange(4).reshape(2, 2), + "x": np.array([[0, 0], [1, 1]]), + "y": np.array([[0, 0], [1, 1]]), + }, + "do not form a complete Cartesian grid", + ), + ( + { + "signal": np.arange(4).reshape(2, 2), + "x": np.array([[0, 0], [1, 1]]), + "y": np.array([[0, 1], [0, 1]]), + "inferred": np.arange(3), + }, + "Parameter contains 3 values, but 4 were expected", + ), + ], +) +def test_xarray_data_set_direct_rejects_invalid_grid( + direct_export_dataset: DataSet, + data: dict[str, np.ndarray], + error: str, +) -> None: + with pytest.raises(ValueError, match=error): + _xarray_data_set_direct(direct_export_dataset, "signal", data) + + +def test_xarray_data_set_direct_skips_missing_inferred_data( + direct_export_dataset: DataSet, +) -> None: + data = { + "signal": np.arange(4).reshape(2, 2), + "x": np.array([[0, 0], [1, 1]]), + "y": np.array([[0, 1], [0, 1]]), + } + + xarray_dataset = _xarray_data_set_direct(direct_export_dataset, "signal", data) + + assert set(xarray_dataset.coords) == {"x", "y"} + assert set(xarray_dataset.data_vars) == {"signal"} + + def test_multi_index_options_incomplete_grid( mock_dataset_grid_incomplete: DataSet, ) -> None: From 5aea2ca76fff9a6960d242f7e013df887932ead8 Mon Sep 17 00:00:00 2001 From: "Jens H. Nielsen" Date: Thu, 27 Aug 2026 10:46:12 +0200 Subject: [PATCH 3/3] Fix inferred parameter export when index cannot be expanded `Series(...).to_xarray()` expands a pandas index into one dimension per index level. That is only compatible with the target xarray dataset when the dataset dimensions are exactly the levels of the index. It is not the case when the dataset uses a single `multi_index` dimension or the flat index created by `DataFrame.reset_index()`, and a non unique MultiIndex cannot be converted at all. Only take the reindexing path when the index is unique and the dataset dimensions match the index level names. In the remaining cases the data is already in index order and can be reshaped directly. Adds tests covering export of an inferred parameter with a `multi_index` dimension and with a non unique MultiIndex. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 92aac34d-1a70-4ca9-900c-4ab563a731df --- .../dataset/exporters/export_to_xarray.py | 13 ++- tests/dataset/test_dataset_export.py | 84 +++++++++++++++++++ 2 files changed, 96 insertions(+), 1 deletion(-) diff --git a/src/qcodes/dataset/exporters/export_to_xarray.py b/src/qcodes/dataset/exporters/export_to_xarray.py index 3e564dc08f38..528aca68c7da 100644 --- a/src/qcodes/dataset/exporters/export_to_xarray.py +++ b/src/qcodes/dataset/exporters/export_to_xarray.py @@ -84,6 +84,17 @@ def _add_inferred_data_vars( dep_names = {dep.name for dep in deps} dims = tuple(d for d in xr_dataset.dims) + # ``Series.to_xarray()`` expands the index into one dimension per index + # level. That is only compatible with the target dataset if the dataset + # dimensions are exactly the levels of the index. It is not the case when + # the dataset uses a single ``multi_index`` dimension or the flat index + # created by ``DataFrame.reset_index()``. A non unique MultiIndex cannot be + # converted at all. In those cases the data is already in index order so it + # can be reshaped directly. + index_is_expandable = ( + index is not None and index.is_unique and set(dims) == set(index.names) + ) + for inf in inferred: if inf.name in dep_names: continue @@ -108,7 +119,7 @@ def _add_inferred_data_vars( expected_shape = tuple(xr_dataset.sizes[d] for d in dims) expected_size = prod(expected_shape) if flat.shape[0] == expected_size: - if index is not None: + if index is not None and index_is_expandable: # If an index is provided, we should align the inferred data with the index. # This is necessary because data may be reordered when transforming from a pandas DataFrame to an xarray Dataset. # Passing an index allows the original data ordering to be preserved on reconstruction. diff --git a/tests/dataset/test_dataset_export.py b/tests/dataset/test_dataset_export.py index 16dfc2155a37..217dc94c99e3 100644 --- a/tests/dataset/test_dataset_export.py +++ b/tests/dataset/test_dataset_export.py @@ -305,6 +305,62 @@ def _make_mock_dataset_non_grid(experiment: Experiment) -> DataSet: return dataset +@pytest.fixture(name="mock_dataset_non_grid_inferred") +def _make_mock_dataset_non_grid_inferred(experiment: Experiment) -> DataSet: + """Non grid dataset where an inferred parameter is inferred from z.""" + dataset = new_data_set("dataset") + xparam = ParamSpecBase("x", "numeric") + yparam = ParamSpecBase("y", "numeric") + zparam = ParamSpecBase("z", "numeric") + tparam = ParamSpecBase("t", "numeric") + idps = InterDependencies_( + dependencies={zparam: (xparam, yparam)}, + inferences={tparam: (zparam,)}, + ) + dataset.set_interdependencies(idps) + + num_samples = 50 + + rng = np.random.default_rng(1234) + + x_vals = rng.random(num_samples) * 10 + y_vals = 20 + rng.random(num_samples) * 5 + + dataset.mark_started() + + for i, (x, y) in enumerate(zip(x_vals, y_vals)): + dataset.add_results([{"x": x, "y": y, "z": x + y, "t": float(i)}]) + dataset.mark_completed() + return dataset + + +@pytest.fixture(name="mock_dataset_non_unique_index_inferred") +def _make_mock_dataset_non_unique_index_inferred(experiment: Experiment) -> DataSet: + """Dataset with a non unique MultiIndex and an inferred parameter.""" + dataset = new_data_set("dataset") + xparam = ParamSpecBase("x", "numeric") + yparam = ParamSpecBase("y", "numeric") + zparam = ParamSpecBase("z", "numeric") + tparam = ParamSpecBase("t", "numeric") + idps = InterDependencies_( + dependencies={zparam: (xparam, yparam)}, + inferences={tparam: (zparam,)}, + ) + dataset.set_interdependencies(idps) + + num_samples = 20 + # every (x, y) pair is measured twice making the index non unique + x_vals = np.repeat(np.arange(num_samples // 2, dtype=float), 2) + y_vals = np.repeat(np.arange(num_samples // 2, dtype=float), 2) + + dataset.mark_started() + + for i, (x, y) in enumerate(zip(x_vals, y_vals)): + dataset.add_results([{"x": x, "y": y, "z": x + y, "t": float(i)}]) + dataset.mark_completed() + return dataset + + @pytest.fixture(name="mock_dataset_non_grid_in_mem") def _make_mock_dataset_non_grid_in_mem(experiment: Experiment) -> DataSetProtocol: meas = Measurement(exp=experiment, name="in_mem_ds") @@ -1760,6 +1816,34 @@ def test_multi_index_options_non_grid(mock_dataset_non_grid: DataSet) -> None: assert xds_always.sizes == {"multi_index": 50} +@pytest.mark.parametrize("use_multi_index", ["auto", "always"]) +def test_multi_index_export_with_inferred_parameter( + mock_dataset_non_grid_inferred: DataSet, use_multi_index: str +) -> None: + """Inferred parameters must export correctly when a MultiIndex dim is used.""" + xds = mock_dataset_non_grid_inferred.to_xarray_dataset( + use_multi_index=use_multi_index # pyright: ignore[reportArgumentType] + ) + + assert xds.sizes == {"multi_index": 50} + assert "t" in xds.data_vars + assert xds["t"].dims == ("multi_index",) + np.testing.assert_array_equal(xds["t"].values, np.arange(50, dtype=float)) + + +def test_non_unique_multi_index_export_with_inferred_parameter( + mock_dataset_non_unique_index_inferred: DataSet, +) -> None: + """A non unique MultiIndex must not break export of inferred parameters.""" + xds = mock_dataset_non_unique_index_inferred.to_xarray_dataset() + + assert "t" in xds.data_vars + assert xds["t"].dims == xds["z"].dims + np.testing.assert_array_equal( + np.asarray(xds["t"].values).ravel(), np.arange(20, dtype=float) + ) + + def test_multi_index_wrong_option(mock_dataset_non_grid: DataSet) -> None: with pytest.raises(ValueError, match="Invalid value for use_multi_index"): mock_dataset_non_grid.to_xarray_dataset(use_multi_index=True) # pyright: ignore[reportArgumentType]