diff --git a/nion/data/annotated_array/__init__.py b/nion/data/annotated_array/__init__.py new file mode 100644 index 0000000..dbae22a --- /dev/null +++ b/nion/data/annotated_array/__init__.py @@ -0,0 +1,53 @@ +"""Public annotated-array API surface. + +Import this package as the single entry point for annotated-array related functionality: + + from nion.data import annotated_array as aa +""" + +from .primitives import fft +from .primitives import ifft + +from ._implementation import AffineCalibration +from ._implementation import AnnotatedArray +from ._implementation import ArrayDescriptor +from ._implementation import ArrayHeader +from ._implementation import ArrayMetadata +from ._implementation import ArrayProtocol +from ._implementation import Axis +from ._implementation import AxisGroup +from ._implementation import Calibration +from ._implementation import CalibrationSet +from ._implementation import collapse_scalar_axis_groups +from ._implementation import CoordinateCalibration +from ._implementation import ExtensionRecord +from ._implementation import ValueType +from ._implementation import from_data_and_metadata +from ._implementation import infer_value_type +from ._implementation import to_data_and_metadata +from ._implementation import zeros_annotated_array + +__all__ = [ + "AffineCalibration", + "AnnotatedArray", + "ArrayDescriptor", + "ArrayHeader", + "ArrayMetadata", + "ArrayProtocol", + "Axis", + "AxisGroup", + "Calibration", + "CalibrationSet", + "collapse_scalar_axis_groups", + "CoordinateCalibration", + "ExtensionRecord", + "ValueType", + "fft", + "from_data_and_metadata", + "ifft", + "infer_value_type", + "to_data_and_metadata", + "zeros_annotated_array", +] + + diff --git a/nion/data/AnnotatedArray.py b/nion/data/annotated_array/_implementation.py similarity index 94% rename from nion/data/AnnotatedArray.py rename to nion/data/annotated_array/_implementation.py index dca12fa..4c90d79 100644 --- a/nion/data/AnnotatedArray.py +++ b/nion/data/annotated_array/_implementation.py @@ -539,7 +539,35 @@ def _legacy_timestamp_to_created( return utc_timestamp.astimezone(datetime.timezone.utc) -def from_data_and_metadata(xdata: DataAndMetadata.DataAndMetadata) -> AnnotatedArray: +def _collapse_scalar_axis_groups_for_descriptor(descriptor: ArrayDescriptor) -> ArrayDescriptor: + non_scalar_axis_groups = tuple(axis_group for axis_group in descriptor.axis_groups if axis_group.rank > 0) + if len(non_scalar_axis_groups) == len(descriptor.axis_groups): + return descriptor + + axis_groups = non_scalar_axis_groups if non_scalar_axis_groups else (AxisGroup(),) + return ArrayDescriptor( + axis_groups=axis_groups, + intensity_calibrations=descriptor.intensity_calibrations, + value_type=descriptor.value_type, + ) + + +def collapse_scalar_axis_groups(array: AnnotatedArray) -> AnnotatedArray: + """Return an AnnotatedArray with rank-0 axis groups removed. + + If every axis group is scalar, a single scalar group is retained. + """ + collapsed_descriptor = _collapse_scalar_axis_groups_for_descriptor(array.descriptor) + if collapsed_descriptor == array.descriptor: + return array + return AnnotatedArray(data=array.data, descriptor=collapsed_descriptor, metadata=array.metadata) + + +def from_data_and_metadata( + xdata: DataAndMetadata.DataAndMetadata, + *, + collapse_scalar_axis_groups: bool = False, +) -> AnnotatedArray: """Convert a legacy DataAndMetadata instance into an AnnotatedArray.""" data = xdata.data if data is None: @@ -576,7 +604,12 @@ def from_data_and_metadata(xdata: DataAndMetadata.DataAndMetadata) -> AnnotatedA created = _legacy_timestamp_to_created(xdata.timestamp, xdata.timezone, xdata.timezone_offset) metadata = ArrayMetadata(created=created, attributes=dict(xdata.metadata)) - return AnnotatedArray(data=data, descriptor=descriptor, metadata=metadata) + annotated_array = AnnotatedArray(data=data, descriptor=descriptor, metadata=metadata) + if collapse_scalar_axis_groups: + collapsed_descriptor = _collapse_scalar_axis_groups_for_descriptor(annotated_array.descriptor) + if collapsed_descriptor != annotated_array.descriptor: + return AnnotatedArray(data=annotated_array.data, descriptor=collapsed_descriptor, metadata=annotated_array.metadata) + return annotated_array def to_data_and_metadata(annotated_array: AnnotatedArray) -> DataAndMetadata.DataAndMetadata: diff --git a/nion/data/annotated_array/primitives.py b/nion/data/annotated_array/primitives.py new file mode 100644 index 0000000..5e92e01 --- /dev/null +++ b/nion/data/annotated_array/primitives.py @@ -0,0 +1,275 @@ +"""Annotated-array processing primitives. + +This module is the home for low-level processing operations on annotated arrays. +Operation implementations can be added here and re-exported from +`nion.data.annotated_array` as they are introduced. +""" + +from __future__ import annotations + +import math +import typing + +import numpy +import numpy.typing +import scipy.fft + +from nion.data.annotated_array._implementation import ( + AffineCalibration, + AnnotatedArray, + ArrayDescriptor, + AxisGroup, + CoordinateCalibration, + ValueType, +) + + +# --------------------------------------------------------------------------- +# Internal calibration helpers +# --------------------------------------------------------------------------- + +def _spatial_to_frequency_axis_group(axis_group: AxisGroup) -> AxisGroup: + """Return a new AxisGroup with all calibrations transformed to frequency domain. + + For each axis of size *N* with spatial calibration ``(scale=s, unit=u)`` + the frequency calibration is:: + + scale_freq = 1 / (s * N) + offset_freq = (-0.5 - N // 2) / (s * N) # DC at centre after fftshift + unit_freq = "1/" + u # reciprocal unit + + All calibration keys are preserved unchanged (same keys in, same keys out). + Each calibration is independently transformed under its original key. + The primary calibration key is unchanged. + """ + def transform_coord_calibration(coord_cal: CoordinateCalibration) -> CoordinateCalibration: + freq_calibrations: list[AffineCalibration] = [] + for axis_index, axis in enumerate(axis_group.axes): + n = axis.size + spatial_cal = coord_cal.calibrations[axis_index] + s = spatial_cal.scale if isinstance(spatial_cal, AffineCalibration) else 1.0 + u = spatial_cal.unit if isinstance(spatial_cal, AffineCalibration) else "" + freq_calibrations.append(AffineCalibration( + scale=1.0 / (s * n), + offset=(-0.5 - n // 2) / (s * n), + unit=("1/" + u) if u else "", + )) + return CoordinateCalibration(calibrations=tuple(freq_calibrations)) + + new_calibrations = { + key: transform_coord_calibration(coord_cal) + for key, coord_cal in axis_group.coordinate_calibrations.items() + } + + return AxisGroup( + axes=axis_group.axes, + coordinate_system_id=axis_group.coordinate_system_id, + coordinate_calibrations=new_calibrations, + primary_calibration_key=axis_group.primary_calibration_key, + ) + + +def _frequency_to_spatial_axis_group(axis_group: AxisGroup) -> AxisGroup: + """Return a new AxisGroup with all calibrations transformed back to spatial domain. + + For each frequency calibration ``(scale=s_freq, unit=u_freq)``:: + + scale_spatial = 1 / (s_freq * N) + offset_spatial = 0 + unit_spatial = u_freq[2:] if u_freq.startswith("1/") else "" + + All calibration keys are preserved unchanged (same keys in, same keys out). + The primary calibration key is unchanged. + """ + def transform_coord_calibration(coord_cal: CoordinateCalibration) -> CoordinateCalibration: + spatial_calibrations: list[AffineCalibration] = [] + for axis_index, axis in enumerate(axis_group.axes): + n = axis.size + freq_cal = coord_cal.calibrations[axis_index] + if isinstance(freq_cal, AffineCalibration) and freq_cal.scale != 0.0: + s_freq = freq_cal.scale + u_freq = freq_cal.unit + s_spatial = 1.0 / (s_freq * n) + u_spatial = u_freq[2:] if u_freq.startswith("1/") else "" + else: + s_spatial, u_spatial = 1.0, "" + spatial_calibrations.append(AffineCalibration(scale=s_spatial, offset=0.0, unit=u_spatial)) + return CoordinateCalibration(calibrations=tuple(spatial_calibrations)) + + new_calibrations = { + key: transform_coord_calibration(coord_cal) + for key, coord_cal in axis_group.coordinate_calibrations.items() + } + + return AxisGroup( + axes=axis_group.axes, + coordinate_system_id=axis_group.coordinate_system_id, + coordinate_calibrations=new_calibrations, + primary_calibration_key=axis_group.primary_calibration_key, + ) + + +# --------------------------------------------------------------------------- +# Public API +# --------------------------------------------------------------------------- + +def fft(array: AnnotatedArray) -> AnnotatedArray: + """Compute the forward FFT of an :class:`AnnotatedArray`. + + The transform is applied to the array's only axis group. + + Energy normalisation + ~~~~~~~~~~~~~~~~~~~~ + The scaling factor ``1 / sqrt(N)`` (1-D) or ``1 / sqrt(N * M)`` (2-D) is + applied so that the RMS value is preserved: + + .. code-block:: python + + numpy.sqrt(numpy.mean(numpy.abs(data)**2)) + == numpy.sqrt(numpy.mean(numpy.abs(fft(xdata).data)**2)) + + DC at centre + ~~~~~~~~~~~~ + The result is passed through :func:`scipy.fft.fftshift` along the signal axes + so that the zero-frequency component lies at the array centre. + + Calibration + ~~~~~~~~~~~ + All calibrations in the signal :class:`AxisGroup` are transformed to the + frequency domain. The keys of all calibrations are preserved, so the user-defined + calibration names remain unchanged. For example, if the input has calibrations + keyed ``"spatial"`` and ``"angular"``, the output will have calibrations keyed + ``"spatial"`` and ``"angular"`` (with their scales/offsets/units transformed to + frequency). The primary calibration key is also preserved. + + Args: + array: Input :class:`AnnotatedArray` with a scalar or complex datum. + RGB and RGBA value types are not supported. + + Returns: + :class:`AnnotatedArray` with complex datum (``complex128``) and + frequency-domain calibrations on the signal :class:`AxisGroup`. + + Raises: + ValueError: If there is not exactly one axis group, if that group rank + is not 1 or 2, or if the value type is not ``SCALAR`` or + ``COMPLEX``. + """ + axis_groups = array.descriptor.axis_groups + if len(axis_groups) != 1: + raise ValueError(f"fft: expected exactly one axis group, got {len(axis_groups)}") + + signal_group = axis_groups[-1] + rank = signal_group.rank + + if rank not in (1, 2): + raise ValueError(f"fft: signal rank must be 1 or 2, got {rank}") + + value_type = array.descriptor.value_type + if value_type not in (ValueType.SCALAR, ValueType.COMPLEX): + raise ValueError( + f"fft: unsupported value type {value_type!r}; " + "only SCALAR and COMPLEX are supported" + ) + + data = numpy.asarray(array.data) + signal_shape = signal_group.shape # e.g. (N,) or (N, M) + signal_axes = tuple(range(-rank, 0)) # e.g. (-1,) or (-2, -1) + + if rank == 1: + n = signal_shape[0] + scaling = 1.0 / math.sqrt(n) + result_data: numpy.typing.NDArray[numpy.complexfloating[typing.Any, typing.Any]] = ( + scipy.fft.fftshift(scipy.fft.fft(data, axis=-1) * scaling, axes=signal_axes) + ) + else: + n, m = signal_shape + scaling = 1.0 / math.sqrt(n * m) + result_data = scipy.fft.fftshift( + scipy.fft.fft2(data, axes=signal_axes) * scaling, + axes=signal_axes, + ) + + new_signal_group = _spatial_to_frequency_axis_group(signal_group) + new_axis_groups = axis_groups[:-1] + (new_signal_group,) + new_descriptor = ArrayDescriptor( + axis_groups=new_axis_groups, + intensity_calibrations=array.descriptor.intensity_calibrations, + value_type=ValueType.COMPLEX, + ) + return AnnotatedArray(data=result_data, descriptor=new_descriptor, metadata=array.metadata) + + +def ifft(array: AnnotatedArray) -> AnnotatedArray: + """Compute the inverse FFT of an :class:`AnnotatedArray`. + + The transform is applied to the array's only axis group. + + Energy normalisation + ~~~~~~~~~~~~~~~~~~~~ + The inverse scaling factor ``sqrt(N)`` (1-D) or ``sqrt(N * M)`` (2-D) is + applied to be the exact inverse of :func:`fft`. + + DC at centre + ~~~~~~~~~~~~ + The input is assumed to have its DC component at the array centre + (produced by :func:`fft`); :func:`scipy.fft.ifftshift` is applied along + the signal axes before the inverse transform. + + Calibration round-trip + ~~~~~~~~~~~~~~~~~~~~~~ + All calibrations in the signal :class:`AxisGroup` are transformed back to the + spatial domain. The keys of all calibrations are preserved exactly, so if the + frequency-domain array has calibrations keyed ``"spatial"`` and ``"angular"``, + the result will also have calibrations with those same keys (now with spatial + scales/offsets/units). The primary calibration key is also preserved. + + Args: + array: :class:`AnnotatedArray` with a complex datum in frequency + space (DC at centre). + + Returns: + :class:`AnnotatedArray` with complex datum and spatial-domain + calibrations on the signal :class:`AxisGroup`. + + Raises: + ValueError: If there is not exactly one axis group, or if that group + rank is not 1 or 2. + """ + axis_groups = array.descriptor.axis_groups + if len(axis_groups) != 1: + raise ValueError(f"ifft: expected exactly one axis group, got {len(axis_groups)}") + + signal_group = axis_groups[-1] + rank = signal_group.rank + + if rank not in (1, 2): + raise ValueError(f"ifft: signal rank must be 1 or 2, got {rank}") + + data = numpy.asarray(array.data) + signal_shape = signal_group.shape + signal_axes = tuple(range(-rank, 0)) + + if rank == 1: + n = signal_shape[0] + scaling = math.sqrt(n) + result_data = scipy.fft.ifft( + scipy.fft.ifftshift(data, axes=signal_axes) * scaling, + axis=-1, + ) + else: + n, m = signal_shape + scaling = math.sqrt(n * m) + result_data = scipy.fft.ifft2( + scipy.fft.ifftshift(data, axes=signal_axes) * scaling, + axes=signal_axes, + ) + + new_signal_group = _frequency_to_spatial_axis_group(signal_group) + new_axis_groups = axis_groups[:-1] + (new_signal_group,) + new_descriptor = ArrayDescriptor( + axis_groups=new_axis_groups, + intensity_calibrations=array.descriptor.intensity_calibrations, + value_type=ValueType.COMPLEX, + ) + return AnnotatedArray(data=result_data, descriptor=new_descriptor, metadata=array.metadata) diff --git a/nion/data/annotated_array/types.py b/nion/data/annotated_array/types.py new file mode 100644 index 0000000..ebd2646 --- /dev/null +++ b/nion/data/annotated_array/types.py @@ -0,0 +1,6 @@ +"""Annotated-array operation support types. + +This module is reserved for parameter classes, enums, and helper structures used +by annotated-array primitives. +""" + diff --git a/nion/data/test/AnnotatedArray_test.py b/nion/data/test/AnnotatedArray_test.py index d79d188..fc7dffe 100644 --- a/nion/data/test/AnnotatedArray_test.py +++ b/nion/data/test/AnnotatedArray_test.py @@ -7,7 +7,7 @@ import h5py import numpy -from nion.data import AnnotatedArray +from nion.data import annotated_array from nion.data import Calibration from nion.data import DataAndMetadata @@ -30,9 +30,9 @@ class TestAnnotatedArray(unittest.TestCase): ) def test_calibration_set_uses_explicit_calibration_accessors(self) -> None: - primary = AnnotatedArray.AffineCalibration(unit="nm") - alternate = AnnotatedArray.AffineCalibration(unit="rad") - calibrations = AnnotatedArray.CalibrationSet.from_calibration(primary, key="primary").with_calibration("alternate", alternate) + primary = annotated_array.AffineCalibration(unit="nm") + alternate = annotated_array.AffineCalibration(unit="rad") + calibrations = annotated_array.CalibrationSet.from_calibration(primary, key="primary").with_calibration("alternate", alternate) self.assertTrue(calibrations.has_calibration("alternate")) self.assertIs(primary, calibrations.primary_calibration) @@ -40,19 +40,19 @@ def test_calibration_set_uses_explicit_calibration_accessors(self) -> None: self.assertIs(alternate, calibrations.with_primary_calibration("alternate").primary_calibration) def test_array_descriptor_describes_shape_and_rank(self) -> None: - collection = AnnotatedArray.AxisGroup.from_1d_size(3) - signal = AnnotatedArray.AxisGroup.from_2d_size((4, 5)) - descriptor = AnnotatedArray.ArrayDescriptor((collection, signal)) + collection = annotated_array.AxisGroup.from_1d_size(3) + signal = annotated_array.AxisGroup.from_2d_size((4, 5)) + descriptor = annotated_array.ArrayDescriptor((collection, signal)) self.assertEqual((3, 4, 5), descriptor.shape) self.assertEqual(3, descriptor.ndim) def test_array_descriptor_uses_identity_intensity_calibration_when_no_primary_is_designated(self) -> None: - descriptor = AnnotatedArray.ArrayDescriptor((AnnotatedArray.AxisGroup.from_1d_size(3),)) + descriptor = annotated_array.ArrayDescriptor((annotated_array.AxisGroup.from_1d_size(3),)) calibration = descriptor.get_intensity_calibration() - self.assertIsInstance(calibration, AnnotatedArray.AffineCalibration) - affine_calibration = typing.cast(AnnotatedArray.AffineCalibration, calibration) + self.assertIsInstance(calibration, annotated_array.AffineCalibration) + affine_calibration = typing.cast(annotated_array.AffineCalibration, calibration) self.assertEqual(1.0, affine_calibration.scale) self.assertEqual(0.0, affine_calibration.offset) self.assertEqual("", affine_calibration.unit) @@ -61,8 +61,8 @@ def test_array_descriptor_uses_identity_intensity_calibration_when_no_primary_is descriptor.get_intensity_calibration("missing") def test_axis_group_size_factories_accept_optional_coordinate_calibrations(self) -> None: - single_calibration = AnnotatedArray.CoordinateCalibration(calibrations=(AnnotatedArray.AffineCalibration(unit="nm"),)) - group_1d = AnnotatedArray.AxisGroup.from_1d_size( + single_calibration = annotated_array.CoordinateCalibration(calibrations=(annotated_array.AffineCalibration(unit="nm"),)) + group_1d = annotated_array.AxisGroup.from_1d_size( 3, coordinate_calibrations={"spatial": single_calibration}, primary_calibration_key="spatial", @@ -70,10 +70,10 @@ def test_axis_group_size_factories_accept_optional_coordinate_calibrations(self) self.assertEqual(("spatial",), group_1d.calibration_keys) self.assertEqual("spatial", group_1d.primary_calibration_key) - map_calibration = AnnotatedArray.CoordinateCalibration( - calibrations=(AnnotatedArray.AffineCalibration(unit="nm"), AnnotatedArray.AffineCalibration(unit="nm")) + map_calibration = annotated_array.CoordinateCalibration( + calibrations=(annotated_array.AffineCalibration(unit="nm"), annotated_array.AffineCalibration(unit="nm")) ) - group_2d = AnnotatedArray.AxisGroup.from_2d_size( + group_2d = annotated_array.AxisGroup.from_2d_size( (2, 3), coordinate_calibrations={"camera": map_calibration}, primary_calibration_key="camera", @@ -83,16 +83,16 @@ def test_axis_group_size_factories_accept_optional_coordinate_calibrations(self) def test_array_descriptor_requires_valid_axis_group_layout(self) -> None: with self.assertRaisesRegex(ValueError, "at least one"): - AnnotatedArray.ArrayDescriptor(()) + annotated_array.ArrayDescriptor(()) - scalar_group = AnnotatedArray.AxisGroup() - vector_group = AnnotatedArray.AxisGroup.from_1d_size(3) + scalar_group = annotated_array.AxisGroup() + vector_group = annotated_array.AxisGroup.from_1d_size(3) with self.assertRaisesRegex(ValueError, "Only the final"): - AnnotatedArray.ArrayDescriptor((scalar_group, vector_group)) + annotated_array.ArrayDescriptor((scalar_group, vector_group)) def test_array_metadata_controls_extension_access(self) -> None: - extension = AnnotatedArray.ExtensionRecord("org.nion.test", 1, "value=42") - metadata = AnnotatedArray.ArrayMetadata(extensions=(extension,)) + extension = annotated_array.ExtensionRecord("org.nion.test", 1, "value=42") + metadata = annotated_array.ArrayMetadata(extensions=(extension,)) self.assertEqual("org.nion.test", extension.extension_type_id) self.assertEqual(("org.nion.test",), metadata.extension_type_ids) @@ -101,7 +101,7 @@ def test_array_metadata_controls_extension_access(self) -> None: with self.assertRaises(KeyError): metadata.get_extension("org.nion.missing") - replacement = AnnotatedArray.ExtensionRecord("org.nion.test", 2, "value=43") + replacement = annotated_array.ExtensionRecord("org.nion.test", 2, "value=43") replaced_metadata = metadata.with_extension(replacement) self.assertEqual(2, replaced_metadata.get_extension("org.nion.test").schema_version) self.assertEqual(1, metadata.get_extension("org.nion.test").schema_version) @@ -109,20 +109,20 @@ def test_array_metadata_controls_extension_access(self) -> None: def test_array_metadata_validates_extensions(self) -> None: with self.assertRaisesRegex(ValueError, "must not be empty"): - AnnotatedArray.ExtensionRecord("", 1, "") + annotated_array.ExtensionRecord("", 1, "") with self.assertRaisesRegex(ValueError, "must be positive"): - AnnotatedArray.ExtensionRecord("org.nion.test", 0, "") + annotated_array.ExtensionRecord("org.nion.test", 0, "") with self.assertRaisesRegex(TypeError, "must be str"): - AnnotatedArray.ExtensionRecord("org.nion.test", 1, bytearray()) # type: ignore[arg-type] + annotated_array.ExtensionRecord("org.nion.test", 1, bytearray()) # type: ignore[arg-type] - extension = AnnotatedArray.ExtensionRecord("org.nion.test", 1, "") + extension = annotated_array.ExtensionRecord("org.nion.test", 1, "") with self.assertRaisesRegex(ValueError, "Duplicate"): - AnnotatedArray.ArrayMetadata(extensions=(extension, extension)) + annotated_array.ArrayMetadata(extensions=(extension, extension)) def test_array_metadata_snapshots_attributes_and_requires_timezone(self) -> None: source = {"note": "original"} created = datetime.datetime(2026, 7, 16, tzinfo=datetime.timezone.utc) - metadata = AnnotatedArray.ArrayMetadata(created=created, attributes=source) + metadata = annotated_array.ArrayMetadata(created=created, attributes=source) source["note"] = "changed" self.assertEqual("original", metadata.attributes["note"]) @@ -131,11 +131,11 @@ def test_array_metadata_snapshots_attributes_and_requires_timezone(self) -> None with self.assertRaises(TypeError): metadata.attributes["note"] = "changed" # type: ignore[index] with self.assertRaisesRegex(ValueError, "timezone-aware"): - AnnotatedArray.ArrayMetadata(created=datetime.datetime(2026, 7, 16)) + annotated_array.ArrayMetadata(created=datetime.datetime(2026, 7, 16)) def test_array_metadata_created_retains_iana_zone_and_offset(self) -> None: created = datetime.datetime(2026, 7, 16, 12, tzinfo=zoneinfo.ZoneInfo("America/Los_Angeles")) - metadata = AnnotatedArray.ArrayMetadata(created=created) + metadata = annotated_array.ArrayMetadata(created=created) self.assertEqual("America/Los_Angeles", getattr(metadata.created.tzinfo, "key", None)) self.assertEqual(datetime.timedelta(hours=-7), metadata.created.utcoffset()) @@ -144,21 +144,21 @@ def test_array_metadata_default_created_has_iana_timezone(self) -> None: # tzlocal.get_localzone() returns a zoneinfo.ZoneInfo with a proper IANA key # (e.g. "America/New_York"). A plain datetime.timezone fixed-offset object # would not satisfy the isinstance check and would have no .key attribute. - metadata = AnnotatedArray.ArrayMetadata() + metadata = annotated_array.ArrayMetadata() self.assertIsInstance(metadata.created.tzinfo, zoneinfo.ZoneInfo) self.assertIsNotNone(metadata.created.tzinfo.key) # type: ignore[union-attr] def test_array_header_can_be_passed_without_data(self) -> None: - descriptor = AnnotatedArray.ArrayDescriptor((AnnotatedArray.AxisGroup.from_2d_size((4, 5)),)) - metadata = AnnotatedArray.ArrayMetadata(attributes={"note": "test"}) - header = AnnotatedArray.ArrayHeader(descriptor, "float32", metadata) + descriptor = annotated_array.ArrayDescriptor((annotated_array.AxisGroup.from_2d_size((4, 5)),)) + metadata = annotated_array.ArrayMetadata(attributes={"note": "test"}) + header = annotated_array.ArrayHeader(descriptor, "float32", metadata) self.assertEqual((4, 5), header.shape) self.assertEqual(numpy.dtype(numpy.float32), header.dtype) - first = AnnotatedArray.AnnotatedArray.from_header(numpy.zeros(header.shape, dtype=header.dtype), header) - second = AnnotatedArray.AnnotatedArray.from_header(numpy.ones(header.shape, dtype=header.dtype), header) + first = annotated_array.AnnotatedArray.from_header(numpy.zeros(header.shape, dtype=header.dtype), header) + second = annotated_array.AnnotatedArray.from_header(numpy.ones(header.shape, dtype=header.dtype), header) self.assertEqual(header, first.header) self.assertIsNot(header, first.header) self.assertEqual(first.header, second.header) @@ -166,24 +166,24 @@ def test_array_header_can_be_passed_without_data(self) -> None: self.assertIs(metadata, first.metadata) def test_annotated_array_validates_data_against_header(self) -> None: - descriptor = AnnotatedArray.ArrayDescriptor((AnnotatedArray.AxisGroup.from_1d_size(3),)) - header = AnnotatedArray.ArrayHeader(descriptor, numpy.float32) + descriptor = annotated_array.ArrayDescriptor((annotated_array.AxisGroup.from_1d_size(3),)) + header = annotated_array.ArrayHeader(descriptor, numpy.float32) with self.assertRaisesRegex(ValueError, "shape"): - AnnotatedArray.AnnotatedArray(numpy.zeros((4,), dtype=numpy.float32), descriptor) + annotated_array.AnnotatedArray(numpy.zeros((4,), dtype=numpy.float32), descriptor) with self.assertRaisesRegex(ValueError, "dtype"): - AnnotatedArray.AnnotatedArray.from_header(numpy.zeros((3,), dtype=numpy.float64), header) + annotated_array.AnnotatedArray.from_header(numpy.zeros((3,), dtype=numpy.float64), header) def test_zeros_annotated_array_constructs_matching_shape_and_dtype(self) -> None: - group = AnnotatedArray.AxisGroup.from_2d_size((2, 3)) - array = AnnotatedArray.zeros_annotated_array((group,), dtype=numpy.float32) + group = annotated_array.AxisGroup.from_2d_size((2, 3)) + array = annotated_array.zeros_annotated_array((group,), dtype=numpy.float32) self.assertEqual((2, 3), array.data.shape) self.assertEqual(numpy.dtype(numpy.float32), array.header.dtype) def test_annotated_array_is_numpy_passable(self) -> None: - group = AnnotatedArray.AxisGroup.from_1d_size(4) - array = AnnotatedArray.zeros_annotated_array((group,), dtype=numpy.float64) + group = annotated_array.AxisGroup.from_1d_size(4) + array = annotated_array.zeros_annotated_array((group,), dtype=numpy.float64) # AnnotatedArray itself is directly usable with numpy functions via __array__ self.assertEqual(0.0, float(numpy.sum(array))) @@ -193,8 +193,8 @@ def test_annotated_array_is_numpy_passable(self) -> None: self.assertEqual(numpy.dtype(numpy.float64), as_array.dtype) def test_annotated_array_data_is_numpy_passable(self) -> None: - group = AnnotatedArray.AxisGroup.from_2d_size((2, 3)) - array = AnnotatedArray.zeros_annotated_array((group,), dtype=numpy.float32) + group = annotated_array.AxisGroup.from_2d_size((2, 3)) + array = annotated_array.zeros_annotated_array((group,), dtype=numpy.float32) # data satisfies ArrayProtocol including __array__, so it is directly usable with numpy self.assertEqual(0.0, float(numpy.sum(array.data))) @@ -209,8 +209,8 @@ def test_annotated_array_accepts_h5py_dataset_as_data(self) -> None: buf = io.BytesIO() with h5py.File(buf, "w") as f: ds = f.create_dataset("data", data=numpy.arange(6, dtype=numpy.float32).reshape(2, 3)) - group = AnnotatedArray.AxisGroup.from_2d_size((2, 3)) - annotated = AnnotatedArray.AnnotatedArray(ds, AnnotatedArray.ArrayDescriptor((group,))) + group = annotated_array.AxisGroup.from_2d_size((2, 3)) + annotated = annotated_array.AnnotatedArray(ds, annotated_array.ArrayDescriptor((group,))) # shape and dtype are read directly from the dataset without materialising it self.assertEqual((2, 3), annotated.data.shape) @@ -223,12 +223,12 @@ def test_annotated_array_accepts_h5py_dataset_as_data(self) -> None: numpy.testing.assert_array_equal(result, numpy.arange(6, dtype=numpy.float32).reshape(2, 3)) def test_from_data_and_metadata_covers_sequence_collection_and_datum_variations(self) -> None: - def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) -> tuple[AnnotatedArray.AffineCalibration, ...]: - calibrations = list[AnnotatedArray.AffineCalibration]() - for axis_group in annotated_array.descriptor.axis_groups: + def flatten_affine_calibrations(array_value: annotated_array.AnnotatedArray) -> tuple[annotated_array.AffineCalibration, ...]: + calibrations = list[annotated_array.AffineCalibration]() + for axis_group in array_value.descriptor.axis_groups: for axis_index in range(axis_group.rank): - calibration = axis_group.get_calibration(axis_index) if axis_group.primary_calibration_key else AnnotatedArray.AffineCalibration() - calibrations.append(typing.cast(AnnotatedArray.AffineCalibration, calibration)) + calibration = axis_group.get_calibration(axis_index) if axis_group.primary_calibration_key else annotated_array.AffineCalibration() + calibrations.append(typing.cast(annotated_array.AffineCalibration, calibration)) return tuple(calibrations) for is_sequence, collection_rank, datum_rank in self._descriptor_variants: @@ -248,11 +248,11 @@ def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) timezone_offset="-0700", ) - annotated = AnnotatedArray.from_data_and_metadata(xdata) + annotated = annotated_array.from_data_and_metadata(xdata) expected_axis_group_ranks = (1, datum_rank) if is_sequence and collection_rank == 0 else (1, collection_rank, datum_rank) if is_sequence else (datum_rank,) if collection_rank == 0 else (collection_rank, datum_rank) self.assertEqual(expected_axis_group_ranks, tuple(axis_group.rank for axis_group in annotated.descriptor.axis_groups)) self.assertEqual(xdata.data_shape, annotated.descriptor.shape) - self.assertEqual("counts", typing.cast(AnnotatedArray.AffineCalibration, annotated.get_intensity_calibration()).unit) + self.assertEqual("counts", typing.cast(annotated_array.AffineCalibration, annotated.get_intensity_calibration()).unit) flattened_annotated_calibrations = flatten_affine_calibrations(annotated) self.assertEqual(len(xdata.dimensional_calibrations), len(flattened_annotated_calibrations)) @@ -264,13 +264,45 @@ def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) self.assertEqual(xdata.metadata["descriptor_variant"], annotated.metadata.attributes["descriptor_variant"]) + def test_collapse_scalar_axis_groups_drops_rank_zero_groups(self) -> None: + xdata = DataAndMetadata.new_data_and_metadata( + data=numpy.arange(16, dtype=numpy.float32), + data_descriptor=DataAndMetadata.DataDescriptor(True, 0, 0), + ) + annotated = annotated_array.from_data_and_metadata(xdata) + self.assertEqual((1, 0), tuple(axis_group.rank for axis_group in annotated.descriptor.axis_groups)) + + collapsed = annotated_array.collapse_scalar_axis_groups(annotated) + self.assertEqual((1,), tuple(axis_group.rank for axis_group in collapsed.descriptor.axis_groups)) + self.assertEqual(annotated.descriptor.shape, collapsed.descriptor.shape) + + def test_collapse_scalar_axis_groups_keeps_single_scalar_group(self) -> None: + scalar_descriptor = annotated_array.ArrayDescriptor((annotated_array.AxisGroup(),)) + scalar_array = annotated_array.AnnotatedArray(data=numpy.asarray(5.0), descriptor=scalar_descriptor) + + collapsed = annotated_array.collapse_scalar_axis_groups(scalar_array) + self.assertEqual(1, len(collapsed.descriptor.axis_groups)) + self.assertEqual(0, collapsed.descriptor.axis_groups[0].rank) + + def test_from_data_and_metadata_optional_collapse_scalar_axis_groups(self) -> None: + xdata = DataAndMetadata.new_data_and_metadata( + data=numpy.arange(16, dtype=numpy.float32), + data_descriptor=DataAndMetadata.DataDescriptor(True, 0, 0), + ) + + annotated = annotated_array.from_data_and_metadata(xdata) + collapsed = annotated_array.from_data_and_metadata(xdata, collapse_scalar_axis_groups=True) + + self.assertEqual((1, 0), tuple(axis_group.rank for axis_group in annotated.descriptor.axis_groups)) + self.assertEqual((1,), tuple(axis_group.rank for axis_group in collapsed.descriptor.axis_groups)) + def test_to_data_and_metadata_covers_sequence_collection_and_datum_variations(self) -> None: - def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) -> tuple[AnnotatedArray.AffineCalibration, ...]: - calibrations = list[AnnotatedArray.AffineCalibration]() - for axis_group in annotated_array.descriptor.axis_groups: + def flatten_affine_calibrations(array_value: annotated_array.AnnotatedArray) -> tuple[annotated_array.AffineCalibration, ...]: + calibrations = list[annotated_array.AffineCalibration]() + for axis_group in array_value.descriptor.axis_groups: for axis_index in range(axis_group.rank): - calibration = axis_group.get_calibration(axis_index) if axis_group.primary_calibration_key else AnnotatedArray.AffineCalibration() - calibrations.append(typing.cast(AnnotatedArray.AffineCalibration, calibration)) + calibration = axis_group.get_calibration(axis_index) if axis_group.primary_calibration_key else annotated_array.AffineCalibration() + calibrations.append(typing.cast(annotated_array.AffineCalibration, calibration)) return tuple(calibrations) for is_sequence, collection_rank, datum_rank in self._descriptor_variants: @@ -283,42 +315,42 @@ def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) dim_sizes = tuple(range(2, 2 + sum(expected_axis_group_ranks))) dim_index = 0 - axis_groups = list[AnnotatedArray.AxisGroup]() + axis_groups = list[annotated_array.AxisGroup]() for axis_group_rank in expected_axis_group_ranks: - axes = list[AnnotatedArray.Axis]() - calibrations = list[AnnotatedArray.AffineCalibration]() + axes = list[annotated_array.Axis]() + calibrations = list[annotated_array.AffineCalibration]() for _ in range(axis_group_rank): size = dim_sizes[dim_index] calibration_index = dim_index - axes.append(AnnotatedArray.Axis(label=f"a{calibration_index}", size=size)) - calibrations.append(AnnotatedArray.AffineCalibration(offset=calibration_index + 0.5, scale=calibration_index + 1.25, unit=f"ua{calibration_index}")) + axes.append(annotated_array.Axis(label=f"a{calibration_index}", size=size)) + calibrations.append(annotated_array.AffineCalibration(offset=calibration_index + 0.5, scale=calibration_index + 1.25, unit=f"ua{calibration_index}")) dim_index += 1 axis_groups.append( - AnnotatedArray.AxisGroup( + annotated_array.AxisGroup( axes=tuple(axes), - coordinate_calibrations={"calibrated": AnnotatedArray.CoordinateCalibration(calibrations=tuple(calibrations))}, + coordinate_calibrations={"calibrated": annotated_array.CoordinateCalibration(calibrations=tuple(calibrations))}, primary_calibration_key="calibrated", ) ) - descriptor = AnnotatedArray.ArrayDescriptor( + descriptor = annotated_array.ArrayDescriptor( axis_groups=tuple(axis_groups), - intensity_calibrations=AnnotatedArray.CalibrationSet.from_calibration( - AnnotatedArray.AffineCalibration(offset=4.0, scale=2.0, unit="counts"), + intensity_calibrations=annotated_array.CalibrationSet.from_calibration( + annotated_array.AffineCalibration(offset=4.0, scale=2.0, unit="counts"), "calibrated", ), ) - metadata = AnnotatedArray.ArrayMetadata( + metadata = annotated_array.ArrayMetadata( created=datetime.datetime(2026, 7, 16, 12, tzinfo=zoneinfo.ZoneInfo("America/Los_Angeles")), attributes={"descriptor_variant": f"annotated-{int(is_sequence)}-{collection_rank}-{datum_rank}"}, ) - annotated = AnnotatedArray.AnnotatedArray( + annotated = annotated_array.AnnotatedArray( data=numpy.arange(int(numpy.prod(dim_sizes)), dtype=numpy.float32).reshape(dim_sizes), descriptor=descriptor, metadata=metadata, ) - xdata = AnnotatedArray.to_data_and_metadata(annotated) + xdata = annotated_array.to_data_and_metadata(annotated) self.assertEqual(is_sequence, xdata.is_sequence) self.assertEqual(collection_rank, xdata.collection_dimension_count) self.assertEqual(datum_rank, xdata.datum_dimension_count) @@ -333,7 +365,7 @@ def flatten_affine_calibrations(annotated_array: AnnotatedArray.AnnotatedArray) self.assertEqual(annotated_calibration.scale, legacy_calibration.scale) self.assertEqual(annotated_calibration.unit, legacy_calibration.units) - round_tripped = AnnotatedArray.from_data_and_metadata(xdata) + round_tripped = annotated_array.from_data_and_metadata(xdata) self.assertEqual(tuple(axis_group.rank for axis_group in annotated.descriptor.axis_groups), tuple(axis_group.rank for axis_group in round_tripped.descriptor.axis_groups)) self.assertEqual(annotated.descriptor.shape, round_tripped.descriptor.shape) self.assertEqual(annotated.metadata.attributes["descriptor_variant"], round_tripped.metadata.attributes["descriptor_variant"]) @@ -354,18 +386,15 @@ def test_data_and_metadata_timezone_round_trip_through_annotated_array(self) -> timezone_offset="-0700", ) - annotated = AnnotatedArray.from_data_and_metadata(xdata) + annotated = annotated_array.from_data_and_metadata(xdata) self.assertEqual("America/Los_Angeles", getattr(annotated.metadata.created.tzinfo, "key", None)) self.assertEqual(datetime.timedelta(hours=-7), annotated.metadata.created.utcoffset()) - round_tripped = AnnotatedArray.to_data_and_metadata(annotated) + round_tripped = annotated_array.to_data_and_metadata(annotated) self.assertEqual(xdata.timestamp, round_tripped.timestamp) self.assertEqual("America/Los_Angeles", round_tripped.timezone) self.assertEqual("-0700", round_tripped.timezone_offset) - if __name__ == "__main__": unittest.main() - - diff --git a/nion/data/test/annotated_array_primitives_test.py b/nion/data/test/annotated_array_primitives_test.py new file mode 100644 index 0000000..032a804 --- /dev/null +++ b/nion/data/test/annotated_array_primitives_test.py @@ -0,0 +1,442 @@ +"""Tests for annotated_array.primitives — FFT / IFFT. + +These tests are the canonical human-readable specification of the expected +behaviour. Each test focuses on one observable property and is written to +be self-explanatory without requiring any other context. +""" + +import math +import typing +import unittest + +import numpy +import numpy.testing + +from nion.data import annotated_array +from nion.data.annotated_array import primitives + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_1d_array( + data: numpy.ndarray, + scale: float = 1.0, + offset: float = 0.0, + unit: str = "", +) -> annotated_array.AnnotatedArray: + """Return a 1-D AnnotatedArray with a single affine spatial calibration. + + The value_type is inferred from the array dtype so that both real and + complex test inputs are accepted without additional boilerplate. + """ + calibration = annotated_array.CoordinateCalibration( + calibrations=(annotated_array.AffineCalibration(scale=scale, offset=offset, unit=unit),) + ) + signal_group = annotated_array.AxisGroup.from_1d_size( + data.shape[0], + coordinate_calibrations={"spatial": calibration}, + primary_calibration_key="spatial", + ) + value_type = annotated_array.infer_value_type(data.dtype) + descriptor = annotated_array.ArrayDescriptor((signal_group,), value_type=value_type) + return annotated_array.AnnotatedArray(data=data, descriptor=descriptor) + + +def _make_2d_array( + data: numpy.ndarray, + scale: tuple[float, float] = (1.0, 1.0), + unit: str = "", +) -> annotated_array.AnnotatedArray: + """Return a 2-D AnnotatedArray with a single affine spatial calibration. + + The value_type is inferred from the array dtype. + """ + calibration = annotated_array.CoordinateCalibration( + calibrations=( + annotated_array.AffineCalibration(scale=scale[0], unit=unit), + annotated_array.AffineCalibration(scale=scale[1], unit=unit), + ) + ) + signal_group = annotated_array.AxisGroup.from_2d_size( + (data.shape[0], data.shape[1]), + coordinate_calibrations={"spatial": calibration}, + primary_calibration_key="spatial", + ) + value_type = annotated_array.infer_value_type(data.dtype) + descriptor = annotated_array.ArrayDescriptor((signal_group,), value_type=value_type) + return annotated_array.AnnotatedArray(data=data, descriptor=descriptor) + + +# --------------------------------------------------------------------------- +# FFT — output data +# --------------------------------------------------------------------------- + +class TestFftOutputData(unittest.TestCase): + + def test_fft_1d_output_is_complex(self) -> None: + """The output of fft on a real 1-D array must be complex.""" + src = _make_1d_array(numpy.ones(8, dtype=numpy.float64)) + result = primitives.fft(src) + self.assertTrue(numpy.iscomplexobj(result.data)) + + def test_fft_2d_output_is_complex(self) -> None: + """The output of fft on a real 2-D array must be complex.""" + src = _make_2d_array(numpy.ones((4, 8), dtype=numpy.float64)) + result = primitives.fft(src) + self.assertTrue(numpy.iscomplexobj(result.data)) + + def test_fft_1d_preserves_rms_energy(self) -> None: + """RMS of input equals RMS of output (Parseval / energy-normalised FFT).""" + rng = numpy.random.default_rng(0) + data = rng.standard_normal(64) + src = _make_1d_array(data) + result = primitives.fft(src) + rms_in = numpy.sqrt(numpy.mean(numpy.abs(data) ** 2)) + rms_out = numpy.sqrt(numpy.mean(numpy.abs(result.data) ** 2)) + numpy.testing.assert_allclose(rms_in, rms_out, rtol=1e-12) + + def test_fft_2d_preserves_rms_energy(self) -> None: + """RMS of input equals RMS of output (Parseval / energy-normalised FFT).""" + rng = numpy.random.default_rng(1) + data = rng.standard_normal((16, 32)) + src = _make_2d_array(data) + result = primitives.fft(src) + rms_in = numpy.sqrt(numpy.mean(numpy.abs(data) ** 2)) + rms_out = numpy.sqrt(numpy.mean(numpy.abs(result.data) ** 2)) + numpy.testing.assert_allclose(rms_in, rms_out, rtol=1e-12) + + def test_fft_1d_dc_component_is_at_centre(self) -> None: + """For a constant (DC-only) signal the single non-zero bin must be at the array centre.""" + n = 16 + data = numpy.ones(n, dtype=numpy.float64) + src = _make_1d_array(data) + result = primitives.fft(src) + magnitude = numpy.abs(numpy.asarray(result.data)) + centre = n // 2 + # N ones scaled by 1/sqrt(N) → DC bin = N/sqrt(N) = sqrt(N). + self.assertAlmostEqual(magnitude[centre], math.sqrt(n)) + # All other bins must be (essentially) zero. + mask = numpy.ones(n, dtype=bool) + mask[centre] = False + numpy.testing.assert_allclose(magnitude[mask], 0.0, atol=1e-12) + + def test_fft_2d_dc_component_is_at_centre(self) -> None: + """For a constant 2-D image the only non-zero bin must be at the image centre.""" + rows, cols = 8, 12 + data = numpy.ones((rows, cols), dtype=numpy.float64) + src = _make_2d_array(data) + result = primitives.fft(src) + magnitude = numpy.abs(numpy.asarray(result.data)) + cy, cx = rows // 2, cols // 2 + self.assertGreater(magnitude[cy, cx], 0.5) + # All other bins must be (essentially) zero. + mask = numpy.ones((rows, cols), dtype=bool) + mask[cy, cx] = False + numpy.testing.assert_allclose(magnitude[mask], 0.0, atol=1e-12) + + +# --------------------------------------------------------------------------- +# FFT — output calibrations +# --------------------------------------------------------------------------- + +class TestFftOutputCalibrations(unittest.TestCase): + + def test_fft_1d_frequency_scale_is_reciprocal(self) -> None: + """freq_scale = 1 / (spatial_scale * N).""" + n = 32 + s = 0.5 # spatial scale (e.g. 0.5 nm/pixel) + src = _make_1d_array(numpy.zeros(n), scale=s, unit="nm") + result = primitives.fft(src) + + signal_group = result.descriptor.axis_groups[-1] + freq_cal = signal_group.get_calibration(0) # primary (frequency) calibration + self.assertIsInstance(freq_cal, annotated_array.AffineCalibration) + freq_aff = typing.cast(annotated_array.AffineCalibration, freq_cal) + expected_scale = 1.0 / (s * n) + self.assertAlmostEqual(freq_aff.scale, expected_scale) + + def test_fft_1d_frequency_offset_places_dc_at_centre(self) -> None: + """freq_offset = (-0.5 - N//2) / (spatial_scale * N).""" + n = 32 + s = 0.5 + src = _make_1d_array(numpy.zeros(n), scale=s, unit="nm") + result = primitives.fft(src) + + signal_group = result.descriptor.axis_groups[-1] + freq_cal = signal_group.get_calibration(0) + self.assertIsInstance(freq_cal, annotated_array.AffineCalibration) + freq_aff = typing.cast(annotated_array.AffineCalibration, freq_cal) + expected_offset = (-0.5 - n // 2) / (s * n) + self.assertAlmostEqual(freq_aff.offset, expected_offset) + + def test_fft_1d_frequency_unit_is_reciprocal(self) -> None: + """The frequency unit is '1/'.""" + src = _make_1d_array(numpy.zeros(16), scale=1.0, unit="nm") + result = primitives.fft(src) + signal_group = result.descriptor.axis_groups[-1] + freq_cal = signal_group.get_calibration(0) + self.assertIsInstance(freq_cal, annotated_array.AffineCalibration) + freq_aff = typing.cast(annotated_array.AffineCalibration, freq_cal) + self.assertEqual("1/nm", freq_aff.unit) + + def test_fft_2d_frequency_unit_is_reciprocal(self) -> None: + """Both axes of a 2-D FFT result have reciprocal units.""" + src = _make_2d_array(numpy.zeros((8, 12)), scale=(1.0, 1.0), unit="nm") + result = primitives.fft(src) + signal_group = result.descriptor.axis_groups[-1] + freq_cal_row = signal_group.get_calibration(0) + freq_cal_col = signal_group.get_calibration(1) + self.assertIsInstance(freq_cal_row, annotated_array.AffineCalibration) + self.assertIsInstance(freq_cal_col, annotated_array.AffineCalibration) + freq_row_aff = typing.cast(annotated_array.AffineCalibration, freq_cal_row) + freq_col_aff = typing.cast(annotated_array.AffineCalibration, freq_cal_col) + self.assertEqual("1/nm", freq_row_aff.unit) + self.assertEqual("1/nm", freq_col_aff.unit) + + def test_fft_2d_frequency_calibration_applied_to_each_axis(self) -> None: + """Both axes of a 2-D FFT result have independent frequency calibrations.""" + rows, cols = 8, 16 + s_row, s_col = 0.25, 0.5 + src = _make_2d_array(numpy.zeros((rows, cols)), scale=(s_row, s_col), unit="nm") + result = primitives.fft(src) + + signal_group = result.descriptor.axis_groups[-1] + freq_cal_row = signal_group.get_calibration(0) + freq_cal_col = signal_group.get_calibration(1) + self.assertIsInstance(freq_cal_row, annotated_array.AffineCalibration) + self.assertIsInstance(freq_cal_col, annotated_array.AffineCalibration) + freq_row_aff = typing.cast(annotated_array.AffineCalibration, freq_cal_row) + freq_col_aff = typing.cast(annotated_array.AffineCalibration, freq_cal_col) + self.assertAlmostEqual(freq_row_aff.scale, 1.0 / (s_row * rows)) + self.assertAlmostEqual(freq_col_aff.scale, 1.0 / (s_col * cols)) + + def test_fft_primary_calibration_key_is_preserved(self) -> None: + """After FFT the primary coordinate calibration key is unchanged.""" + src = _make_1d_array(numpy.zeros(16), scale=1.0, unit="nm") + result = primitives.fft(src) + signal_group = result.descriptor.axis_groups[-1] + self.assertEqual("spatial", signal_group.primary_calibration_key) + + def test_fft_calibration_key_has_frequency_values(self) -> None: + """After FFT each calibration key is present with frequency-domain values.""" + src = _make_1d_array(numpy.zeros(16), scale=0.3, unit="nm") + result = primitives.fft(src) + signal_group = result.descriptor.axis_groups[-1] + self.assertIn("spatial", signal_group.calibration_keys) + spatial_cal = signal_group.get_calibration(0, key="spatial") + self.assertIsInstance(spatial_cal, annotated_array.AffineCalibration) + spatial_aff = typing.cast(annotated_array.AffineCalibration, spatial_cal) + self.assertAlmostEqual(spatial_aff.scale, 1.0 / (0.3 * 16)) + self.assertEqual("1/nm", spatial_aff.unit) + + def test_fft_intensity_calibration_is_unchanged(self) -> None: + """FFT must not alter the intensity calibration.""" + from nion.data import Calibration + from nion.data import DataAndMetadata + xdata = DataAndMetadata.new_data_and_metadata( + numpy.zeros(16), + intensity_calibration=Calibration.Calibration(0.0, 2.5, "counts"), + ) + aa = annotated_array.from_data_and_metadata(xdata) + result = primitives.fft(aa) + intensity = result.descriptor.get_intensity_calibration() + self.assertIsInstance(intensity, annotated_array.AffineCalibration) + intensity_aff = typing.cast(annotated_array.AffineCalibration, intensity) + self.assertAlmostEqual(intensity_aff.scale, 2.5) + self.assertEqual("counts", intensity_aff.unit) + + +# --------------------------------------------------------------------------- +# FFT — input validation +# --------------------------------------------------------------------------- + +class TestFftInputValidation(unittest.TestCase): + + def test_fft_rejects_multiple_axis_groups(self) -> None: + """FFT requires exactly one axis group.""" + from nion.data.annotated_array._implementation import Axis + navigation_group = annotated_array.AxisGroup(axes=(Axis("n", 4),)) + signal_group = annotated_array.AxisGroup(axes=(Axis("x", 8),)) + descriptor = annotated_array.ArrayDescriptor((navigation_group, signal_group)) + array = annotated_array.AnnotatedArray(data=numpy.zeros((4, 8)), descriptor=descriptor) + with self.assertRaises(ValueError): + primitives.fft(array) + + def test_fft_rejects_signal_rank_other_than_1_or_2(self) -> None: + """FFT on a 3-D signal axis group must raise ValueError.""" + from nion.data.annotated_array._implementation import Axis + signal_group = annotated_array.AxisGroup(axes=(Axis("x", 4), Axis("y", 4), Axis("z", 4))) + descriptor = annotated_array.ArrayDescriptor((signal_group,)) + array = annotated_array.AnnotatedArray(data=numpy.zeros((4, 4, 4)), descriptor=descriptor) + with self.assertRaises(ValueError): + primitives.fft(array) + + def test_fft_accepts_complex_input(self) -> None: + """FFT on a complex input must succeed and produce a complex output.""" + data = numpy.ones(8, dtype=numpy.complex128) + src = _make_1d_array(data) + result = primitives.fft(src) + self.assertTrue(numpy.iscomplexobj(result.data)) + + +# --------------------------------------------------------------------------- +# IFFT — output data +# --------------------------------------------------------------------------- + +class TestIfftOutputData(unittest.TestCase): + + def test_ifft_1d_round_trip_recovers_original_data(self) -> None: + """ifft(fft(x)) must recover x to floating-point precision.""" + rng = numpy.random.default_rng(2) + data = rng.standard_normal(64) + src = _make_1d_array(data) + recovered = primitives.ifft(primitives.fft(src)) + numpy.testing.assert_allclose(numpy.real(recovered.data), data, atol=1e-12) + + +class TestIfftInputValidation(unittest.TestCase): + + def test_ifft_rejects_multiple_axis_groups(self) -> None: + """IFFT requires exactly one axis group.""" + from nion.data.annotated_array._implementation import Axis + navigation_group = annotated_array.AxisGroup(axes=(Axis("n", 4),)) + signal_group = annotated_array.AxisGroup(axes=(Axis("x", 8),)) + descriptor = annotated_array.ArrayDescriptor((navigation_group, signal_group), value_type=annotated_array.ValueType.COMPLEX) + array = annotated_array.AnnotatedArray(data=numpy.zeros((4, 8), dtype=numpy.complex128), descriptor=descriptor) + with self.assertRaises(ValueError): + primitives.ifft(array) + + def test_ifft_2d_round_trip_recovers_original_data(self) -> None: + """ifft(fft(x)) must recover x to floating-point precision.""" + rng = numpy.random.default_rng(3) + data = rng.standard_normal((16, 32)) + src = _make_2d_array(data) + recovered = primitives.ifft(primitives.fft(src)) + numpy.testing.assert_allclose(numpy.real(recovered.data), data, atol=1e-12) + + +# --------------------------------------------------------------------------- +# IFFT — calibration round-trip +# --------------------------------------------------------------------------- + +class TestIfftCalibrationRoundTrip(unittest.TestCase): + + def test_ifft_restores_spatial_scale(self) -> None: + """After ifft(fft(x)) the signal AxisGroup primary scale must equal the original.""" + s = 0.4 + src = _make_1d_array(numpy.zeros(32), scale=s, unit="nm") + recovered = primitives.ifft(primitives.fft(src)) + signal_group = recovered.descriptor.axis_groups[-1] + cal = signal_group.get_calibration(0) + self.assertIsInstance(cal, annotated_array.AffineCalibration) + cal_aff = typing.cast(annotated_array.AffineCalibration, cal) + self.assertAlmostEqual(cal_aff.scale, s) + + def test_ifft_restores_spatial_unit(self) -> None: + """After ifft(fft(x)) the primary unit must match the original unit.""" + src = _make_1d_array(numpy.zeros(32), scale=0.4, unit="nm") + recovered = primitives.ifft(primitives.fft(src)) + signal_group = recovered.descriptor.axis_groups[-1] + cal = signal_group.get_calibration(0) + self.assertIsInstance(cal, annotated_array.AffineCalibration) + cal_aff = typing.cast(annotated_array.AffineCalibration, cal) + self.assertEqual("nm", cal_aff.unit) + + def test_ifft_restores_original_calibration_key(self) -> None: + """The primary calibration key is restored to its pre-FFT name.""" + src = _make_1d_array(numpy.zeros(32), scale=0.4, unit="nm") + recovered = primitives.ifft(primitives.fft(src)) + signal_group = recovered.descriptor.axis_groups[-1] + self.assertEqual("spatial", signal_group.primary_calibration_key) + + def test_ifft_2d_restores_both_axes(self) -> None: + """Both axis calibrations are recovered correctly after a 2-D round-trip.""" + rows, cols = 8, 16 + s_row, s_col = 0.25, 0.5 + src = _make_2d_array(numpy.zeros((rows, cols)), scale=(s_row, s_col), unit="nm") + recovered = primitives.ifft(primitives.fft(src)) + signal_group = recovered.descriptor.axis_groups[-1] + cal_row = signal_group.get_calibration(0) + cal_col = signal_group.get_calibration(1) + self.assertIsInstance(cal_row, annotated_array.AffineCalibration) + self.assertIsInstance(cal_col, annotated_array.AffineCalibration) + cal_row_aff = typing.cast(annotated_array.AffineCalibration, cal_row) + cal_col_aff = typing.cast(annotated_array.AffineCalibration, cal_col) + self.assertAlmostEqual(cal_row_aff.scale, s_row) + self.assertAlmostEqual(cal_col_aff.scale, s_col) + + def test_ifft_derived_spatial_scale_from_frequency_calibration(self) -> None: + """When an AnnotatedArray with frequency calibration is constructed directly + (not via fft) ifft must derive the spatial calibration as 1/(scale_freq * N).""" + n = 32 + scale_freq = 0.1 # 1/nm per pixel + freq_cal = annotated_array.CoordinateCalibration( + calibrations=(annotated_array.AffineCalibration(scale=scale_freq, unit="1/nm"),) + ) + signal_group = annotated_array.AxisGroup.from_1d_size( + n, + coordinate_calibrations={"frequency": freq_cal}, + primary_calibration_key="frequency", + ) + descriptor = annotated_array.ArrayDescriptor((signal_group,), value_type=annotated_array.ValueType.COMPLEX) + array = annotated_array.AnnotatedArray( + data=numpy.zeros(n, dtype=numpy.complex128), + descriptor=descriptor, + ) + result = primitives.ifft(array) + spatial_cal = result.descriptor.axis_groups[-1].get_calibration(0) + self.assertIsInstance(spatial_cal, annotated_array.AffineCalibration) + spatial_aff = typing.cast(annotated_array.AffineCalibration, spatial_cal) + expected_scale = 1.0 / (scale_freq * n) + self.assertAlmostEqual(spatial_aff.scale, expected_scale) + self.assertEqual("nm", spatial_aff.unit) + + def test_fft_ifft_preserves_all_calibration_keys(self) -> None: + """FFT and IFFT preserve the exact set of calibration keys (same keys in, same keys out). + + Each key's calibration is independently transformed to frequency domain on FFT + and back to spatial domain on IFFT. + """ + spatial_cal = annotated_array.AffineCalibration(scale=0.5, offset=0.0, unit="nm") + angular_cal = annotated_array.AffineCalibration(scale=0.1, offset=0.0, unit="radians") + coord_cals = { + "spatial": annotated_array.CoordinateCalibration((spatial_cal,)), + "angular": annotated_array.CoordinateCalibration((angular_cal,)), + } + signal_group = annotated_array.AxisGroup.from_1d_size( + 16, coordinate_calibrations=coord_cals, primary_calibration_key="spatial", + ) + descriptor = annotated_array.ArrayDescriptor((signal_group,)) + array = annotated_array.AnnotatedArray(data=numpy.ones(16), descriptor=descriptor) + + # FFT: same two keys, now in frequency domain. + result_fft = primitives.fft(array) + fft_group = result_fft.descriptor.axis_groups[-1] + self.assertEqual(set(fft_group.coordinate_calibrations), {"spatial", "angular"}) + self.assertEqual("spatial", fft_group.primary_calibration_key) + + spatial_freq = typing.cast(annotated_array.AffineCalibration, fft_group.get_calibration(0, key="spatial")) + angular_freq = typing.cast(annotated_array.AffineCalibration, fft_group.get_calibration(0, key="angular")) + self.assertAlmostEqual(spatial_freq.scale, 1.0 / (0.5 * 16)) + self.assertAlmostEqual(angular_freq.scale, 1.0 / (0.1 * 16)) + self.assertEqual("1/nm", spatial_freq.unit) + self.assertEqual("1/radians", angular_freq.unit) + + # IFFT: same two keys, back to spatial domain. + result_ifft = primitives.ifft(result_fft) + ifft_group = result_ifft.descriptor.axis_groups[-1] + self.assertEqual(set(ifft_group.coordinate_calibrations), {"spatial", "angular"}) + self.assertEqual("spatial", ifft_group.primary_calibration_key) + + spatial_back = typing.cast(annotated_array.AffineCalibration, ifft_group.get_calibration(0, key="spatial")) + angular_back = typing.cast(annotated_array.AffineCalibration, ifft_group.get_calibration(0, key="angular")) + self.assertAlmostEqual(spatial_back.scale, 0.5) + self.assertAlmostEqual(angular_back.scale, 0.1) + self.assertEqual("nm", spatial_back.unit) + self.assertEqual("radians", angular_back.unit) + + +if __name__ == "__main__": + unittest.main()