|
17 | 17 |
|
18 | 18 | import copy |
19 | 19 | import functools |
20 | | -from collections.abc import Generator, Iterable, Sequence |
21 | 20 | from contextlib import suppress |
22 | 21 | from itertools import islice |
23 | 22 | from math import prod |
24 | 23 | from pathlib import Path |
25 | | -from typing import TYPE_CHECKING, Literal, NamedTuple, overload |
| 24 | +from typing import TYPE_CHECKING, NamedTuple |
26 | 25 |
|
27 | 26 | import h5py |
28 | 27 | import matplotlib.pyplot as mpl |
|
32 | 31 | from matplotlib.colors import to_hex as mpl_to_hex |
33 | 32 | from matplotlib.lines import lineStyles |
34 | 33 | from matplotlib.markers import MarkerStyle |
35 | | -from more_itertools import first, locate, nth, nth_product, sort_together, unzip |
| 34 | +from more_itertools import first, locate, nth, nth_product |
36 | 35 | from qtpy.QtCore import QModelIndex, Qt, Signal, Slot |
37 | 36 | from qtpy.QtGui import QColor, QStandardItem, QStandardItemModel |
38 | 37 |
|
39 | 38 | from MDANSE.IO.IOUtils import summarise_array |
40 | 39 | from MDANSE.MLogging import LOG |
41 | 40 |
|
42 | 41 | if TYPE_CHECKING: |
43 | | - from collections.abc import Iterable |
| 42 | + from collections.abc import Generator, Iterable, Sequence |
44 | 43 |
|
45 | 44 | NUMBERS_FOR_SLICE = 3 |
46 | 45 | NUMBERS_FOR_RANGE = 2 |
@@ -619,7 +618,7 @@ def curves_vs_axis( |
619 | 618 | x_axis = self.x_axis(axis_label) |
620 | 619 |
|
621 | 620 | if self._data.ndim == 1: |
622 | | - yield axis_label, (x_axis, self.data) |
| 621 | + yield "", (x_axis, self.data) |
623 | 622 | return |
624 | 623 |
|
625 | 624 | data_shape = self._data.shape |
@@ -686,67 +685,94 @@ def curves_vs_axis( |
686 | 685 |
|
687 | 686 | def planes_vs_axis( |
688 | 687 | self, |
689 | | - axis_number: int, |
690 | | - max_limit: int = 1, |
691 | | - ) -> Generator[tuple[str, npt.NDArray[np.floating]]]: |
| 688 | + main_axis: str, |
| 689 | + max_limit: int = 9, |
| 690 | + ) -> Generator[tuple[str, npt.NDArray[np.floating], tuple[str, str]]]: |
692 | 691 | """Prepare for plotting 2D subsets of an ND array. |
693 | 692 |
|
694 | 693 | Parameters |
695 | 694 | ---------- |
696 | | - axis_number : int |
697 | | - index of the axis perpendicular to the plotted array |
| 695 | + main_axis : str |
| 696 | + Label of the axis perpendicular to the plotted array. |
698 | 697 |
|
699 | 698 | Yields |
700 | 699 | ------ |
701 | | - str |
| 700 | + main_label : str |
702 | 701 | Grid label. |
703 | | - npt.NDArray[np.floating] |
| 702 | + image_array : npt.NDArray[np.floating] |
704 | 703 | 2D array. |
705 | | -
|
| 704 | + axis_labels : tuple[str, ...] |
| 705 | + Labels for each axis. |
706 | 706 | """ |
| 707 | + main_axis_index = self.main_axis_index(main_axis) |
| 708 | + other_labels = self._axis_labels(main_axis) |
| 709 | + |
707 | 710 | match self._data.ndim: |
708 | 711 | case 1: |
709 | 712 | pass |
710 | | - case 2 if axis_number == 1: |
711 | | - yield self._labels["medium"], self.data.T |
| 713 | + case 2 if main_axis_index == 1: |
| 714 | + yield self._labels["medium"], self.data.T, (main_axis, other_labels[0]) |
712 | 715 | case 2: |
713 | | - yield self._labels["medium"], self.data |
| 716 | + yield self._labels["medium"], self.data, (main_axis, other_labels[0]) |
714 | 717 | case 3: |
715 | 718 | perpendicular_axis_name, perpendicular_axis = nth( |
716 | | - self._axes.items(), axis_number, default=(None, None) |
| 719 | + self._axes.items(), main_axis_index, default=(None, None) |
717 | 720 | ) |
718 | 721 |
|
719 | 722 | if perpendicular_axis is None: |
720 | 723 | return |
721 | 724 |
|
722 | | - reordered_view = np.moveaxis(self.data, axis_number, 0) |
| 725 | + reordered_view = np.moveaxis(self.data, main_axis_index, 0) |
723 | 726 |
|
724 | 727 | for plane_number in self.curve_ind(max_limit): |
| 728 | + if plane_number > len(reordered_view): |
| 729 | + continue |
| 730 | + |
725 | 731 | yield ( |
726 | 732 | f"{self._labels['minimal']}:{perpendicular_axis_name}={perpendicular_axis[plane_number]}", |
727 | 733 | reordered_view[plane_number], |
| 734 | + other_labels, |
728 | 735 | ) |
729 | 736 | case _: |
730 | 737 | raise NotImplementedError( |
731 | 738 | f"Cannot handle {self._data.ndim}-dimensional data." |
732 | 739 | ) |
733 | 740 |
|
734 | | - def main_axis_index(self, main_axis: str | None, *, default: int) -> int: |
| 741 | + def main_axis_index( |
| 742 | + self, main_axis: str | None, *, default: int | None = None |
| 743 | + ) -> int: |
735 | 744 | """Find index of main axis. |
736 | 745 |
|
737 | 746 | Parameters |
738 | 747 | ---------- |
739 | 748 | main_axis : str |
740 | 749 | Main axis name to search for. |
741 | | - default : int |
| 750 | + default : int, optional |
742 | 751 | Index if ``main_axis`` not found. |
743 | 752 |
|
744 | 753 | Returns |
745 | 754 | ------- |
746 | 755 | int |
747 | 756 | Index of main axis. |
| 757 | +
|
| 758 | + Raises |
| 759 | + ------ |
| 760 | + ValueError |
| 761 | + Axis not found and no default. |
748 | 762 | """ |
749 | | - return first(locate(self._axes, pred=lambda x: x == main_axis), default) |
| 763 | + ind = first(locate(self._axes, pred=main_axis.__eq__), default) |
| 764 | + if ind is None: |
| 765 | + raise ValueError( |
| 766 | + f"Cannot find axis {main_axis} in {','.join(self._axes.keys())}" |
| 767 | + ) |
| 768 | + return ind |
| 769 | + |
| 770 | + def _axis_labels(self, main_axis: str) -> tuple[str] | tuple[str, str]: |
| 771 | + main_axis_index = self.main_axis_index(main_axis) |
| 772 | + |
| 773 | + return tuple( |
| 774 | + label for i, label in enumerate(self._axes) if i != main_axis_index |
| 775 | + ) |
750 | 776 |
|
751 | 777 |
|
752 | 778 | plotting_column_labels = [ |
@@ -1057,18 +1083,16 @@ def delete_dataset(self, index: QModelIndex): |
1057 | 1083 | self._datasets.pop(dkey, None) |
1058 | 1084 |
|
1059 | 1085 | def planes( |
1060 | | - self, default_axis: int = 0, planes_per_dataset: int | None = None |
1061 | | - ) -> Generator[tuple[PlotArgs, str, npt.NDArray[np.floating]]]: |
| 1086 | + self, default_axis: int | None = None, planes_per_dataset: int | None = None |
| 1087 | + ) -> Generator[tuple[PlotArgs, str, npt.NDArray[np.floating], tuple[str, str]]]: |
1062 | 1088 | for databundle in self.datasets().values(): |
1063 | 1089 | ds = databundle.dataset |
1064 | 1090 |
|
1065 | | - for label, plane in islice( |
1066 | | - ds.planes_vs_axis( |
1067 | | - ds.main_axis_index(databundle.main_axis, default=default_axis) |
1068 | | - ), |
| 1091 | + for label, plane, axis_labels in islice( |
| 1092 | + ds.planes_vs_axis(databundle.main_axis), |
1069 | 1093 | planes_per_dataset, |
1070 | 1094 | ): |
1071 | | - yield databundle, label, plane |
| 1095 | + yield databundle, label, plane, axis_labels |
1072 | 1096 |
|
1073 | 1097 | def curves( |
1074 | 1098 | self, curves_per_dataset: int | None = None |
|
0 commit comments