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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 28 additions & 9 deletions python/image.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,13 @@ ElectronCountedDataPyArray electronCount(Reader* reader,
return electronCount(reader, options.toCpp());
}

// Explicitly instantiate version for py::array_t
template std::vector<STEMImage> createSTEMImages(
const std::vector<std::vector<py::array_t<uint32_t>>>& sparseData,
const std::vector<int>& innerRadii, const std::vector<int>& outerRadii,
Dimensions2D scanDimensions, Dimensions2D frameDimensions,
CoordinatesDouble2D center);

// Retain the original exported integer specialization.
template std::vector<STEMImage> createSTEMImages(
const std::vector<std::vector<py::array_t<uint32_t>>>& sparseData,
const std::vector<int>& innerRadii, const std::vector<int>& outerRadii,
Expand All @@ -223,13 +229,26 @@ template std::vector<STEMImage> createSTEMImages(

} // namespace stempy

vector<STEMImage> 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 <typename Center>
std::enable_if_t<std::is_same<Center, CoordinatesDouble2D>::value,
vector<STEMImage>> createSTEMImages(const ElectronCountedDataPyArray& array,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Coordinates2D coords)
Center center)
{
return createSTEMImages(array.data, innerRadii, outerRadii,
array.scanDimensions, array.frameDimensions, coords);
array.scanDimensions, array.frameDimensions, center);
}

vector<STEMImage> createSTEMImages(const ElectronCountedDataPyArray& array,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Coordinates2D center)
{
return createSTEMImages(array, innerRadii, outerRadii,
CoordinatesDouble2D(center));
}

template <typename... Params>
Expand Down Expand Up @@ -392,27 +411,27 @@ PYBIND11_MODULE(_image, m)
m.def("create_stem_images",
(vector<STEMImage>(*)(StreamReader::iterator, StreamReader::iterator,
const vector<int>&, const vector<int>&,
Dimensions2D, Coordinates2D)) &
Dimensions2D, CoordinatesDouble2D)) &
createSTEMImages<StreamReader::iterator>,
py::call_guard<py::gil_scoped_release>());
m.def(
"create_stem_images",
(vector<STEMImage>(*)(SectorStreamReader::iterator,
SectorStreamReader::iterator, const vector<int>&,
const vector<int>&, Dimensions2D, Coordinates2D)) &
const vector<int>&, Dimensions2D, CoordinatesDouble2D)) &
createSTEMImages<SectorStreamReader::iterator>,
py::call_guard<py::gil_scoped_release>());
m.def("create_stem_images",
(vector<STEMImage>(*)(
const std::vector<std::vector<py::array_t<uint32_t>>>&,
const vector<int>&, const vector<int>&, Dimensions2D, Dimensions2D,
Coordinates2D)) &
CoordinatesDouble2D)) &
createSTEMImages,
py::call_guard<py::gil_scoped_release>());
m.def(
"create_stem_images",
(vector<STEMImage>(*)(const ElectronCountedDataPyArray&, const vector<int>&,
const vector<int>&, Coordinates2D)) &
const vector<int>&, CoordinatesDouble2D)) &
createSTEMImages,
py::call_guard<py::gil_scoped_release>());
m.def("calculate_average", &calculateAverage<StreamReader::iterator>,
Expand Down Expand Up @@ -585,7 +604,7 @@ PYBIND11_MODULE(_image, m)
m.def("create_stem_images",
(vector<STEMImage>(*)(PyReader::iterator, PyReader::iterator,
const vector<int>&, const vector<int>&,
Dimensions2D, Coordinates2D)) &
Dimensions2D, CoordinatesDouble2D)) &
createSTEMImages<PyReader::iterator>,
py::call_guard<py::gil_scoped_release>());
m.def("maximum_diffraction_pattern",
Expand Down
70 changes: 40 additions & 30 deletions python/stempy/image/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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]
Expand Down Expand Up @@ -532,15 +542,15 @@ 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
comy0 = np.sum(y) / mm0
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)
Expand All @@ -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(
Expand Down Expand Up @@ -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,)
Expand All @@ -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
Expand All @@ -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]
Expand Down
60 changes: 55 additions & 5 deletions stempy/image.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -159,12 +159,15 @@ void _runCalculateSTEMValues(const uint16_t data[],
}
} // end namespace

template <typename InputIt>
vector<STEMImage> 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 <typename InputIt, typename Center>
std::enable_if_t<std::is_same<Center, CoordinatesDouble2D>::value,
vector<STEMImage>> createSTEMImages(InputIt first, InputIt last,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Dimensions2D scanDimensions,
Coordinates2D center)
Center center)
{
if (first == last) {
ostringstream msg;
Expand Down Expand Up @@ -332,10 +335,14 @@ std::vector<int> createSTEMHistogram(const STEMImage& inImage,
return frequencies;
}

vector<STEMImage> 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 <typename Center>
std::enable_if_t<std::is_same<Center, CoordinatesDouble2D>::value,
vector<STEMImage>> createSTEMImages(const ElectronCountedData& data,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Coordinates2D center)
Center center)
{
return createSTEMImages(data.data, innerRadii, outerRadii,
data.scanDimensions, data.frameDimensions, center);
Expand Down Expand Up @@ -625,7 +632,50 @@ Image<double> maximumDiffractionPattern(InputIt first, InputIt last)
return maximumDiffractionPattern(first, last, dark);
}

template <typename InputIt>
vector<STEMImage> createSTEMImages(InputIt first, InputIt last,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Dimensions2D scanDimensions,
Coordinates2D center)
{
return createSTEMImages(first, last, innerRadii, outerRadii,
scanDimensions, CoordinatesDouble2D(center));
}

vector<STEMImage> createSTEMImages(const ElectronCountedData& data,
const vector<int>& innerRadii,
const vector<int>& outerRadii,
Coordinates2D center)
{
return createSTEMImages(data, innerRadii, outerRadii,
CoordinatesDouble2D(center));
}

template vector<STEMImage> createSTEMImages(const ElectronCountedData&,
const vector<int>&, const vector<int>&, CoordinatesDouble2D);

// Instantiate the ones that can be used
template vector<STEMImage> createSTEMImages<StreamReader::iterator>(
StreamReader::iterator first, StreamReader::iterator last,
const vector<int>& innerRadii, const vector<int>& outerRadii,
Dimensions2D scanDimensions, CoordinatesDouble2D center);

template vector<STEMImage> createSTEMImages<PyReader::iterator>(
PyReader::iterator first, PyReader::iterator last,
const vector<int>& innerRadii, const vector<int>& outerRadii,
Dimensions2D scanDimensions, CoordinatesDouble2D center);

template vector<STEMImage> createSTEMImages<vector<Block>::iterator>(
vector<Block>::iterator first, vector<Block>::iterator last,
const vector<int>& innerRadii, const vector<int>& outerRadii,
Dimensions2D scanDimensions, CoordinatesDouble2D center);

template vector<STEMImage> createSTEMImages<SectorStreamReader::iterator>(
SectorStreamReader::iterator first, SectorStreamReader::iterator last,
const vector<int>& innerRadii, const vector<int>& outerRadii,
Dimensions2D scanDimensions, CoordinatesDouble2D center);

template vector<STEMImage> createSTEMImages<StreamReader::iterator>(
StreamReader::iterator first, StreamReader::iterator last,
const vector<int>& innerRadii, const vector<int>& outerRadii,
Expand Down
Loading
Loading