diff --git a/python/image.cpp b/python/image.cpp index 894caadf..1bb71361 100644 --- a/python/image.cpp +++ b/python/image.cpp @@ -214,7 +214,13 @@ ElectronCountedDataPyArray electronCount(Reader* reader, return electronCount(reader, options.toCpp()); } -// Explicitly instantiate version for py::array_t +template std::vector createSTEMImages( + const std::vector>>& sparseData, + const std::vector& innerRadii, const std::vector& outerRadii, + Dimensions2D scanDimensions, Dimensions2D frameDimensions, + CoordinatesDouble2D center); + +// Retain the original exported integer specialization. template std::vector createSTEMImages( const std::vector>>& sparseData, const std::vector& innerRadii, const std::vector& outerRadii, @@ -223,13 +229,26 @@ template std::vector createSTEMImages( } // namespace stempy -vector createSTEMImages(const ElectronCountedDataPyArray& array, +// Only typed double pairs select this overload. Bare brace-initialized centers +// cannot deduce Center and continue to use the original integer overload. +template +std::enable_if_t::value, + vector> createSTEMImages(const ElectronCountedDataPyArray& array, const vector& innerRadii, const vector& outerRadii, - Coordinates2D coords) + Center center) { return createSTEMImages(array.data, innerRadii, outerRadii, - array.scanDimensions, array.frameDimensions, coords); + array.scanDimensions, array.frameDimensions, center); +} + +vector createSTEMImages(const ElectronCountedDataPyArray& array, + const vector& innerRadii, + const vector& outerRadii, + Coordinates2D center) +{ + return createSTEMImages(array, innerRadii, outerRadii, + CoordinatesDouble2D(center)); } template @@ -392,27 +411,27 @@ PYBIND11_MODULE(_image, m) m.def("create_stem_images", (vector(*)(StreamReader::iterator, StreamReader::iterator, const vector&, const vector&, - Dimensions2D, Coordinates2D)) & + Dimensions2D, CoordinatesDouble2D)) & createSTEMImages, py::call_guard()); m.def( "create_stem_images", (vector(*)(SectorStreamReader::iterator, SectorStreamReader::iterator, const vector&, - const vector&, Dimensions2D, Coordinates2D)) & + const vector&, Dimensions2D, CoordinatesDouble2D)) & createSTEMImages, py::call_guard()); m.def("create_stem_images", (vector(*)( const std::vector>>&, const vector&, const vector&, Dimensions2D, Dimensions2D, - Coordinates2D)) & + CoordinatesDouble2D)) & createSTEMImages, py::call_guard()); m.def( "create_stem_images", (vector(*)(const ElectronCountedDataPyArray&, const vector&, - const vector&, Coordinates2D)) & + const vector&, CoordinatesDouble2D)) & createSTEMImages, py::call_guard()); m.def("calculate_average", &calculateAverage, @@ -585,7 +604,7 @@ PYBIND11_MODULE(_image, m) m.def("create_stem_images", (vector(*)(PyReader::iterator, PyReader::iterator, const vector&, const vector&, - Dimensions2D, Coordinates2D)) & + Dimensions2D, CoordinatesDouble2D)) & createSTEMImages, py::call_guard()); m.def("maximum_diffraction_pattern", diff --git a/python/stempy/image/__init__.py b/python/stempy/image/__init__.py index 2f796214..778420f4 100644 --- a/python/stempy/image/__init__.py +++ b/python/stempy/image/__init__.py @@ -38,7 +38,7 @@ def create_stem_images(input, inner_radii, outer_radii, scan_dimensions=(0, 0), :param center: the center of the images, where the order is (x, y). If set to (-1, -1), the center will be set to (scan_dimensions[0] / 2, scan_dimensions[1] / 2). - :type center: tuple of ints of length 2 + :type center: pair of ints or floats, or numpy.ndarray of shape (2, 1) :param frame_dimensions: the dimensions of each frame, where the order is (width, height). Only used for input of type numpy.ndarray, in which case its presence implies @@ -52,6 +52,16 @@ def create_stem_images(input, inner_radii, outer_radii, scan_dimensions=(0, 0), :return: A numpy array of the STEM images. :rtype: numpy.ndarray """ + # Extract scalars explicitly, NumPy no longer converts one-element arrays. + center_array = np.asarray(center, dtype=np.float64) + if center_array.shape not in ((2,), (2, 1)): + raise ValueError('center must have shape (2,) or (2, 1) in (x, y) order') + + if not np.all(np.isfinite(center_array)): + raise ValueError('center coordinates must be finite') + + center = tuple(float(value) for value in center_array.reshape(2)) + # Ensure the inner and outer radii are tuples or lists if not isinstance(inner_radii, (tuple, list)): inner_radii = [inner_radii] @@ -532,7 +542,7 @@ def _com_sparse_v0(array, crop_to=None, init_center=None, replace_nans=True): x = ev // array.frame_shape[0] y = ev % array.frame_shape[1] mm0 = len(ev) - + if init_center is None: # Initialize center as full frame COM comx0 = np.sum(x) / mm0 @@ -540,7 +550,7 @@ def _com_sparse_v0(array, crop_to=None, init_center=None, replace_nans=True): else: comx0 = init_center[0] comy0 = init_center[1] - + if crop_to is not None: # Crop around the initial center r = np.sqrt((x - comx0)**2 + (y - comy0)**2) @@ -561,18 +571,18 @@ def _com_sparse_v0(array, crop_to=None, init_center=None, replace_nans=True): # Center of mass of the full frame comx = np.sum(x) / mm0 comy = np.sum(y) / mm0 - + com[:, scan_position] = (comy, comx) # save the comx and comy. Needs to be reversed (comy, comx) else: com[:, scan_position] = (np.nan, np.nan) # empty frame - + com = com.reshape((2, *array.scan_shape)) - + if replace_nans: com_mean = np.nanmean(com, axis=(1,2)) np.nan_to_num(com[0,:,:], nan=com_mean[0], copy=False) np.nan_to_num(com[1,:,:], nan=com_mean[1], copy=False) - + return com def com_sparse( @@ -806,28 +816,28 @@ def _electron_counted_metadata_to_dict(metadata): def virtual_darkfield(array, centers_x, centers_y, radii): """Calculate a virtual dark field image from a set of round virtual apertures in diffraction space. Each aperture is defined by a center and radius and the final image is the sum of all of them. - + :param array: The SparseArray :type array: SparseArray - + :param centers_x: The center of each round aperture as the row locations :type centers_x: number or iterable - + :param centers_y: The center of each round aperture as the column locations :type centers_y: number or iterable - + :param radii: The radius of each aperture. :type radii: number or iterable - + :rtype: np.ndarray - + :example: >>> sp = stempy.io.load_electron_counts('file.h5') >>> df2 = stempy.image.virtual_darkfield(sp, (288, 260), (288, 160), (10, 10)) # 2 apertures >>> df1 = stempy.image.virtual_darkfield(sp, 260, 160, 10) # 1 aperture - + """ - + # Change to iterable if single value if isinstance(centers_x, (int, float)): centers_x = (centers_x,) @@ -845,30 +855,30 @@ def virtual_darkfield(array, centers_x, centers_y, radii): dist = np.sqrt((ev_rows - cc_1)**2 + (ev_cols - cc_0)**2) rs_image[ii] += len(np.where(dist < rr)[0]) rs_image = rs_image.reshape(array.scan_shape) - + return rs_image def plot_virtual_darkfield(image, centers_x, centers_y, radii, axes=None): """Plot circles on the diffraction pattern corresponding to the position and size of virtual dark field apertures. This has the same center and radii inputs as stempy.image.virtual_darkfield so users can check their input is physically correct. - + :param image: The diffraction pattern to plot over :type image: np.ndarray, 2D - + :param centers_x: The center of each round aperture as the row locations :type centers_x: iterable - + :param centers_y: The center of each round aperture as the column locations :type centers_y: iterable - + :param radii: The radius of each aperture. :type radii: iterable - + :param axes: A matplotlib axes instance to use for the plotting. If None then a new plot is created. :type axes: matplotlib.axes._subplots.AxesSubplot - + :rtype: matplotlib.axes._subplots.AxesSubplot - + :example: >>> sp = stempy.io.load_electron_counts('file.h5') >>> stempy.image.plot_virtual_darkfield(sp.sum(axis=(0, 1), 260, 160, 10) # 1 aperture @@ -884,28 +894,28 @@ def plot_virtual_darkfield(image, centers_x, centers_y, radii, axes=None): centers_y = (centers_y,) if isinstance(radii, (int, float)): radii = (radii,) - + if not axes: fg, axes = plt.subplots(1, 1) - + axes.imshow(image, cmap='magma', norm=LogNorm()) - + # Place a circle at each apertue location for cc_0, cc_1, rr in zip(centers_x, centers_y, radii): C = Circle((cc_0, cc_1), rr, fc='none', ec='c') axes.add_patch(C) - + return axes def mask_real_space(array, mask): """Calculate a diffraction pattern from an arbitrary set of positions defined in a mask in real space - + :param array: The sparse dataset :type array: SparseArray - + :param mask: The mask to apply with 0 for probe positions to ignore and 1 for probe positions to include in the sum. Must have the same scan shape as array :type mask: np.ndarray - + :rtype: np.ndarray """ assert array.scan_shape[0] == mask.shape[0] and array.scan_shape[1] == mask.shape[1] diff --git a/stempy/image.cpp b/stempy/image.cpp index b2bb1179..e5074b7c 100644 --- a/stempy/image.cpp +++ b/stempy/image.cpp @@ -159,12 +159,15 @@ void _runCalculateSTEMValues(const uint16_t data[], } } // end namespace -template -vector createSTEMImages(InputIt first, InputIt last, +// Only typed double pairs select this overload. Bare brace-initialized centers +// cannot deduce Center and continue to use the original integer overload. +template +std::enable_if_t::value, + vector> createSTEMImages(InputIt first, InputIt last, const vector& innerRadii, const vector& outerRadii, Dimensions2D scanDimensions, - Coordinates2D center) + Center center) { if (first == last) { ostringstream msg; @@ -332,10 +335,14 @@ std::vector createSTEMHistogram(const STEMImage& inImage, return frequencies; } -vector createSTEMImages(const ElectronCountedData& data, +// Only typed double pairs select this overload. Bare brace-initialized centers +// cannot deduce Center and continue to use the original integer overload. +template +std::enable_if_t::value, + vector> createSTEMImages(const ElectronCountedData& data, const vector& innerRadii, const vector& outerRadii, - Coordinates2D center) + Center center) { return createSTEMImages(data.data, innerRadii, outerRadii, data.scanDimensions, data.frameDimensions, center); @@ -625,7 +632,50 @@ Image maximumDiffractionPattern(InputIt first, InputIt last) return maximumDiffractionPattern(first, last, dark); } +template +vector createSTEMImages(InputIt first, InputIt last, + const vector& innerRadii, + const vector& outerRadii, + Dimensions2D scanDimensions, + Coordinates2D center) +{ + return createSTEMImages(first, last, innerRadii, outerRadii, + scanDimensions, CoordinatesDouble2D(center)); +} + +vector createSTEMImages(const ElectronCountedData& data, + const vector& innerRadii, + const vector& outerRadii, + Coordinates2D center) +{ + return createSTEMImages(data, innerRadii, outerRadii, + CoordinatesDouble2D(center)); +} + +template vector createSTEMImages(const ElectronCountedData&, + const vector&, const vector&, CoordinatesDouble2D); + // Instantiate the ones that can be used +template vector createSTEMImages( + StreamReader::iterator first, StreamReader::iterator last, + const vector& innerRadii, const vector& outerRadii, + Dimensions2D scanDimensions, CoordinatesDouble2D center); + +template vector createSTEMImages( + PyReader::iterator first, PyReader::iterator last, + const vector& innerRadii, const vector& outerRadii, + Dimensions2D scanDimensions, CoordinatesDouble2D center); + +template vector createSTEMImages::iterator>( + vector::iterator first, vector::iterator last, + const vector& innerRadii, const vector& outerRadii, + Dimensions2D scanDimensions, CoordinatesDouble2D center); + +template vector createSTEMImages( + SectorStreamReader::iterator first, SectorStreamReader::iterator last, + const vector& innerRadii, const vector& outerRadii, + Dimensions2D scanDimensions, CoordinatesDouble2D center); + template vector createSTEMImages( StreamReader::iterator first, StreamReader::iterator last, const vector& innerRadii, const vector& outerRadii, diff --git a/stempy/image.h b/stempy/image.h index bf95608c..0934bb88 100644 --- a/stempy/image.h +++ b/stempy/image.h @@ -5,6 +5,7 @@ #include #include +#include namespace stempy { @@ -48,11 +49,13 @@ namespace stempy { using STEMImage = Image; // Create STEM Images from raw data - template - std::vector createSTEMImages( + // Only typed double pairs select this overload. Bare brace-initialized centers + // cannot deduce Center and continue to use the original integer overload. + template + std::enable_if_t::value, + std::vector> createSTEMImages( InputIt first, InputIt last, const std::vector& innerRadii, - const std::vector& outerRadii, Dimensions2D scanDimensions = { 0, 0 }, - Coordinates2D center = { -1, -1 }); + const std::vector& outerRadii, Dimensions2D scanDimensions, Center center); // Create STEM Images from sparse data template @@ -73,12 +76,14 @@ namespace stempy { } } - template - std::vector createSTEMImages( + // Only typed double pairs select this overload. Bare brace-initialized centers + // cannot deduce Center and continue to use the original integer overload. + template + std::enable_if_t::value, + std::vector> createSTEMImages( const std::vector>& sparseData, const std::vector& innerRadii, const std::vector& outerRadii, - Dimensions2D scanDimensions = { 0, 0 }, - Dimensions2D frameDimensions = { 0, 0 }, Coordinates2D center = { -1, -1 }) + Dimensions2D scanDimensions, Dimensions2D frameDimensions, Center center) { if (innerRadii.empty() || outerRadii.empty()) { std::ostringstream msg; @@ -111,6 +116,34 @@ namespace stempy { // Create STEM Images from electron counted sparse data struct ElectronCountedData; + // Only typed double pairs select this overload. Bare brace-initialized centers + // cannot deduce Center and continue to use the original integer overload. + template + std::enable_if_t::value, + std::vector> createSTEMImages( + const ElectronCountedData& sparseData, const std::vector& innerRadii, + const std::vector& outerRadii, Center center); + + // Integer signatures retained for source and binary compatibility. The + // constrained overloads above require a typed double pair, leaving bare + // brace-initialized centers and default arguments on the integer API. + template + std::vector createSTEMImages( + InputIt first, InputIt last, const std::vector& innerRadii, + const std::vector& outerRadii, Dimensions2D scanDimensions = { 0, 0 }, + Coordinates2D center = { -1, -1 }); + + template + std::vector createSTEMImages( + const std::vector>& sparseData, + const std::vector& innerRadii, const std::vector& outerRadii, + Dimensions2D scanDimensions = { 0, 0 }, + Dimensions2D frameDimensions = { 0, 0 }, Coordinates2D center = { -1, -1 }) + { + return createSTEMImages(sparseData, innerRadii, outerRadii, + scanDimensions, frameDimensions, CoordinatesDouble2D(center)); + } + std::vector createSTEMImages( const ElectronCountedData& sparseData, const std::vector& innerRadii, const std::vector& outerRadii, Coordinates2D center = { -1, -1 }); diff --git a/stempy/mask.cpp b/stempy/mask.cpp index 47bf6b94..6f745af5 100644 --- a/stempy/mask.cpp +++ b/stempy/mask.cpp @@ -8,6 +8,17 @@ namespace stempy { uint16_t* createAnnularMask(Dimensions2D dimensions, int innerRadius, int outerRadius, Coordinates2D center) +{ + return createAnnularMask(dimensions, innerRadius, outerRadius, + CoordinatesDouble2D(center)); +} + +// Only typed double pairs select this overload. Bare brace-initialized centers +// cannot deduce Center and continue to use the original integer overload. +template +std::enable_if_t::value, uint16_t*> +createAnnularMask(Dimensions2D dimensions, int innerRadius, + int outerRadius, Center center) { auto numberOfElements = dimensions.first * dimensions.second; auto mask = new uint16_t[numberOfElements](); @@ -33,4 +44,6 @@ uint16_t* createAnnularMask(Dimensions2D dimensions, int innerRadius, return mask; } +template uint16_t* createAnnularMask(Dimensions2D, int, int, CoordinatesDouble2D); + } diff --git a/stempy/mask.h b/stempy/mask.h index 7a095b3b..d174848c 100644 --- a/stempy/mask.h +++ b/stempy/mask.h @@ -7,11 +7,20 @@ #include #include #include +#include namespace stempy { +// Preserve the integer API, including brace-initialized centers. uint16_t* createAnnularMask(Dimensions2D dimensions, int innerRadius, int outerRadius, Coordinates2D center = { -1, -1 }); + +// Only typed double pairs select this overload. Bare brace-initialized centers +// cannot deduce Center and continue to use the original integer overload. +template +std::enable_if_t::value, uint16_t*> +createAnnularMask(Dimensions2D dimensions, int innerRadius, + int outerRadius, Center center); } #endif diff --git a/stempy/reader.h b/stempy/reader.h index 4292a915..faa64f7b 100644 --- a/stempy/reader.h +++ b/stempy/reader.h @@ -31,6 +31,8 @@ namespace stempy { // Convention is (x, y) using Coordinates2D = std::pair; +using CoordinatesDouble2D = std::pair; + // Convention is (width, height) using Dimensions2D = std::pair; diff --git a/tests/test_image.py b/tests/test_image.py index 2b8925dd..efe7d3e9 100644 --- a/tests/test_image.py +++ b/tests/test_image.py @@ -3,10 +3,67 @@ import numpy as np -from stempy.image import com_dense, com_sparse, radial_sum_sparse +from stempy.image import com_dense, com_sparse, radial_sum_sparse, create_stem_images from stempy.io.sparse_array import SparseArray +@pytest.mark.parametrize("center", [ + (2, 1), (2.0, 1.0), np.array([2, 1]), + (2.75, 1.25), np.array([[2.75], [1.25]]), (-1, -1), +]) +@pytest.mark.parametrize("sparse", [False, True]) +def test_create_stem_images_center(center, sparse): + # Unequal coordinates detect x/y swaps, fractional centers must not truncate. + frames = np.ones((32, 5, 5), dtype=np.uint16) + if sparse: + data = np.empty((32, 1), dtype=object) + for i in range(32): + data[i, 0] = np.arange(25, dtype=np.uint32) + source = SparseArray(data=data, scan_shape=(4, 8), frame_shape=(5, 5)) + else: + source = frames + + images = create_stem_images(source, [0, 1], [1, 2], + scan_dimensions=(8, 4), center=center) + x, y = np.asarray(center).reshape(2) + if x < 0: + x = 3 # Preserve the existing rounded default on odd-sized frames. + if y < 0: + y = 3 + yy, xx = np.indices((5, 5)) + distances = (xx - x) ** 2 + (yy - y) ** 2 + expected = [np.count_nonzero((distances >= lo ** 2) & (distances < hi ** 2)) + for lo, hi in [(0, 1), (1, 2)]] + + # Sparse integration sums the existing 0xFFFF mask values per electron. + if sparse: + expected = np.array(expected) * 0xFFFF + + assert images.shape == (2, 4, 8) + np.testing.assert_array_equal(images, np.broadcast_to( + np.array(expected)[:, None, None], images.shape)) + + +def test_create_stem_images_com_dense_center(): + frame = np.zeros((5, 5), dtype=np.uint16) + frame[1, 2] = 1 + frame[2, 3] = 3 + center = com_dense(frame) + np.testing.assert_array_equal(center, [[2.75], [1.75]]) + images = create_stem_images(np.repeat(frame[None], 32, axis=0), 0, 1, + scan_dimensions=(8, 4), center=center) + np.testing.assert_array_equal(images, np.full((1, 4, 8), 3)) + np.testing.assert_array_equal(center, [[2.75], [1.75]]) + + +@pytest.mark.parametrize("center", [(1,), (1, 2, 3), [[1, 2]], + (np.nan, 1), (1, np.inf)]) +def test_create_stem_images_invalid_center(center): + with pytest.raises(ValueError, match='center'): + create_stem_images(np.ones((32, 5, 5), dtype=np.uint16), 0, 1, + scan_dimensions=(8, 4), center=center) + + @pytest.mark.parametrize("version", [0, 1]) def test_com_sparse(sparse_array_small, full_array_small, version): # Do a basic test of com_sparse() with multiple frames per scan @@ -50,17 +107,17 @@ def test_radial_sum_sparse(sparse_array_10x10): @pytest.mark.parametrize("version", [0, 1]) def test_com_sparse_parameters(simulate_sparse_array, version): - + sp = simulate_sparse_array #((100,100), (100,100), (30,70), (0.8), (10)) - + # Test no inputs. This should be the full frame COM com0 = com_sparse(sp, version=version) assert round(com0[0,].mean()) == 30 - + # Test crop_to input. Initial COM should be full frame COM com1 = com_sparse(sp, crop_to=10, version=version) assert round(com1[0,].mean()) == 30 - + # Test crop_to input as tuple. Initial COM should be full frame COM com1 = com_sparse(sp, crop_to=(10, 5), version=version) assert round(com1[0,].mean()) == 30