diff --git a/src/cyclebane/graph.py b/src/cyclebane/graph.py index 7b3a930..78a1656 100644 --- a/src/cyclebane/graph.py +++ b/src/cyclebane/graph.py @@ -10,6 +10,7 @@ import networkx as nx from .node_values import IndexName, IndexValue, NodeValues +from .value_array import Grouping def _get_unique_sink(graph: nx.DiGraph) -> Hashable: @@ -111,6 +112,12 @@ def _node_with_indices(node: Hashable, indices: tuple[IndexName, ...]) -> Mapped return MappedNode(name=node, indices=indices) +def _node_name(node: Hashable) -> Hashable: + if isinstance(node, MappedNode): + return node.name + return node + + def _node_indices(node: Hashable) -> tuple[IndexName, ...]: if isinstance(node, MappedNode): return node.indices @@ -265,6 +272,9 @@ def map(self, node_values: MappingToArrayLike) -> Graph: node_values=self._node_values.merge(new_values), ) + def groupby(self, node: Hashable) -> GroupbyGraph: + return GroupbyGraph(self.graph, node_values=self._node_values, node=node) + def reduce( self, key: None | Hashable = None, @@ -273,6 +283,7 @@ def reduce( axis: None | int = None, name: None | Hashable = None, attrs: None | dict[str, Any] = None, + _extra_index_name: None | IndexName = None, ) -> Graph: """ Reduce over the given index or axis previously created with :py:meth:`map`. @@ -314,6 +325,11 @@ def reduce( new_index = tuple(value for i, value in enumerate(indices) if i != axis) else: new_index = None + if _extra_index_name is not None: + if new_index is None: + new_index = (_extra_index_name,) + else: + new_index = (*new_index, _extra_index_name) if name in self.graph: raise ValueError(f"Node '{name}' already exists in the graph.") @@ -324,7 +340,9 @@ def reduce( return Graph(graph, node_values=self._node_values) - def _from_orig_key(self, key: Hashable) -> Hashable: + def _from_orig_key( + self, key: Hashable, match_index: None | Hashable = None + ) -> Hashable: # Graph.map relabels nodes to include index names, which can be inconvenient # for the user. Is this convenience of finding the node by its original name # worth the complexity and a good idea? @@ -334,6 +352,8 @@ def _from_orig_key(self, key: Hashable) -> Hashable: for node in self.graph.nodes if isinstance(node, MappedNode) and node.name == key ] + if match_index is not None: + matches = [node for node in matches if match_index in node.indices] if len(matches) == 0: raise KeyError(f"Node '{key}' does not exist in the graph.") if len(matches) > 1: @@ -354,23 +374,37 @@ def to_networkx(self, value_attr: str = 'value') -> nx.DiGraph: value_attr: The name of the attribute on nodes that holds the array-like object. """ - graph = self.graph + graph = self.graph.copy() + + # Maintain a list of actual node values, without groupings, since we only want + # to set the former (user-provided) on (input) nodes. + node_values = self._node_values.copy() + groupby_graphs = [] + # Handle groupby/reduce operations. The regular iterative node duplication does + # not work in this case. We have to handle the graph edges that correspond to + # a particular groupby/reduce operation in isolation, or else we get broken + # result in the presence of multiple (chained or not) groupby operations. The + # resulting graphs that correspond to the grouping are later composed with the + # rest of the graph. + for key, values in self._node_values.items(): + if (grouping := values.get_grouping()) is not None: + del node_values[key] + key = self._from_orig_key(key, match_index=grouping.group_index_name) + # Note there should be only a single predecessor for the grouping node. + groupby_graph = graph.subgraph([*graph.predecessors(key), key]).copy() + # Remove edges, or the loop for the regular map/reduce will add + # all-to-all edges between these nodes + graph.remove_edges_from(groupby_graph.edges) + groupby_graphs.append(self._make_groupby_graph(grouping, groupby_graph)) + + # Handle regular map/reduce operations for index_name, index in reversed(self.indices.items()): - # Find all nodes with this index - nodes = [ - node - for node in graph.nodes() - if index_name - in _node_indices(node.name if isinstance(node, NodeName) else node) - ] - # Make a copy for each index value - graphs = [ - _rename_successors( - graph, successors=nodes, index=IndexValues((index_name,), (i,)) - ) - for i in index - ] + graphs = _clone_graph(graph, index_name, index) graph = nx.compose_all(graphs) + + if groupby_graphs: + graph = nx.compose_all([*groupby_graphs, graph]) + # Replace all MappingNodes with their name new_names = { node: NodeName(node.name.name, node.index) @@ -383,12 +417,29 @@ def to_networkx(self, value_attr: str = 'value') -> nx.DiGraph: for node in graph.nodes: if ( isinstance(node, NodeName) - and (node_values := self._node_values.get(node.name)) is not None + and (value_array := node_values.get(node.name)) is not None ): - graph.nodes[node][value_attr] = node_values.sel(node.index.to_tuple()) + graph.nodes[node][value_attr] = value_array.sel(node.index.to_tuple()) return graph + def _make_groupby_graph( + self, grouping: Grouping, groupby_graph: nx.DiGraph + ) -> nx.DiGraph: + for index_name, index in reversed(self.indices.items()): + if index_name == grouping.index_name: + continue + graphs = _clone_graph(groupby_graph, index_name, index) + if index_name == grouping.group_index_name: + subgraphs = [ + _clone_graph(group_graph, grouping.index_name, idx) + for idx, group_graph in zip(grouping.indices, graphs, strict=True) + ] + # Flatten nested list of graphs + graphs = [g for sublist in subgraphs for g in sublist] + groupby_graph = nx.compose_all(graphs) + return groupby_graph + def __getitem__(self, key: Hashable | slice) -> Graph: """ Get the branch of the graph rooted at the given node. @@ -474,4 +525,90 @@ def __setitem__(self, branch: Hashable | slice, other: Graph) -> None: # Delay setting graph until we know no step fails self._node_values = self._node_values.merge(other._node_values) + # Remove node values of the branch, if they exist + if _node_name(branch) in self._node_values: + del self._node_values[_node_name(branch)] + + # Ensure we preserve the node values of the branch, if it exists. This step is + # necessary since __setitem__ effectively renames the sink node of the input + # graph to the branch name. + if _node_name(sink) in self._node_values: + node_values = self._node_values[_node_name(sink)] + del self._node_values[_node_name(sink)] + self._node_values[_node_name(branch)] = node_values + self.graph = graph + + +class GroupbyGraph: + """ + A graph that has been grouped by a specific index. + + This is a specialized graph that is used to represent the result of a groupby + operation on a Cyclebane graph. It allows for operations on the grouped data, + such as aggregation or summarization. + """ + + # TODO Should we support a custom new dim name here, instead of using `node`? + def __init__(self, graph: nx.DiGraph, node_values: NodeValues, node: Hashable): + self._graph = graph + self._node_values = node_values + values_to_group_by = node_values[node] + self._group_index_name = node + self._index_name = values_to_group_by.index_names[0] + self._groups = values_to_group_by.group(index_name=node) + + # TODO Require specifying index!? + def reduce( + self, + key: None | Hashable = None, + *, + name: None | Hashable = None, + attrs: None | dict[str, Any] = None, + ) -> Graph: + """ + Reduce the grouped graph over the given index or axis. + + Parameters + ---------- + key: + The name of the source node to reduce. This is the original name prior to + mapping. If not given, tries to find a unique sink node. + name: + The name of the new node. If not given, a unique name is generated. + attrs: + Attributes to set on the new node(s). + """ + # Generate name here since we want to store grouping on the new "reduce" node. + name = name or _get_new_node_name(self._graph) + # Why do we store the grouping here? This works well with existing mechanisms, + # e.g., __getitem__, which needs to decided what subset of node values to keep + # when returning a subgraph. + node_values = self._node_values.merge({name: self._groups}) + graph = Graph(self._graph, node_values=node_values) + return graph.reduce( + key=key, + index=self._index_name, + name=name, + attrs=attrs, + _extra_index_name=self._group_index_name, + ) + + +def _clone_graph( + graph: nx.DiGraph, index_name: IndexName, index: Iterable[IndexValue] +) -> list[nx.DiGraph]: + # Find all nodes with this index + nodes = [ + node + for node in graph.nodes() + if index_name + in _node_indices(node.name if isinstance(node, NodeName) else node) + ] + # Make a copy for each index value + return [ + _rename_successors( + graph, successors=nodes, index=IndexValues((index_name,), (i,)) + ) + for i in index + ] diff --git a/src/cyclebane/node_values.py b/src/cyclebane/node_values.py index 0830fdb..274a81f 100644 --- a/src/cyclebane/node_values.py +++ b/src/cyclebane/node_values.py @@ -2,16 +2,11 @@ # Copyright (c) 2024 Scipp contributors (https://github.com/scipp) from __future__ import annotations -from abc import ABC, abstractmethod from collections.abc import Hashable, Iterable, Iterator, Mapping, Sequence -from types import ModuleType -from typing import TYPE_CHECKING, Any, ClassVar, TypeVar +from typing import Any, TypeVar -if TYPE_CHECKING: - import numpy - import pandas - import scipp - import xarray +from . import value_array_adapters # noqa: F401 +from .value_array import ValueArray IndexName = Hashable IndexValue = Hashable @@ -19,357 +14,6 @@ T = TypeVar('T', bound='ValueArray') -class ValueArray(ABC): - """ - Abstract base class for a series of values with an index that can be sliced. - - Used by :py:class:`NodeValues` to store the values of a given node in a graph. The - abstraction allows for the use of different data structures to store the values of - nodes in a graph, such as pandas.DataFrame, xarray.DataArray, numpy.ndarray, or - simple Python iterables. - """ - - _registry: ClassVar = [] - - def __init_subclass__(cls) -> None: - super().__init_subclass__() - ValueArray._registry.append(cls) - - @staticmethod - def from_array_like(values: Any, *, axis_zero: int = 0) -> ValueArray: - # Reversed to ensure SequenceAdapter is tried last, as it is the most general - # SequenceAdapter is defined right after this class so it is registered first - for subclass in reversed(ValueArray._registry): - if (a := subclass.try_from(values, axis_zero=axis_zero)) is not None: - return a - raise ValueError(f'Cannot create ValueArray from {values}') - - @staticmethod - @abstractmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> ValueArray | None: ... - - def __eq__(self, other: object) -> bool: - if type(self) is not type(other): - return NotImplemented - return self._equal(other) - - def __ne__(self, other: object) -> bool: - return not self == other - - @abstractmethod - def _equal(self: T, other: T) -> bool: ... - - @abstractmethod - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - """Return data by selecting from index with given name and index value.""" - - def loc(self, key: dict[IndexName, slice]) -> ValueArray: - if not all(isinstance(i, slice) for i in key.values()): - raise ValueError('ValueArray.loc only accepts slices, not integers') - if not set(key).issubset(set(self.index_names)): - raise ValueError( - f'ValueArray.loc got {key.keys()}, not a subset of {self.index_names}' - ) - return self[key] - - @abstractmethod - def __getitem__(self, key: dict[IndexName, slice]) -> ValueArray: - pass - - @property - @abstractmethod - def shape(self) -> tuple[int, ...]: - pass - - @property - @abstractmethod - def index_names(self) -> tuple[IndexName, ...]: - pass - - @property - @abstractmethod - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - pass - - -class SequenceAdapter(ValueArray): - def __init__( - self, - values: Sequence[Any], - *, - index: Iterable[IndexValue] | None = None, - axis_zero: int = 0, - ): - self._values = values - self._index = index or range(len(values)) - self._axis_zero = axis_zero - - @staticmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> SequenceAdapter | None: - return SequenceAdapter(obj, axis_zero=axis_zero) - - def _equal(self, other: SequenceAdapter) -> bool: - return ( - self._values == other._values - and self._index == other._index - and self._axis_zero == other._axis_zero - ) - - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - if len(key) != 1: - raise ValueError('SequenceAdapter only supports single index') - _, i = key[0] - return self._values[self._index.index(i)] - - def __getitem__(self, key: dict[IndexName, slice]) -> SequenceAdapter: - _, i = next(iter(key.items())) - return SequenceAdapter( - self._values[i], index=self._index[i], axis_zero=self._axis_zero - ) - - @property - def shape(self) -> tuple[int, ...]: - return (len(self._values),) - - @property - def index_names(self) -> tuple[IndexName, ...]: - return (f'dim_{self._axis_zero}',) - - @property - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - return {f'dim_{self._axis_zero}': self._index} - - -class PandasSeriesAdapter(ValueArray): - def __init__(self, series: pandas.Series, *, axis_zero: int = 0): - self._series = series - self._axis_zero = axis_zero - - @staticmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> PandasSeriesAdapter | None: - try: - import pandas - except ModuleNotFoundError: - return None - if isinstance(obj, pandas.Series): - return PandasSeriesAdapter(obj, axis_zero=axis_zero) - - def _equal(self, other: PandasSeriesAdapter) -> bool: - return ( - self._series.equals(other._series) and self._axis_zero == other._axis_zero - ) - - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - if len(key) != 1: - raise ValueError('PandasSeriesAdapter only supports single index') - index_name, i = key[0] - if index_name != self.index_names[0]: - raise ValueError( - f'Unexpected index name {index_name} for PandasSeriesAdapter with ' - f'index names {self.index_names}' - ) - return self._series.loc[i] - - def __getitem__(self, key: dict[IndexName, slice]) -> PandasSeriesAdapter: - _, i = next(iter(key.items())) - return PandasSeriesAdapter(self._series[i], axis_zero=self._axis_zero) - - @property - def shape(self) -> tuple[int, ...]: - return (len(self._series),) - - @property - def index_names(self) -> tuple[IndexName, ...]: - index_name = ( - self._series.index.name - if self._series.index.name is not None - else f'dim_{self._axis_zero}' - ) - return (index_name,) - - @property - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - return {self.index_names[0]: self._series.index} - - -class XarrayDataArrayAdapter(ValueArray): - def __init__( - self, - data_array: xarray.DataArray, - ): - default_indices = { - dim: range(size) - for dim, size in data_array.sizes.items() - if dim not in data_array.coords - } - self._data_array = data_array.assign_coords(default_indices) - - @staticmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> XarrayDataArrayAdapter | None: - try: - import xarray - - if isinstance(obj, xarray.DataArray): - return XarrayDataArrayAdapter(obj) - except ModuleNotFoundError: - pass - - def _equal(self, other: XarrayDataArrayAdapter) -> bool: - return self._data_array.identical(other._data_array) - - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - return self._data_array.sel(dict(key)) - - def __getitem__(self, key: dict[IndexName, slice]) -> XarrayDataArrayAdapter: - return XarrayDataArrayAdapter(self._data_array.isel(key)) - - @property - def shape(self) -> tuple[int, ...]: - return self._data_array.shape - - @property - def index_names(self) -> tuple[IndexName, ...]: - return tuple(self._data_array.dims) - - @property - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - return { - dim: self._data_array.coords[dim].values for dim in self._data_array.dims - } - - -class ScippDataArrayAdapter(ValueArray): - def __init__(self, data_array: scipp.DataArray, scipp: ModuleType): - default_indices = { - dim: scipp.arange(dim, size, unit=None) - for dim, size in data_array.sizes.items() - if dim not in data_array.coords - } - self._data_array = data_array.assign_coords(default_indices) - self._scipp = scipp - - @staticmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> ScippDataArrayAdapter | None: - try: - import scipp - - if isinstance(obj, scipp.Variable): - return ScippDataArrayAdapter(scipp.DataArray(obj), scipp=scipp) - if isinstance(obj, scipp.DataArray): - return ScippDataArrayAdapter(obj, scipp=scipp) - except ModuleNotFoundError: - pass - - def _equal(self, other: ScippDataArrayAdapter) -> bool: - return self._scipp.identical(self._data_array, other._data_array) - - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - values = self._data_array - for dim, value in key: - # Reconstruct label, to use label-based indexing instead of positional - if isinstance(value, tuple): - value, unit = value - else: - unit = None - label = self._scipp.scalar(value, unit=unit) - # Scipp indexing uses a comma to separate dimension label from the index, - # unlike Numpy and other libraries where it separates the indices for - # different axes. - values = values[dim, label] - return values - - def __getitem__(self, key: dict[IndexName, slice]) -> ScippDataArrayAdapter: - values = self._data_array - for dim, i in key: - values = values[dim, i] - return ScippDataArrayAdapter(values, scipp=self._scipp) - - @property - def shape(self) -> tuple[int, ...]: - return self._data_array.shape - - @property - def index_names(self) -> tuple[IndexName, ...]: - return tuple(self._data_array.dims) - - def _index_for_dim(self, dim: str) -> list[tuple[Any, scipp.Unit]]: - # Work around some NetworkX errors. Probably scipp.Variable lacks functionality. - # For now we return a list of tuples, where the first element is the value and - # the second is the unit. - coord = self._data_array.coords[dim] - unit = coord.unit - if unit is None: - return coord.values - unit = str(unit) - return [(value, unit) for value in coord.values] - - @property - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - return {dim: self._index_for_dim(dim) for dim in self._data_array.dims} - - -class NumpyArrayAdapter(ValueArray): - def __init__( - self, - array: numpy.ndarray, - *, - indices: dict[IndexName, Iterable[IndexValue]] | None = None, - axis_zero: int = 0, - ): - import numpy as np - - self._array = np.asarray(array) - if indices is None: - indices = { - f'dim_{i + axis_zero}': range(size) - for i, size in enumerate(self._array.shape) - } - self._indices = indices - self._axis_zero = axis_zero - - @staticmethod - def try_from(obj: Any, *, axis_zero: int = 0) -> NumpyArrayAdapter | None: - try: - import numpy - except ModuleNotFoundError: - return None - if isinstance(obj, numpy.ndarray): - return NumpyArrayAdapter(obj, axis_zero=axis_zero) - - def _equal(self, other: NumpyArrayAdapter) -> bool: - return ( - (self._array == other._array).all() - and self._indices == other._indices - and self._axis_zero == other._axis_zero - ) - - def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: - index_tuple = tuple(self._indices[k].index(i) for k, i in key) - return self._array[index_tuple] - - def __getitem__(self, key: dict[IndexName, slice]) -> NumpyArrayAdapter: - return NumpyArrayAdapter( - self._array[tuple(key.get(k, slice(None)) for k in self._indices)], - indices={ - index_name: (index_values[key.get(index_name, slice(None))]) - for index_name, index_values in self._indices.items() - }, - axis_zero=self._axis_zero, - ) - - @property - def shape(self) -> tuple[int, ...]: - return self._array.shape - - @property - def index_names(self) -> tuple[IndexName, ...]: - return tuple(self._indices) - - @property - def indices(self) -> dict[IndexName, Iterable[IndexValue]]: - return self._indices - - class NodeValues(Mapping[Hashable, ValueArray]): """ A collection of pandas.DataFrame-like objects with distinct indices. @@ -384,6 +28,14 @@ def __init__(self, values: Mapping[Any, ValueArray]): merged = self.merge(values) self._values = merged._values + def __contains__(self, key: Hashable) -> bool: + """Return True if the column with the given name exists.""" + return key in self._values + + def copy(self) -> NodeValues: + """Return a copy of the NodeValues.""" + return NodeValues(dict(self._values)) + def __len__(self) -> int: """Return the number of columns.""" return len(self._values) @@ -396,13 +48,35 @@ def __getitem__(self, key: Hashable) -> ValueArray: """Return the column with the given name.""" return self._values[key] + def __delitem__(self, key: Hashable) -> None: + """Remove the column with the given name.""" + if key in self._values: + del self._values[key] + else: + raise KeyError(f'Node "{key}" does not exist in NodeValues.') + def __setitem__(self, key: Hashable, value_array: ValueArray) -> None: """Add a single value array, checking for conflicts.""" # Check if the value array is identical to existing one - existing_value = self._values.get(key) - if existing_value is not None: - if existing_value == value_array: + old_value = self._values.get(key) + if old_value is not None: + if old_value == value_array: return # No change needed + elif old_value.index_names == value_array.index_names: + for old_index, new_index in zip( + old_value.indices.values(), + value_array.indices.values(), + strict=True, + ): + if (len(old_index) != len(new_index)) or any( + i != j for i, j in zip(old_index, new_index, strict=True) + ): + raise ValueError( + f"Node '{key}' has already been mapped with different " + f"indices: existing {old_index} vs new {new_index}" + ) + # If indices match, we can replace the value + self._values[key] = value_array else: raise ValueError(f"Node '{key}' has already been mapped") diff --git a/src/cyclebane/value_array.py b/src/cyclebane/value_array.py new file mode 100644 index 0000000..62eff0a --- /dev/null +++ b/src/cyclebane/value_array.py @@ -0,0 +1,112 @@ +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2025 Scipp contributors (https://github.com/scipp) +from __future__ import annotations + +from abc import ABC, abstractmethod +from collections.abc import Hashable, Iterable +from dataclasses import dataclass +from typing import Any, ClassVar, TypeVar + +IndexName = Hashable +IndexValue = Hashable + +T = TypeVar('T', bound='ValueArray') + + +class ValueArray(ABC): + """ + Abstract base class for a series of values with an index that can be sliced. + + Used by :py:class:`NodeValues` to store the values of a given node in a graph. The + abstraction allows for the use of different data structures to store the values of + nodes in a graph, such as pandas.DataFrame, xarray.DataArray, numpy.ndarray, or + simple Python iterables. + """ + + _registry: ClassVar = [] + + def __init_subclass__(cls) -> None: + super().__init_subclass__() + ValueArray._registry.append(cls) + + @staticmethod + def from_array_like(values: Any, *, axis_zero: int = 0) -> ValueArray: + # Reversed to ensure SequenceAdapter is tried last, as it is the most general + # SequenceAdapter is defined right after this class so it is registered first + for subclass in reversed(ValueArray._registry): + if (a := subclass.try_from(values, axis_zero=axis_zero)) is not None: + return a + raise ValueError(f'Cannot create ValueArray from {values}') + + @staticmethod + @abstractmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> ValueArray | None: ... + + def __eq__(self, other: object) -> bool: + if type(self) is not type(other): + return NotImplemented + return self._equal(other) + + def __ne__(self, other: object) -> bool: + return not self == other + + @abstractmethod + def _equal(self: T, other: T) -> bool: ... + + @abstractmethod + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + """Return data by selecting from index with given name and index value.""" + + def loc(self, key: dict[IndexName, slice]) -> ValueArray: + if not all(isinstance(i, slice) for i in key.values()): + raise ValueError('ValueArray.loc only accepts slices, not integers') + if not set(key).issubset(set(self.index_names)): + raise ValueError( + f'ValueArray.loc got {key.keys()}, not a subset of {self.index_names}' + ) + return self[key] + + @abstractmethod + def __getitem__(self, key: dict[IndexName, slice]) -> ValueArray: + pass + + @property + @abstractmethod + def shape(self) -> tuple[int, ...]: + pass + + @property + @abstractmethod + def index_names(self) -> tuple[IndexName, ...]: + pass + + @property + @abstractmethod + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + pass + + def group(self, index_name: Hashable) -> ValueArray: + """ + Group the values by their indices. + + This method is expected to return a new ValueArray that groups the values by + their indices, allowing for operations like aggregation or summarization. + """ + raise NotImplementedError( + 'ValueArray.group() is only implemented for Pandas series.' + ) + + def get_grouping(self) -> Grouping | None: + """ + If the instance holds grouping information, return it. + + Meant to be overridden by subclasses that support grouping. + """ + return None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Grouping: + indices: Iterable[Iterable[IndexValue]] + index_name: IndexName + group_index_name: IndexName diff --git a/src/cyclebane/value_array_adapters.py b/src/cyclebane/value_array_adapters.py new file mode 100644 index 0000000..e6177b1 --- /dev/null +++ b/src/cyclebane/value_array_adapters.py @@ -0,0 +1,316 @@ +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2024 Scipp contributors (https://github.com/scipp) +from __future__ import annotations + +from collections.abc import Hashable, Iterable, Sequence +from types import ModuleType +from typing import TYPE_CHECKING, Any, TypeVar + +from .value_array import Grouping, ValueArray + +if TYPE_CHECKING: + import numpy + import pandas + import scipp + import xarray + +IndexName = Hashable +IndexValue = Hashable + +T = TypeVar('T', bound='ValueArray') + + +class SequenceAdapter(ValueArray): + def __init__( + self, + values: Sequence[Any], + *, + index: Iterable[IndexValue] | None = None, + axis_zero: int = 0, + ): + self._values = values + self._index = index or range(len(values)) + self._axis_zero = axis_zero + + @staticmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> SequenceAdapter | None: + return SequenceAdapter(obj, axis_zero=axis_zero) + + def _equal(self, other: SequenceAdapter) -> bool: + return ( + self._values == other._values + and self._index == other._index + and self._axis_zero == other._axis_zero + ) + + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + if len(key) != 1: + raise ValueError('SequenceAdapter only supports single index') + _, i = key[0] + return self._values[self._index.index(i)] + + def __getitem__(self, key: dict[IndexName, slice]) -> SequenceAdapter: + _, i = next(iter(key.items())) + return SequenceAdapter( + self._values[i], index=self._index[i], axis_zero=self._axis_zero + ) + + @property + def shape(self) -> tuple[int, ...]: + return (len(self._values),) + + @property + def index_names(self) -> tuple[IndexName, ...]: + return (f'dim_{self._axis_zero}',) + + @property + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + return {f'dim_{self._axis_zero}': self._index} + + +class PandasSeriesAdapter(ValueArray): + def __init__( + self, series: pandas.Series, *, axis_zero: int = 0, _is_grouping: bool = False + ): + self._series = series + self._axis_zero = axis_zero + self._is_grouping = _is_grouping + + @staticmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> PandasSeriesAdapter | None: + try: + import pandas + except ModuleNotFoundError: + return None + if isinstance(obj, pandas.Series): + return PandasSeriesAdapter(obj, axis_zero=axis_zero) + + def _equal(self, other: PandasSeriesAdapter) -> bool: + return ( + self._series.equals(other._series) and self._axis_zero == other._axis_zero + ) + + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + if len(key) != 1: + raise ValueError('PandasSeriesAdapter only supports single index') + index_name, i = key[0] + if index_name != self.index_names[0]: + raise ValueError( + f'Unexpected index name {index_name} for PandasSeriesAdapter with ' + f'index names {self.index_names}' + ) + return self._series.loc[i] + + def __getitem__(self, key: dict[IndexName, slice]) -> PandasSeriesAdapter: + _, i = next(iter(key.items())) + return PandasSeriesAdapter(self._series[i], axis_zero=self._axis_zero) + + @property + def shape(self) -> tuple[int, ...]: + return (len(self._series),) + + @property + def index_names(self) -> tuple[IndexName, ...]: + index_name = ( + self._series.index.name + if self._series.index.name is not None + else f'dim_{self._axis_zero}' + ) + return (index_name,) + + @property + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + return {self.index_names[0]: self._series.index} + + def group(self, index_name: Hashable) -> PandasSeriesAdapter: + inner_index = self.index_names[0] + groupby = self._series.rename_axis(inner_index).groupby(self._series) + groups = type(self._series)(groupby.groups) + groups.index.rename(index_name, inplace=True) + return PandasSeriesAdapter(groups, _is_grouping=True) + + def get_grouping(self) -> Grouping | None: + if self._is_grouping: + return Grouping( + indices=self._series, + group_index_name=self.index_names[0], + index_name=next(iter(self._series)).name, + ) + + +class XarrayDataArrayAdapter(ValueArray): + def __init__( + self, + data_array: xarray.DataArray, + ): + default_indices = { + dim: range(size) + for dim, size in data_array.sizes.items() + if dim not in data_array.coords + } + self._data_array = data_array.assign_coords(default_indices) + + @staticmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> XarrayDataArrayAdapter | None: + try: + import xarray + + if isinstance(obj, xarray.DataArray): + return XarrayDataArrayAdapter(obj) + except ModuleNotFoundError: + pass + + def _equal(self, other: XarrayDataArrayAdapter) -> bool: + return self._data_array.identical(other._data_array) + + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + return self._data_array.sel(dict(key)) + + def __getitem__(self, key: dict[IndexName, slice]) -> XarrayDataArrayAdapter: + return XarrayDataArrayAdapter(self._data_array.isel(key)) + + @property + def shape(self) -> tuple[int, ...]: + return self._data_array.shape + + @property + def index_names(self) -> tuple[IndexName, ...]: + return tuple(self._data_array.dims) + + @property + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + return { + dim: self._data_array.coords[dim].values for dim in self._data_array.dims + } + + +class ScippDataArrayAdapter(ValueArray): + def __init__(self, data_array: scipp.DataArray, scipp: ModuleType): + default_indices = { + dim: scipp.arange(dim, size, unit=None) + for dim, size in data_array.sizes.items() + if dim not in data_array.coords + } + self._data_array = data_array.assign_coords(default_indices) + self._scipp = scipp + + @staticmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> ScippDataArrayAdapter | None: + try: + import scipp + + if isinstance(obj, scipp.Variable): + return ScippDataArrayAdapter(scipp.DataArray(obj), scipp=scipp) + if isinstance(obj, scipp.DataArray): + return ScippDataArrayAdapter(obj, scipp=scipp) + except ModuleNotFoundError: + pass + + def _equal(self, other: ScippDataArrayAdapter) -> bool: + return self._scipp.identical(self._data_array, other._data_array) + + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + values = self._data_array + for dim, value in key: + # Reconstruct label, to use label-based indexing instead of positional + if isinstance(value, tuple): + value, unit = value + else: + unit = None + label = self._scipp.scalar(value, unit=unit) + # Scipp indexing uses a comma to separate dimension label from the index, + # unlike Numpy and other libraries where it separates the indices for + # different axes. + values = values[dim, label] + return values + + def __getitem__(self, key: dict[IndexName, slice]) -> ScippDataArrayAdapter: + values = self._data_array + for dim, i in key: + values = values[dim, i] + return ScippDataArrayAdapter(values, scipp=self._scipp) + + @property + def shape(self) -> tuple[int, ...]: + return self._data_array.shape + + @property + def index_names(self) -> tuple[IndexName, ...]: + return tuple(self._data_array.dims) + + def _index_for_dim(self, dim: str) -> list[tuple[Any, scipp.Unit]]: + # Work around some NetworkX errors. Probably scipp.Variable lacks functionality. + # For now we return a list of tuples, where the first element is the value and + # the second is the unit. + coord = self._data_array.coords[dim] + unit = coord.unit + if unit is None: + return coord.values + unit = str(unit) + return [(value, unit) for value in coord.values] + + @property + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + return {dim: self._index_for_dim(dim) for dim in self._data_array.dims} + + +class NumpyArrayAdapter(ValueArray): + def __init__( + self, + array: numpy.ndarray, + *, + indices: dict[IndexName, Iterable[IndexValue]] | None = None, + axis_zero: int = 0, + ): + import numpy as np + + self._array = np.asarray(array) + if indices is None: + indices = { + f'dim_{i + axis_zero}': range(size) + for i, size in enumerate(self._array.shape) + } + self._indices = indices + self._axis_zero = axis_zero + + @staticmethod + def try_from(obj: Any, *, axis_zero: int = 0) -> NumpyArrayAdapter | None: + try: + import numpy + except ModuleNotFoundError: + return None + if isinstance(obj, numpy.ndarray): + return NumpyArrayAdapter(obj, axis_zero=axis_zero) + + def _equal(self, other: NumpyArrayAdapter) -> bool: + return ( + (self._array == other._array).all() + and self._indices == other._indices + and self._axis_zero == other._axis_zero + ) + + def sel(self, key: tuple[tuple[IndexName, IndexValue], ...]) -> Any: + index_tuple = tuple(self._indices[k].index(i) for k, i in key) + return self._array[index_tuple] + + def __getitem__(self, key: dict[IndexName, slice]) -> NumpyArrayAdapter: + return NumpyArrayAdapter( + self._array[tuple(key.get(k, slice(None)) for k in self._indices)], + indices={ + index_name: (index_values[key.get(index_name, slice(None))]) + for index_name, index_values in self._indices.items() + }, + axis_zero=self._axis_zero, + ) + + @property + def shape(self) -> tuple[int, ...]: + return self._array.shape + + @property + def index_names(self) -> tuple[IndexName, ...]: + return tuple(self._indices) + + @property + def indices(self) -> dict[IndexName, Iterable[IndexValue]]: + return self._indices diff --git a/tests/graph_test.py b/tests/graph_test.py index faac437..b25448c 100644 --- a/tests/graph_test.py +++ b/tests/graph_test.py @@ -696,6 +696,27 @@ def test_setitem_preserves_nodes_that_are_ancestors_of_unrelated_node() -> None: nx.utils.graphs_equal(graph.to_networkx(), g) +def test_setitem_preserves_node_values_of_sink_nodes() -> None: + g = nx.DiGraph() + g.add_edge('a', 'b') + g.add_edge('b', 'c') + + graph = cb.Graph(g) + mapped = graph.map({'a': [1, 2, 3]}) + # Special case: The graph we are setting has mapped node values associated with its + # sink node. The setitem effectively renames a to b in the graph, this ensures we + # are also renaming/preserving the associated node values. This is different from + # regular node attributes since mapped node values are not stored as node data in + # the underlying NetworkX graph, but in a separate data structure. + mapped['b'] = mapped['a'] + + result = mapped.to_networkx() + assert result.nodes[idx('b', 0)] == {'value': 1} + assert result.nodes[idx('b', 1)] == {'value': 2} + assert result.nodes[idx('b', 2)] == {'value': 3} + assert len(result.nodes) == 3 * 2 + + def test_getitem_returns_graph_containing_only_key_and_ancestors() -> None: g = nx.DiGraph() g.add_edge('a', 'b') @@ -902,16 +923,14 @@ def test_setitem_allows_compatible_node_values(node_values) -> None: assert len(mapped.index_names) == 1 -def test_setitem_raises_if_node_values_equivalent_but_of_different_type() -> None: +def test_setitem_allows_changing_node_values() -> None: g = nx.DiGraph() g.add_edge('a', 'b') graph = cb.Graph(g) mapped1 = graph.map({'a': [1, 2]}).reduce('b', name='d') - mapped2 = graph.map({'a': np.array([1, 2])}).reduce('b', name='d') - # One could imagine treating this as equivalent, but we are strict in the - # comparison. - with pytest.raises(ValueError, match="Node 'a' has already been mapped"): - mapped1['x'] = mapped2['d'] + mapped2 = graph.map({'a': [1, 3]}).reduce('b', name='d') + mapped1['x'] = mapped2['d'] + assert len(mapped1.index_names) == 1 def test_setitem_raises_if_node_values_incompatible() -> None: diff --git a/tests/groupby_test.py b/tests/groupby_test.py new file mode 100644 index 0000000..b11bd49 --- /dev/null +++ b/tests/groupby_test.py @@ -0,0 +1,180 @@ +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2025 Scipp contributors (https://github.com/scipp) +from collections.abc import Hashable + +import networkx as nx +import pandas as pd + +import cyclebane as cb + + +def idx( + name: str, *index: Hashable, offset=None, dims: tuple[str, ...] = ('dim_0', 'dim_1') +) -> cb.graph.NodeName: + """Helper to create a NodeName with a tuple of indices.""" + return cb.graph.NodeName( + name, + cb.graph.IndexValues(dims[offset : len(index) + (offset or 0)], tuple(index)), + ) + + +def test_basic_map_groupby_reduce_gives_correct_graph_structure() -> None: + g = nx.DiGraph() + g.add_edge('a', 'c') + g.add_edge('b', 'c') + df = pd.DataFrame({'a': [11, 22, 33], 'b': ['a', 'a', 'b']}) + + graph = cb.Graph(g) + mapped = graph.map(df) + grouped = mapped.groupby('b').reduce('c', name='d') + + result = grouped.to_networkx() + + # Nodes before grouping + assert result.nodes[idx('a', 0)] == {'value': 11} + assert result.nodes[idx('b', 0)] == {'value': 'a'} + assert result.nodes[idx('c', 0)] == {} + # Nodes after grouping + assert result.nodes[idx('d', 'a', dims=('b',))] == {} + + # Edges to grouped node + assert result.has_edge(idx('c', 0), idx('d', 'a', dims=('b',))) + assert result.has_edge(idx('c', 1), idx('d', 'a', dims=('b',))) + assert result.has_edge(idx('c', 2), idx('d', 'b', dims=('b',))) + # No cross-group edges + assert not result.has_edge(idx('c', 0), idx('d', 'b', dims=('b',))) + assert not result.has_edge(idx('c', 1), idx('d', 'b', dims=('b',))) + assert not result.has_edge(idx('c', 2), idx('d', 'a', dims=('b',))) + + +def test_group_twice_in_same_path() -> None: + g1 = nx.DiGraph() + g1.add_edge('a', 'c') + g1.add_edge('param1', 'c') + g1.add_edge('c', 'd') + + g2 = nx.DiGraph() + g2.add_edge('e', 'f') + g2.add_edge('param2', 'f') + + grouped = ( + cb.Graph(g1) + .map(pd.DataFrame({'a': [11, 22, 33, 44], 'param1': ['x', 'x', 'y', 'z']})) + .groupby('param1') + .reduce('d', name='grouped-d') + ) + mapped = cb.Graph(g2).map( + pd.DataFrame( + {'e': [1, 2, 3], 'param2': [0, 1, 1], 'param1': ['x', 'y', 'z']} + ).set_index('param1') + ) + + mapped['e'] = grouped + grouped_twice = mapped.groupby('param2').reduce('f', name='grouped-f') + + result = grouped_twice.to_networkx() + + # Nodes from second grouping (grouped-f) + assert result.nodes[idx('grouped-f', 0, dims=('param2',))] == {} + assert result.nodes[idx('grouped-f', 1, dims=('param2',))] == {} + + # Nodes from first grouping / mapping over param1 + assert result.nodes[idx('param2', 'x', dims=('param1',))] == {'value': 0} + assert result.nodes[idx('param2', 'y', dims=('param1',))] == {'value': 1} + assert result.nodes[idx('param2', 'z', dims=('param1',))] == {'value': 1} + assert result.nodes[idx('f', 'x', dims=('param1',))] == {} + assert result.nodes[idx('f', 'y', dims=('param1',))] == {} + assert result.nodes[idx('f', 'z', dims=('param1',))] == {} + # No value on 'e', was replaced by grouped-d links + assert result.nodes[idx('e', 'x', dims=('param1',))] == {} + assert result.nodes[idx('e', 'y', dims=('param1',))] == {} + assert result.nodes[idx('e', 'z', dims=('param1',))] == {} + assert idx('grouped-d', 'x', dims=('param1',)) not in result.nodes + assert idx('grouped-d', 'y', dims=('param1',)) not in result.nodes + assert idx('grouped-d', 'z', dims=('param1',)) not in result.nodes + + # Nodes from mapping over dim_0 + assert result.nodes[idx('a', 0)] == {'value': 11} + assert result.nodes[idx('a', 1)] == {'value': 22} + assert result.nodes[idx('a', 2)] == {'value': 33} + assert result.nodes[idx('a', 3)] == {'value': 44} + assert result.nodes[idx('param1', 0)] == {'value': 'x'} + assert result.nodes[idx('param1', 1)] == {'value': 'x'} + assert result.nodes[idx('param1', 2)] == {'value': 'y'} + assert result.nodes[idx('param1', 3)] == {'value': 'z'} + assert result.nodes[idx('c', 0)] == {} + assert result.nodes[idx('c', 1)] == {} + assert result.nodes[idx('c', 2)] == {} + assert result.nodes[idx('c', 3)] == {} + assert result.nodes[idx('d', 0)] == {} + assert result.nodes[idx('d', 1)] == {} + assert result.nodes[idx('d', 2)] == {} + assert result.nodes[idx('d', 3)] == {} + + # Edges within dim_0 (original graph structure) + assert result.has_edge(idx('a', 0), idx('c', 0)) + assert result.has_edge(idx('a', 1), idx('c', 1)) + assert result.has_edge(idx('a', 2), idx('c', 2)) + assert result.has_edge(idx('a', 3), idx('c', 3)) + assert result.has_edge(idx('param1', 0), idx('c', 0)) + assert result.has_edge(idx('param1', 1), idx('c', 1)) + assert result.has_edge(idx('param1', 2), idx('c', 2)) + assert result.has_edge(idx('param1', 3), idx('c', 3)) + assert result.has_edge(idx('c', 0), idx('d', 0)) + assert result.has_edge(idx('c', 1), idx('d', 1)) + assert result.has_edge(idx('c', 2), idx('d', 2)) + assert result.has_edge(idx('c', 3), idx('d', 3)) + + # Edges within param1 dimension (second graph structure) + assert result.has_edge( + idx('param2', 'x', dims=('param1',)), idx('f', 'x', dims=('param1',)) + ) + assert result.has_edge( + idx('param2', 'y', dims=('param1',)), idx('f', 'y', dims=('param1',)) + ) + assert result.has_edge( + idx('param2', 'z', dims=('param1',)), idx('f', 'z', dims=('param1',)) + ) + + # Edges from dim_0 to param1 grouping (first groupby) + assert result.has_edge(idx('d', 0), idx('e', 'x', dims=('param1',))) + assert result.has_edge(idx('d', 1), idx('e', 'x', dims=('param1',))) + assert result.has_edge(idx('d', 2), idx('e', 'y', dims=('param1',))) + assert result.has_edge(idx('d', 3), idx('e', 'z', dims=('param1',))) + + # Edges from param1 to param2 grouping (second groupby) + assert result.has_edge( + idx('f', 'x', dims=('param1',)), idx('grouped-f', 0, dims=('param2',)) + ) + assert result.has_edge( + idx('f', 'y', dims=('param1',)), idx('grouped-f', 1, dims=('param2',)) + ) + assert result.has_edge( + idx('f', 'z', dims=('param1',)), idx('grouped-f', 1, dims=('param2',)) + ) + + +def test_group_in_different_ways() -> None: + g = nx.DiGraph() + g.add_edge('a', 'b') + g.add_edge('attach2', 'd') + df = pd.DataFrame( + {'a': [11, 22, 33], 'param1': ['a', 'a', 'b'], 'param2': ['x', 'y', 'x']} + ) + + graph = cb.Graph(g) + mapped = graph.map(df) + grouped = mapped.groupby('param1').reduce('b', name='grouped1') + grouped2 = mapped.groupby('param2').reduce('b', name='grouped2') + + # Map helper node over param2 so we have a place where we can attached the grouping + # by param2. + grouped = grouped.map( + pd.DataFrame({'attach2': [None, None], 'param2': ['x', 'y']}).set_index( + 'param2' + ) + ) + grouped['attach2'] = grouped2['grouped2'] + + # with pytest.raises(KeyError, match='dim_0'): + grouped.to_networkx() diff --git a/tests/node_values_test.py b/tests/node_values_test.py index c65c82d..6d64eff 100644 --- a/tests/node_values_test.py +++ b/tests/node_values_test.py @@ -147,15 +147,16 @@ def test_merge_existing_node_equal_but_different_object(self): assert len(merged) == 1 assert merged is not node_values # Should return new object (copy) - def test_merge_existing_node_different_value_raises(self): + def test_merge_existing_node_different_value_works(self): """Re-adding existing node with different value.""" initial_values = {'a': ValueArray.from_array_like([1, 2, 3], axis_zero=0)} node_values = NodeValues(initial_values) new_values = {'a': ValueArray.from_array_like([4, 5, 6], axis_zero=0)} - with pytest.raises(ValueError, match="Node 'a' has already been mapped"): - node_values.merge(new_values) + merged = node_values.merge(new_values) + assert len(merged) == 1 + assert merged['a'] == new_values['a'] # Should update value def test_merge_empty_new_values(self): """Merging empty mapping.""" @@ -194,7 +195,7 @@ def test_merge_multiple_new_nodes_mixed_conflicts_raises(self): ): node_values.merge(new_values) - def test_merge_multiple_new_nodes_one_existing_raises(self): + def test_merge_multiple_new_nodes_one_existing_compatible_index(self): """Multiple new nodes where one already exists.""" initial_values = {'a': ValueArray.from_array_like([1, 2, 3], axis_zero=0)} node_values = NodeValues(initial_values) @@ -202,11 +203,30 @@ def test_merge_multiple_new_nodes_one_existing_raises(self): new_values = { 'a': ValueArray.from_array_like( [4, 5, 6], axis_zero=0 - ), # Exists but different + ), # Exists but different values 'b': ValueArray.from_array_like([7, 8, 9], axis_zero=1), # New } - with pytest.raises(ValueError, match="Node 'a' has already been mapped"): + merged = node_values.merge(new_values) + assert len(merged) == 2 + assert set(merged.keys()) == {'a', 'b'} + assert merged['a'].sel((('dim_0', 0),)) == 4 # Updated value + + def test_merge_multiple_new_nodes_one_existing_raises(self): + """Multiple new nodes where one already exists.""" + initial_values = {'a': ValueArray.from_array_like([1, 2, 3], axis_zero=0)} + node_values = NodeValues(initial_values) + + new_values = { + 'a': ValueArray.from_array_like([4, 5, 6, 8], axis_zero=0)[ + {'dim_0': slice(1, 4)} # Make a slice to obtain an index value conflict + ], # Exists but different + 'b': ValueArray.from_array_like([7, 8, 9], axis_zero=1), # New + } + + with pytest.raises( + ValueError, match="Node 'a' has already been mapped with different indices" + ): node_values.merge(new_values) def test_merge_partial_index_overlap_compatible(self):