Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
147 changes: 111 additions & 36 deletions src/qcodes/dataset/exporters/export_to_xarray.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,13 +67,15 @@ def _add_inferred_data_vars(
name: str,
sub_dict: Mapping[str, npt.NDArray],
xr_dataset: xr.Dataset,
index: pd.Index | pd.MultiIndex | None,
Comment thread
asmull marked this conversation as resolved.
) -> xr.Dataset:
"""Add inferred parameters as data variables to an xarray dataset.

Parameters that are inferred from the top-level measurement parameter
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"]
Expand All @@ -82,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
Expand All @@ -106,7 +119,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 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.
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:
_LOGGER.warning(
"Cannot add inferred parameter '%s' to xarray dataset for '%s' "
Expand Down Expand Up @@ -164,7 +184,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(
Expand All @@ -177,7 +197,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(
Expand All @@ -188,7 +208,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
Expand Down Expand Up @@ -234,64 +254,119 @@ 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

meas_paramspec = dataset.description.interdeps.graph.nodes[name]["value"]
_, 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()

Comment thread
asmull marked this conversation as resolved.
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 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)
Comment thread
asmull marked this conversation as resolved.


def load_to_xarray_dataset_dict(
Expand Down
Loading
Loading