Skip to content

Image and volume manipulation¤

cryojax.ndimage implements routines for image and volume arrays, such coordinate creation, downsampling, filters, and masks. This is a key submodule for supporting cryojax.simulator.

Coordinate systems¤

This documentation is a collection of functions used to work with coordinate systems in cryojax's conventions. The most important functions are make_coordinate_grid and make_frequency_grid.

Creating coordinate systems¤

cryojax.ndimage.make_coordinate_grid(shape: tuple[int, ...], grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0) -> Float[Array, '*shape ndim'] ¤

Create a real-space cartesian coordinate system on a grid.

Arguments:

  • shape: Shape of the grid, with ndim = len(shape).
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.

Returns:

A cartesian coordinate system in real space.


cryojax.ndimage.make_frequency_grid(shape: tuple[int, ...], grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0, outputs_rfftfreqs: bool = True, fftshifted: bool = False) -> Float[Array, '*shape ndim'] ¤

Create a fourier-space cartesian coordinate system on a grid. The zero-frequency component is in the corner.

Arguments:

  • shape: Shape of the grid, with ndim = len(shape).
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.
  • outputs_rfftfreqs: Return a frequency grid for use with jax.numpy.fft.rfftn. shape[-1] is the axis on which the negative frequencies are omitted.

Returns:

A cartesian coordinate system in frequency space.


cryojax.ndimage.make_radial_coordinate_grid(shape: tuple[int, ...], grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0) -> Float[Array, '*shape'] ¤

Create a real-space radial coordinate system on a grid.

This wraps the function make_coordinate_grid to compute the coordinate vector magnitude.

Arguments:

  • shape: Shape of the grid, with ndim = len(shape).
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.

Returns:

A radial coordinate system in real space.


cryojax.ndimage.make_radial_frequency_grid(shape: tuple[int, ...], grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0, outputs_rfftfreqs: bool = True, fftshifted: bool = False) -> Float[Array, '*shape'] ¤

Create a fourier-space radial coordinate system on a grid. The zero-frequency component is in the corner.

This wraps the function make_frequency_grid to compute the frequency magnitude, which is a common use case for things like computing fourier shell correlations and power spectrums.

Arguments:

  • shape: Shape of the grid, with ndim = len(shape).
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.
  • outputs_rfftfreqs: Return a frequency grid for use with jax.numpy.fft.rfftn. shape[-1] is the axis on which the negative frequencies are omitted.

Returns:

A radial coordinate system in frequency space.


cryojax.ndimage.make_frequency_slice(shape: tuple[int, int], grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0, outputs_rfftfreqs: bool = True, fftshifted: bool = True) -> Float[Array, '1 {shape[0]} {shape[1]} 3'] ¤

Create central slice frequency coordinates. By default, returns in the convention required for usage with cryojax.ndimage.sample_fft_slice.

Arguments:

  • shape: Shape of the frequency slice, e.g. shape = (100, 100).
  • grid_spacing: The grid spacing (i.e. voxel size), in units of length.
  • outputs_rfftfreqs: Return a frequency grid for use with jax.numpy.fft.rfftn. shape[-1] is the axis on which the negative frequencies are omitted.

Returns:

The central, \(q_z = 0\) slice of a 3D frequency grid \((q_x, q_y, q_z)\), where zero-frequency component is in the center of the grid.


cryojax.ndimage.make_1d_coordinate_grid(size: int, grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0, fftshifted: bool = False) -> Float[Array, '*shape ndim'] ¤

Create a 1D real-space cartesian coordinate array.

Arguments:

  • size: Size of the coordinate array.
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.

Returns:

A 1D cartesian coordinate array in real space.


cryojax.ndimage.make_1d_frequency_grid(size: int, grid_spacing: float | Float[ndarray, ''] | Float[Array, ''] = 1.0, outputs_rfftfreqs: bool = True, fftshifted: bool = False) -> Float[Array, '*shape ndim'] ¤

Create a 1D fourier-space cartesian coordinate array. If outputs_rfftfreqs = False, the zero-frequency component is in the beginning.

Arguments¤
  • size: Size of the coordinate array.
  • grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.
  • outputs_rfftfreqs: Return a frequency grid for use with jax.numpy.fft.rfftn. shape[-1] is the axis on which the negative frequencies are omitted.

Returns:

A 1D cartesian coordinate array in frequency space.

Transforming coordinate systems¤

cryojax also provides functions that transform between coordinate conventions.

cryojax.ndimage.cartesian_to_polar(coordinate_or_frequency_grid: Float[Array, 'y_dim x_dim 2'], square: bool = False) -> tuple[Inexact[Array, 'y_dim x_dim'], Inexact[Array, 'y_dim x_dim']] ¤

Convert from cartesian to polar coordinates.

Arguments:

  • coordinate_or_frequency_grid: The cartesian coordinate system in real or fourier space.
  • square: If True, return the square of the radial coordinate \(|r|^2\). Otherwise, return \(|r|\).

Returns:

A tuple (r, theta), where r is the radial coordinate system and theta is the angular coordinate system. If square=True, return a tuple (r_squared, theta).

Image transforms (e.g. filters and masks)¤

cryojax.ndimage.AbstractImageTransform

cryojax.ndimage.AbstractImageTransform ¤

Base class for computing and applying an Array to an image.

__init__() ¤

Initialize self. See help(type(self)) for accurate signature.

__call__(image: Inexact[Array, '*batch _ _'] | Inexact[Array, '*batch _ _ _']) -> Inexact[Array, '*batch _ _'] | Inexact[Array, '*batch _ _ _'] ¤

Filters¤

cryojax.ndimage.AbstractFilter

cryojax.ndimage.AbstractFilter(cryojax.ndimage.AbstractImageTransform) ¤

Base class for computing and applying an image filter.

get() -> Float[Array, 'y_dim x_dim'] | Float[Array, 'z_dim y_dim x_dim'] ¤

cryojax.ndimage.LowpassFilter(cryojax.ndimage.AbstractFilter) ¤

Apply a low-pass filter to an image or volume, with a cosine soft-edge.

__init__(frequency_grid: Float[Array, 'y_dim x_dim 2'] | Float[Array, 'z_dim y_dim x_dim 3'], frequency_cutoff_fraction: cryojax.jax_util.FloatLike = 0.95, rolloff_width_fraction: cryojax.jax_util.FloatLike = 0.05) ¤

Arguments:

  • frequency_grid: The frequency grid of the image or volume, in pixel-units.
  • frequency_cutoff_fraction: The cutoff frequency as a fraction of the Nyquist frequency. By default, 0.95.
  • rolloff_width_fraction: The rolloff width as a fraction of the Nyquist frequency. By default, 0.05.
get() -> Inexact[Array, 'y_dim x_dim'] | Inexact[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim']) -> Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the filter to an image or volume, which may carry leading batch dimensions. The filter is broadcast against them.


cryojax.ndimage.HighpassFilter(cryojax.ndimage.AbstractFilter) ¤

Apply a high-pass filter to an image or volume, with a cosine soft-edge.

__init__(frequency_grid: Float[Array, 'y_dim x_dim 2'] | Float[Array, 'z_dim y_dim x_dim 3'], frequency_cutoff_fraction: cryojax.jax_util.FloatLike = 0.95, rolloff_width_fraction: cryojax.jax_util.FloatLike = 0.05) ¤

Arguments:

  • frequency_grid: The frequency grid of the image or volume, in pixel-units.
  • frequency_cutoff_fraction: The cutoff frequency as a fraction of the Nyquist frequency. By default, 0.95.
  • rolloff_width_fraction: The rolloff width as a fraction of the Nyquist frequency. By default, 0.05.
get() -> Inexact[Array, 'y_dim x_dim'] | Inexact[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim']) -> Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the filter to an image or volume, which may carry leading batch dimensions. The filter is broadcast against them.


cryojax.ndimage.WhiteningFilter(cryojax.ndimage.AbstractFilter) ¤

Compute a whitening filter from an image. This is taken to be the inverse square root of the 2D radially averaged power spectrum.

The filter is normalized to preserve the mean and variance of the image it is applied to: the zero-frequency (mean) mode is left unchanged and the remaining modes are rescaled so that a white-noise input maps to the identity filter.

__init__(images: Float[NDArrayLike, '_ _'] | Float[NDArrayLike, '_ _ _'], shape: tuple[int, int] | None = None, *, interp: str = 'linear', squared: bool = False) ¤

Arguments:

  • images: The image (or stack of images) from which to compute the power spectrum.
  • shape: The shape of the resulting filter. This downsamples or upsamples the filter by cropping or padding in real space.
  • interp: The method of interpolating the binned, radially averaged power spectrum onto a 2D grid. Either nearest or linear.
  • squared: If False, the whitening filter is the inverse square root of the image power. If True, the filter is the inverse of the image power.
get() -> Inexact[Array, 'y_dim x_dim'] ¤
__call__(image: Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim']) -> Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the filter to an image or volume, which may carry leading batch dimensions. The filter is broadcast against them.


cryojax.ndimage.CustomFilter(cryojax.ndimage.AbstractFilter) ¤

Pass a custom filter as an array.

__init__(filter: Inexact[NDArrayLike, 'y_dim x_dim'] | Inexact[NDArrayLike, 'z_dim y_dim x_dim']) ¤
get() -> Inexact[Array, 'y_dim x_dim'] | Inexact[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim']) -> Complex[Array, '*batch y_dim x_dim'] | Complex[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the filter to an image or volume, which may carry leading batch dimensions. The filter is broadcast against them.

Other Fourier space operations:

cryojax.ndimage.PhaseShiftFFT(cryojax.ndimage.AbstractImageTransform) ¤

Apply a phase shift to an image in Fourier space, effectively applying an in-plane shift to the image in real space. Only square images are supported.

Apply a translation in real-space

import jax.numpy as jnp
from cryojax.ndimage import PhaseShiftFFT

offset_in_angstroms = jnp.array([50.0, -30.0])
fft = jnp.fft.rfftn(...) # e.g., fft of a real 2D image

shift_fn = PhaseShiftFFT(
    offset=offset_in_angstroms, pixel_size=1.1
)

shifted_fft = shift_fn(fft)
shifted_image = jnp.fft.irfftn(shifted_image_fft)
__init__(offset: Float[NDArrayLike, '2'] | Float[NDArrayLike, '3'], *, pixel_size: cryojax.jax_util.FloatLike = 1.0) ¤

Arguments:

  • offset: The offset by which to shift the image, in pixels or angstroms.
  • pixel_size: The pixel size of the image. Set pixel_size if offset is given in Angstroms, and leave as 1.0 if offset is given in pixel units.
__call__(image: Complex[Array, 'y_dim x_dim'] | Complex[Array, 'z_dim y_dim x_dim']) -> Complex[Array, 'y_dim x_dim'] | Complex[Array, 'z_dim y_dim x_dim'] ¤

Apply the phase shift to the input image in Fourier space.

Arguments:

  • image: The input image in Fourier space.

Returns:

The phase shifted image in Fourier space.

cryojax.ndimage.RotateFFT(cryojax.ndimage.AbstractImageTransform) ¤

Rotate an image in Fourier space using interpolation. Only square, even-dimension images are supported.

Rotation is done by interpolating the image's fourier transform with cryojax.ndimage.map_frequencies.

Example

import jax.numpy as jnp
from cryojax.ndimage import RotateFFT, make_frequency_grid

image = ...  # e.g., a real 2D image of shape (dim, dim)
frequency_grid = make_frequency_grid((dim, dim))  # in pixels

rotation_fn = RotateFFT(
    rotation_angle=45.0, frequency_grid=frequency_grid
)

rotated_image_fft = jnp.fft.irfftn(
    rotation_fn(jnp.fft.rfftn(image)), s=image.shape
)
__init__(rotation_angle: cryojax.jax_util.FloatLike, *, frequency_grid: Float[NDArrayLike, 'y_dim x_dim 2'] | None = None, pixel_size: cryojax.jax_util.FloatLike = 1.0) ¤

Arguments:

  • rotation_angle: The angle by which to rotate the image, in degrees.
  • frequency_grid: The frequency grid, of the half-space (rfft) shape (dim, dim // 2 + 1, 2), as returned by cryojax.ndimage.make_frequency_grid. If not provided, generate on-the-fly.
  • pixel_size: The pixel size of the frequency_grid.
__call__(image: Complex[Array, 'y_dim x_dim']) -> Complex[Array, 'y_dim x_dim'] ¤

Rotate the input image in Fourier space.

Arguments:

image: The image in Fourier space, i.e. the output of jax.numpy.fft.rfftn.

Returns:

The rotated image in Fourier space.

Masks¤

cryojax.ndimage.AbstractMask

cryojax.ndimage.AbstractMask(cryojax.ndimage.AbstractImageTransform) ¤

Base class for computing and applying an image mask.

get() -> Float[Array, 'y_dim x_dim'] | Float[Array, 'z_dim y_dim x_dim'] ¤

cryojax.ndimage.CircularCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a circular mask to an image with a cosine soft-edge.

__init__(coordinate_grid: Float[Array, 'y_dim x_dim 2'], radius: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike, xy_offset: tuple[float, float] | Float[NDArrayLike, '2'] = (0.0, 0.0)) ¤

Arguments:

  • coordinate_grid: The image coordinates.
  • radius: The radius of the circular mask.
  • rolloff_width: The rolloff width of the soft edge.
get() -> Float[Array, 'y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.SphericalCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a spherical mask to a volume with a cosine soft-edge.

__init__(coordinate_grid: Float[Array, 'z_dim y_dim x_dim 3'], radius: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike) ¤

Arguments:

  • coordinate_grid: The volume coordinates.
  • radius: The radius of the spherical mask.
  • rolloff_width: The rolloff width of the soft edge.
get() -> Float[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.SquareCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a square mask to an image with a cosine soft-edge.

__init__(coordinate_grid: Float[Array, 'y_dim x_dim 2'], side_length: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike) ¤

Arguments:

  • coordinate_grid: The image coordinates.
  • side_length: The side length of the square.
  • rolloff_width: The rolloff width of the soft edge.
get() -> Float[Array, 'y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.Rectangular2DCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a rectangular mask in 2D to an image with a cosine soft-edge. Optionally, rotate the rectangle by an angle.

__init__(coordinate_grid: Float[Array, 'y_dim x_dim 2'], x_width: cryojax.jax_util.FloatLike, y_width: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike, rotation_angle: cryojax.jax_util.FloatLike = 0.0) ¤

Arguments:

  • coordinate_grid: The image coordinates.
  • x_width: The width of the rectangle along the x-axis.
  • y_width: The width of the rectangle along the y-axis.
  • rolloff_width: The rolloff width of the soft edge.
  • rotation_angle: The in-plane rotation angle of the rectangle in degrees. By default, 0.0.
get() -> Float[Array, 'y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.Rectangular3DCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a rectangular mask to a volume with a cosine soft-edge.

__init__(coordinate_grid: Float[Array, 'z_dim y_dim x_dim 3'], x_width: cryojax.jax_util.FloatLike, y_width: cryojax.jax_util.FloatLike, z_width: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike) ¤

Arguments:

  • coordinate_grid: The volume coordinates.
  • x_width: The width of the rectangle along the x-axis.
  • y_width: The width of the rectangle along the y-axis.
  • z_width: The width of the rectangle along the z-axis.
  • rolloff_width: The rolloff width of the soft edge.
get() -> Float[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.Cylindrical2DCosineMask(cryojax.ndimage.AbstractMask) ¤

Apply a cylindrical mask to an image with a cosine soft-edge. This implements an infinite in-plane cylinder, rotated at a given angle.

__init__(coordinate_grid: Float[Array, 'y_dim x_dim 2'], radius: cryojax.jax_util.FloatLike, rolloff_width: cryojax.jax_util.FloatLike, rotation_angle: cryojax.jax_util.FloatLike = 0.0, length: cryojax.jax_util.FloatLike | None = None) ¤

Arguments:

  • coordinate_grid: The image coordinates.
  • radius: The radius of the cylinder.
  • rolloff_width: The rolloff width of the soft edge.
  • rotation_angle: The in-plane rotation angle of the cylinder in degrees. By default, 0.0.
  • length: The length of the cylinder. If None, do not mask the cylinder length-wise.
get() -> Float[Array, 'y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.


cryojax.ndimage.CustomMask(cryojax.ndimage.AbstractMask) ¤

Pass a custom mask as an array.

__init__(mask_array: Float[Array, 'y_dim x_dim'] | Float[Array, 'z_dim y_dim x_dim']) ¤
get() -> Float[Array, 'y_dim x_dim'] | Float[Array, 'z_dim y_dim x_dim'] ¤
__call__(image: Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim']) -> Float[Array, '*batch y_dim x_dim'] | Float[Array, '*batch z_dim y_dim x_dim'] ¤

Apply the mask to an image or volume, which may carry leading batch dimensions. The mask is broadcast against them.

Other real-space operations

cryojax.ndimage.ScaleImage(cryojax.ndimage.AbstractImageTransform) ¤

ScaleImage(scale: cryojax.jax_util._typing.FloatLike = 1.0, offset: cryojax.jax_util._typing.FloatLike = 0.0)

__init__(scale: cryojax.jax_util.FloatLike = 1.0, offset: cryojax.jax_util.FloatLike = 0.0) ¤
__call__(image: Inexact[Array, '*batch y_dim x_dim'] | Inexact[Array, '*batch z_dim y_dim x_dim']) -> Inexact[Array, '*batch y_dim x_dim'] | Inexact[Array, '*batch z_dim y_dim x_dim'] ¤

Operators¤

Fourier-space¤

cryojax.ndimage.AbstractFourierOperator

cryojax.ndimage.AbstractFourierOperator ¤

The base class for all fourier-based operators.

By convention, operators should be defined to be dimensionless (up to a scale factor).

To create a subclass,

1) Include the necessary parameters in
   the class definition.
2) Overrwrite the `__call__` method.
__call__(frequencies: Float[Array, '...']) -> Inexact[Array, '...'] ¤

cryojax.ndimage.FourierGaussian(cryojax.ndimage.AbstractFourierOperator) ¤

This operator represents a simple gaussian. Specifically, this is

.. math:: P(k) = \kappa \exp(- \beta k^2 / 4),

where :math:k^2 = k_x^2 + k_y^2 is the length of the wave vector. Here, :math:\beta has dimensions of length squared.

__init__(amplitude: cryojax.jax_util.FloatLike = 1.0, b_factor: cryojax.jax_util.FloatLike = 1.0) ¤

Arguments:

  • amplitude: The amplitude of the operator, equal to \(\kappa\) in the above equation.
  • b_factor: The B-factor of the gaussian, equal to \(\beta\) in the above equation.
__call__(frequencies: Float[Array, '...']) -> Float[Array, '...'] ¤

cryojax.ndimage.PeakedFourierGaussian(cryojax.ndimage.AbstractFourierOperator) ¤

This operator represents a gaussian with a peak at a given frequency shell.

__init__(amplitude: cryojax.jax_util.FloatLike = 1.0, b_factor: cryojax.jax_util.FloatLike = 1.0, radial_peak: cryojax.jax_util.FloatLike = 0.0) ¤

Arguments:

  • amplitude: The amplitude of the operator, equal to \(\kappa\) in the above equation.
  • b_factor: The B-factor of the gaussian, equal to \(\beta\) in the above equation.
  • radial_peak: The frequency shell of the gaussian peak.
__call__(frequencies: Float[Array, '...']) -> Float[Array, '...'] ¤

cryojax.ndimage.FourierConstant(cryojax.ndimage.AbstractFourierOperator) ¤

An operator that is a constant.

__init__(value: cryojax.jax_util.FloatLike) ¤

Arguments:

  • value: The value of the constant
__call__(frequencies: Float[Array, '...']) -> Float[Array, '...'] ¤

cryojax.ndimage.FourierSinc(cryojax.ndimage.AbstractFourierOperator) ¤

The separable sinc function is the Fourier transform of the box function and is commonly used for anti-aliasing applications. In 2D, this is

\[f_{2D}(\vec{q}) = \sinc(q_x w) \sinc(q_y w),\]

and in 3D this is

\[f_{3D}(\vec{q}) = \sinc(q_x w) \sinc(q_y w) \sinc(q_z w)},\]

where \(\sinc(x) = \frac{\sin(\pi x)}{\pi x}\), \(\vec{q} = (q_x, q_y)\) or \(\vec{q} = (q_x, q_y, q_z)\) are spatial frequency coordinates for 2D and 3D respectively, and \(w\) is width of the real-space box function.

__init__(box_width: cryojax.jax_util.FloatLike = 1.0) ¤

Arguments:

  • box_width: If the inverse fourier transform of this class is the rectangular function, its interval is - box_width / 2 to + box_width / 2.
__call__(frequencies: Float[Array, '...']) -> Float[Array, '...'] ¤

cryojax.ndimage.FourierPhaseShifts(cryojax.ndimage.AbstractFourierOperator) ¤

Apply a phase shift the Fourier domain.

__init__(shift: cryojax.jax_util.FloatLike | Float[NDArrayLike, '2'] | Float[NDArrayLike, '3']) ¤

Arguments:

  • shift: The shift to apply in the Fourier domain. The units of this should be the inverse of the units of the frequencies passed at runtime.
__call__(frequencies: Float[Array, '...']) -> Complex[Array, '...'] ¤

cryojax.ndimage.CustomFourierOperator(cryojax.ndimage.AbstractFourierOperator) ¤

An operator that calls a custom function.

__init__(fn: Callable[..., Inexact[Array, '...']], *args: Any, **kwargs: Any) ¤

Arguments:

  • fn: The Callable wrapped into a AbstractFourierOperator. Has signature out = fn(frequencies, *args, **kwargs)
  • args: Passed to fn.
  • kwargs: Passed to fn.
__call__(frequencies: Float[Array, '...']) -> Inexact[Array, '...'] ¤

Real-space¤

cryojax.ndimage.AbstractRealOperator

cryojax.ndimage.AbstractRealOperator ¤

The base class for all real operators.

By convention, operators should be defined to have units of inverse area (up to a scale factor).

To create a subclass,

1) Include the necessary parameters in
   the class definition.
2) Overrwrite the `__call__` method.
__call__(coordinates: Float[Array, '...']) -> Inexact[Array, '...'] ¤

cryojax.ndimage.RealGaussian(cryojax.ndimage.AbstractRealOperator) ¤

This operator is a normalized gaussian in real space

\[g(r) = \frac{\kappa}{2\pi \beta} \exp(- (r - r_0)^2 / (2 \sigma))\]

where \(r^2 = x^2 + y^2\).

__init__(amplitude: cryojax.jax_util.FloatLike = 1.0, variance: cryojax.jax_util.FloatLike = 1.0, offset: cryojax.jax_util.FloatLike | Float[NDArrayLike, '... _'] | Sequence[float] | None = None) ¤

Arguments:

  • amplitude: The amplitude of the operator, equal to \(\kappa\) in the above equation.
  • variance: The variance of the gaussian, equal to \(\sigma\) in the above equation.
  • offset: An offset to the origin, equal to \(r_0\) in the above equation.
__call__(coordinates: Float[Array, '...']) -> Float[Array, '...'] ¤

cryojax.ndimage.RealConstant(cryojax.ndimage.AbstractRealOperator) ¤

An operator that is a constant.

__init__(value: float | Float[NDArrayLike, '...']) ¤

Arguments:

  • value: The value of the constant
__call__(coordinates: Float[Array, '...']) -> Float[Array, ''] ¤

Downsampling¤

cryojax.ndimage.block_reduce_downsample(image_or_volume: Inexact[NDArrayLike, '_ _'] | Inexact[NDArrayLike, '_ _ _'], downsample_factor: int, operation: Callable[[Array, Array], Array] = <function add>, center_correct: bool = True) -> Inexact[Array, '_ _'] | Inexact[Array, '_ _ _'] ¤

Downsample an array by pooling together blocks. Wraps equinox.nn.Pool.

Arguments:

  • image_or_volume: image or volume array to downsample. The shape must be a multiple of downsample_factor
  • downsample_factor: A scale factor at which to downsample image_or_volume by. Must be a value greater than 1.
  • operation: A function such as operation = lambda x, y: f(x, y), where x and y are JAX arrays. See [equinox.nn.Pool] (https://docs.kidger.site/equinox/api/nn/pool/#equinox.nn.Pool) for documentation.
  • center_correct: If True, apply a phase shift in the fourier domain to correct the array center after downsampling. Applies only to even downsample_factor.

Returns:

The downsampled image_or_volume at shape reduced by downsample_factor.

cryojax.ndimage.fourier_crop_downsample(image_or_volume: Inexact[NDArrayLike, '_ _'] | Inexact[NDArrayLike, '_ _ _'], downsample_factor: float | int, outputs_real_space: bool = True, preserve_mean: bool = False, *, outputs_factor: bool = False) -> Inexact[Array, '_ _'] | Inexact[Array, '_ _ _'] | tuple[Inexact[Array, '_ _'] | Inexact[Array, '_ _ _'], tuple[float, ...]] ¤

Downsample an array using fourier cropping.

To make downsample_factor exact (i.e. so that a caller can rescale a pixel/voxel size by exactly downsample_factor, rather than by the ratio implied by naively truncating shape / downsample_factor), the array is first padded (by edge replication) so that its shape is an exact multiple of the new, downsampled shape, then a sub-pixel phase shift is applied before cropping in fourier space so that the downsampled array's center stays anchored to the original (unpadded) array's own center.

This is exact whenever downsample_factor is an integer. For a non-integer downsample_factor, an exact ratio is generally not achievable with integer-pixel padding; the new shape and padding are instead chosen to make the achieved ratio as close to downsample_factor as possible (error shrinking as the output size grows), and this resolved ratio can be recovered with outputs_factor = True.

Arguments:

  • image_or_volume: The image or volume array to downsample.
  • downsample_factor: A scale factor at which to downsample image_or_volume by. Must be a value greater than 1.
  • outputs_real_space: If False, the image_or_volume is returned in fourier space with the zero-frequency component in the corner. For real signals, hermitian symmetry is assumed.
  • preserve_mean: Preserve the mean of the volume after downsampling, rather than the sum.
  • outputs_factor: If True, also return the downsample factor actually resolved on each axis, as a tuple the same length as image_or_volume.ndim. Equal to downsample_factor on every axis when downsample_factor is an integer.

Returns:

The downsampled image_or_volume at shape reduced by downsample_factor. If outputs_factor = True, a (downsampled_array, resolved_downsample_factor) tuple instead.

Interpolation¤

cryojax.ndimage.map_coordinates(input: Array, coordinates: Sequence[Array], order: int = 1, mode: str = 'fill', cval: float | complex = 0.0, unroll: bool = True) -> Array ¤

Interpolate a 2D or 3D array at arbitrary coordinates.

Similar to scipy.ndimage.map_coordinates, but always corresponds to its prefilter=False case: input is convolved with the interpolation kernel directly. To interpolate the fourier transform of a real signal, use cryojax.ndimage.map_frequencies instead.

Arguments:

  • input: The 2D or 3D array to interpolate.
  • coordinates: A sequence of length input.ndim, one coordinate array per axis, in array-axis order and in index units (so input[0, 0] sits at coordinate (0, 0)). Each must be broadcastable to the same shape.
  • order: The order of the interpolation kernel: 1 for linear, 3 for cubic B-spline. Cubic is more accurate, at the cost of reading a 4^ndim rather than a 2^ndim neighborhood per query point.
  • mode : How to extrapolate beyond the edges of input, using JAX's out-of-bounds indexing modes, e.g. "fill" or "clip".
  • cval: The value returned for out-of-bounds coordinates when mode is "fill". Ignored for other modes.
  • unroll: If True (the default), gather the interpolation taps one at a time, which keeps memory use small and predictable at large batch sizes. If False, gather them all at once, which is often substantially faster for order=3 but whose peak memory grows with the number of query points.

Returns:

The interpolated values, with the shape that the coordinate arrays broadcast to.

cryojax.ndimage.map_frequencies(input: Array, frequencies: Sequence[Array], order: int = 1, mode: str = 'fill', unroll: bool = True) -> Array ¤

Interpolate the fourier transform of a real 2D image or 3D volume, at arbitrary frequencies.

A real signal's fourier transform is only stored in the half space, since the other half is redundant. Negative q_x is therefore not stored --- but its value is still exactly known, from the transform's symmetries, and this function recovers it. That matters: roughly a fifth of the frequencies of a rotated grid are negative in q_x.

Arguments:

  • input: The fourier transform of a real 2D image or 3D volume, as prepared by cryojax.ndimage.prepare_sampling_fft. Shape (dim, dim // 2 + 1) or (dim, dim, dim // 2 + 1).
  • frequencies: A sequence of length input.ndim, one frequency array per axis in array-axis order --- (q_y, q_x) or (q_z, q_y, q_x) --- in cycles/pixel, as in cryojax.ndimage.map_coordinates. Each must be broadcastable to the same shape. The last entry is the truncated (rfft) axis, q_x, which may be negative.
  • order: The order of the interpolation kernel: 1 for linear, 3 for cubic B-spline.
  • mode: What to return for frequencies that fall outside the fourier box, e.g. at the corners of a rotated grid. Either "fill" (the default), which returns zero, or "clip", which clamps them onto the edge of the box.
  • unroll: See cryojax.ndimage.map_coordinates.

Returns:

The interpolated values, with the shape that the frequency arrays broadcast to.

Spreading¤

cryojax.ndimage.spread_gaussians_2d(x: Float[Array, 'M'], y: Float[Array, 'M'], amplitude: Float[Array, 'M'], variance: Float[Array, ''] | Float[Array, 'M'], shape: tuple[int, int], *, pixel_size: Float[Array, ''], n_spread: int = 7, use_erf: bool = True, enable_pallas: bool | Mapping[str, bool] | None = None) -> Float[Array, '{shape[0]} {shape[1]}'] ¤

Scatter point strengths onto a 2D grid with an isotropic Gaussian (or pixel-averaged Gaussian) kernel.

This scatters each point's strength onto the n_spread nearest grid points along each axis, weighted by a compactly-supported Gaussian kernel. Differentiable with a custom VJP rule w.r.t. all array arguments.

Arguments:

  • x, y: Physical-unit positions of shape (M,), where 0 corresponds to the real-space center at grid index n // 2 (for both even and odd n).
  • amplitude: The per-point scattering weight, of shape (M,) (e.g. an amplitude times an atom occupancy). Multiplies the (normalized) kernel exactly once, regardless of dimensionality.
  • variance: The variance of the isotropic Gaussian kernel. May be a scalar or a per-point array of shape (M,).
  • shape: The shape (ny, nx) of the output grid.
  • pixel_size: The pixel size of the output grid, in the same units as x, y.
  • n_spread: The width (number of grid points, per axis) of the kernel used to spread each point. Controls speed / accuracy tradeoff: larger n_spread is more accurate but slower. Must be chosen relative to variance and the pixel/voxel size — too small truncates the Gaussian and silently biases the result. See cryojax.ndimage.variance_to_nspread to pick a value for a given variance. Must not exceed the smallest dimension of shape (otherwise a single point's kernel support would wrap around the grid more than once, aliasing the result).
  • use_erf: If True (default), spread the exact average of the Gaussian over a pixel (used to sample the average value within a pixel, rather than its value at a point). If False, spread a point-sampled Gaussian instead.
  • enable_pallas: Whether to use the Pallas/Triton GPU kernel backend instead of the pure-JAX (segment_sum-based) backend, for the forward and backward pass independently. Requires a CUDA GPU; raises if requested without one.

    There is no single best choice for every case -- extensive benchmarking found:

    • The pure-JAX forward pass usually wins outright (the Pallas forward kernel needs an atomic scatter, which contends under realistic atom density). Leave "fwd" False (the default) in most cases.
    • The Pallas backward pass is a real, consistent win in both memory (~10-37x less) and speed (1.25x-8.3x faster) than the pure-JAX analytic backward, because it's a pure gather with no atomic contention. {"bwd": True} -- pure-JAX forward, Pallas backward -- is the sensible starting point if you want to opt into anything.
    • enable_pallas=True (both directions) trades some speed for a flat, M-independent memory profile that also survives jax.vmap-ing this computation over a batch of particles (pure-JAX's memory scales with batch size; Pallas's doesn't) -- worth it specifically when even {"bwd": True}'s memory isn't low enough, e.g. multi-particle refinement with many particles vmapped together.
    • The right choice is also hardware- and scale-dependent (e.g. on Hopper, plain pure-JAX can outperform {"bwd": True} at moderate M, with the crossover shifting by architecture) -- benchmark your own workload if this matters.

    True/False applies to both the forward and backward pass; a dict with "fwd"/"bwd" keys sets them independently (e.g. {"fwd": True} uses Pallas only for the forward pass). None (default) defers to the CRYOJAX_ENABLE_PALLAS environment variable (False if unset). The number of points each Pallas grid program handles is not configurable here; it defaults to a flat 128 (empirically the best overall choice across Ampere, Hopper, and Blackwell, in both 2D and 3D), overridable only via the CRYOJAX_PALLAS_BLOCK_SIZE environment variable if you've benchmarked a better value for your own (GPU, M, n_spread).

Returns:

The grid of shape (ny, nx) with gaussians scattered onto it.

cryojax.ndimage.spread_gaussians_3d(x: Float[Array, 'M'], y: Float[Array, 'M'], z: Float[Array, 'M'], amplitude: Float[Array, 'M'], variance: Float[Array, ''] | Float[Array, 'M'], shape: tuple[int, int, int], *, voxel_size: Float[Array, ''], n_spread: int = 7, use_erf: bool = True, enable_pallas: bool | Mapping[str, bool] | None = None) -> Float[Array, '{shape[0]} {shape[1]} {shape[2]}'] ¤

Scatter point strengths onto a 3D grid with an isotropic Gaussian (or voxel-averaged Gaussian) kernel.

This scatters each point's strength onto the n_spread nearest grid points along each axis, weighted by a compactly-supported Gaussian kernel. Differentiable with a custom VJP rule w.r.t. all array arguments.

Arguments:

  • x, y, z: Physical-unit positions of shape (M,), where 0 corresponds to the real-space center at grid index n // 2 (for both even and odd n).
  • amplitude: The per-point scattering weight, of shape (M,) (e.g. an amplitude times an atom occupancy). Multiplies the (normalized) kernel exactly once, regardless of dimensionality.
  • variance: The variance of the isotropic Gaussian kernel. May be a scalar or a per-point array of shape (M,).
  • shape: The shape (nz, ny, nx) of the output grid.
  • voxel_size: The voxel size of the output grid, in the same units as x, y, z.
  • n_spread: The width (number of grid points, per axis) of the kernel used to spread each point. Controls speed / accuracy tradeoff: larger n_spread is more accurate but slower. Must be chosen relative to variance and the pixel/voxel size — too small truncates the Gaussian and silently biases the result. See cryojax.ndimage.variance_to_nspread to pick a value for a given variance. Must not exceed the smallest dimension of shape (otherwise a single point's kernel support would wrap around the grid more than once, aliasing the result).
  • use_erf: If True (default), spread the exact average of the Gaussian over a voxel (used to sample the average value within a voxel, rather than its value at a point). If False, spread a point-sampled Gaussian instead.
  • enable_pallas: Whether to use the Pallas/Triton GPU kernel backend instead of the pure-JAX (segment_sum-based) backend, for the forward and backward pass independently. Requires a CUDA GPU; raises if requested without one.

    There is no single best choice for every case -- extensive benchmarking found:

    • The pure-JAX forward pass usually wins outright (the Pallas forward kernel needs an atomic scatter, which contends under realistic atom density). Leave "fwd" False (the default) in most cases.
    • The Pallas backward pass is a real, consistent win in both memory (~10-37x less) and speed (1.25x-8.3x faster) than the pure-JAX analytic backward, because it's a pure gather with no atomic contention. {"bwd": True} -- pure-JAX forward, Pallas backward -- is the sensible starting point if you want to opt into anything.
    • enable_pallas=True (both directions) trades some speed for a flat, M-independent memory profile that also survives jax.vmap-ing this computation over a batch of particles (pure-JAX's memory scales with batch size; Pallas's doesn't) -- worth it specifically when even {"bwd": True}'s memory isn't low enough, e.g. multi-particle refinement with many particles vmapped together.
    • The right choice is also hardware- and scale-dependent (e.g. on Hopper, plain pure-JAX can outperform {"bwd": True} at moderate M, with the crossover shifting by architecture) -- benchmark your own workload if this matters.

    True/False applies to both the forward and backward pass; a dict with "fwd"/"bwd" keys sets them independently (e.g. {"fwd": True} uses Pallas only for the forward pass). None (default) defers to the CRYOJAX_ENABLE_PALLAS environment variable (False if unset). The number of points each Pallas grid program handles is not configurable here; it defaults to a flat 128 (empirically the best overall choice across Ampere, Hopper, and Blackwell, in both 2D and 3D), overridable only via the CRYOJAX_PALLAS_BLOCK_SIZE environment variable if you've benchmarked a better value for your own (GPU, M, n_spread).

Returns:

The grid of shape (nz, ny, nx) with gaussians scattered onto it.

cryojax.ndimage.variance_to_nspread(variance: cryojax.jax_util.FloatLike | Float[NDArrayLike, 'M'], pixel_size: cryojax.jax_util.FloatLike, n_sigma: float = 4.0) -> int ¤

Choose an n_spread sufficient to truncate the Gaussian kernel used by cryojax.ndimage.spread_gaussians_2d/cryojax.ndimage.spread_gaussians_3d at n_sigma standard deviations.

n_spread sets array shapes, so it must be a static value rather than depending on variance through tracing; call this ahead of time with concrete values instead of guessing. Too small an n_spread silently truncates the Gaussian and biases the result.

Warning

Not JIT-compatible, and never invokes JAX — variance and pixel_size are handled with plain numpy/math, so passing numpy arrays or python floats never triggers a JAX dispatch or device transfer. Call this once outside of any jax.jit-compiled function (e.g. when choosing n_spread up front from known/expected variances), not from within a jitted call graph.

Arguments:

  • variance: The variance (or per-point array of variances — the largest is used) that will be passed to spread_gaussians_2d/spread_gaussians_3d.
  • pixel_size: The pixel/voxel size of the grid variance will be spread onto.
  • n_sigma: The number of standard deviations of the Gaussian to truncate at.

Returns:

An integer n_spread, at least 2.

Fourier projection-slice extraction¤

cryojax.ndimage.prepare_sampling_fft(real_voxel_grid: Float[NDArrayLike, 'dim dim dim'], *, interp: Literal['linear', 'cubic'] = 'linear', pad_scale: float = 1.0) -> Complex[Array, 'dim dim dim//2+1'] ¤

Transform a real-space voxel grid into the fourier-space array consumed by cryojax.ndimage.sample_fft_slice.

This is the preprocessing that cryojax.simulator.FourierVoxelGridVolume does internally: optional padding, deconvolution of the interpolation kernel, and the transform itself.

Why deconvolution?

Interpolating a fourier voxel grid does not return the volume's true fourier transform, but that of the volume blurred by the interpolation kernel. The blur is known in closed form (sinc^2 for "linear", sinc^4 for "cubic"), so it is divided out of the voxel grid here, before the transform. Slice extraction then reconstructs the true transform, rather than an approximation of it, and costs nothing extra at sampling time.

What is left is aliasing, which shrinks with pad_scale.

Example

import cryojax.ndimage as im

real_voxel_grid = ...  # shape (dim, dim, dim)

sampling_fft = im.prepare_sampling_fft(real_voxel_grid)
frequency_slice = im.make_frequency_slice(
    sampling_fft.shape[:2], fftshifted=True
)
projection_fft = im.sample_fft_slice(sampling_fft, frequency_slice)

Arguments:

  • real_voxel_grid: A cubic, even-dimension voxel grid in real space.
  • interp: The interpolation method the returned grid is prepared for. The same value must be passed to cryojax.ndimage.sample_fft_slice. Either "linear" (the default), or "cubic", which is substantially more accurate at the cost of reading a 4^3 rather than a 2^3 neighborhood per query point.
  • pad_scale: Scale factor at which to Fourier-pad real_voxel_grid before the transform. Must be a value >= 1.0.

Returns:

The prepared fourier voxel grid, of shape (dim, dim, dim // 2 + 1). dim is the (possibly padded) dimension, i.e. real_voxel_grid.shape[0] when pad_scale == 1.0.

cryojax.ndimage.sample_fft_slice(sampling_fft: Complex[NDArrayLike, 'dim dim dim//2+1'], frequency_slice: Float[NDArrayLike, '1 dim dim//2+1 3'] | Float[Array, '1 dim dim 3'], *, interp: Literal['linear', 'cubic'] = 'linear', boundary: str = 'fill', unroll: bool = True) -> Complex[Array, 'dim _'] ¤

Extract a surface from a fourier-space voxel grid using the Fourier projection-slice theorem.

Given a set of 3D frequency coordinates lying on a surface (a central slice, or a curved Ewald sphere surface), interpolate the voxel grid onto those coordinates. The voxel grid is assumed to be stored in the half-space (rfft) convention along its last axis, with the two full axes fftshifted to the center convention (see the example below and cryojax.simulator.FourierVoxelGridVolume).

The output is returned in the rfft/DC-in-corner convention, so it can be passed directly to jax.numpy.fft.irfftn (for a half slice) or jax.numpy.fft.ifftn (for a full Ewald sphere surface).

Preparing the voxel grid and extracting a central slice

The voxel grid must be prepared with the convention used internally by cryojax.simulator.FourierVoxelGridVolume: fftshift the object in real space, then fftshift the two full axes in fourier space. cryojax.ndimage.prepare_sampling_fft does this for you.

import cryojax.ndimage as im

# A cubic, even-dimension voxel grid in real space
real_voxel_grid = ...  # shape (dim, dim, dim)
dim = real_voxel_grid.shape[0]

sampling_fft = im.prepare_sampling_fft(real_voxel_grid)

# The (unrotated) central-slice coordinate system, zero-centered
frequency_slice = im.make_frequency_slice((dim, dim), fftshifted=True)

# Extract the slice and transform back to a real-space projection
projection_fft = im.sample_fft_slice(sampling_fft, frequency_slice)
projection = jnp.fft.irfftn(projection_fft, s=(dim, dim))

Rotating a central slice with cryojax.rotations.SO3

import jax
import cryojax.ndimage as im
from cryojax.rotations import SO3

# A half (rfft) central slice of shape (1, dim, dim//2+1, 3)
frequency_slice = im.make_frequency_slice((dim, dim), fftshifted=True)

# Rotate the coordinate system by a random rotation.
rotation = SO3.sample_uniform(jax.random.key(0))
rotated_slice = rotation.apply(frequency_slice)

projection_fft = im.sample_fft_slice(sampling_fft, rotated_slice)

Arguments:

  • sampling_fft: The fourier-space voxel grid, truncated to the half-space (dim, dim, dim // 2 + 1) and prepared as described above, i.e. by cryojax.ndimage.prepare_sampling_fft.
  • frequency_slice: The 3D frequency coordinates to interpolate onto, in pixel units. This can either be a half (rfft) central slice of shape (1, dim, dim // 2 + 1, 3), as returned by cryojax.ndimage.make_frequency_slice, or a full Ewald sphere surface of shape (1, dim, dim, 3), as returned by cryojax.ndimage.ewald_sphere_from_slice.
  • interp: The interpolation method, either "linear" or "cubic". This must match the interp that sampling_fft was prepared with --- see cryojax.ndimage.prepare_sampling_fft.
  • boundary: What to return for frequencies that fall outside the fourier box, which happens at the corners of a rotated slice and for Ewald sphere surfaces curving out of it. Either "fill" (the default), which returns zero, or "clip", which clamps them onto the edge of the box.
  • unroll: See cryojax.ndimage.map_coordinates. For interp="cubic", unroll=False is often substantially faster on GPU.

Returns:

The extracted surface in the rfft/DC-in-corner convention. Shape (dim, dim // 2 + 1) for a half central slice, or (dim, dim) for a full Ewald sphere surface.

cryojax.ndimage.ewald_sphere_from_slice(frequency_slice: Float[Array, '1 dim dim//2+1 3'], voxel_size: cryojax.jax_util.FloatLike, wavelength: cryojax.jax_util.FloatLike) -> Float[Array, '1 dim dim 3'] ¤

Curve a central slice onto the Ewald sphere surface.

Take a half (rfft) central-slice coordinate system, reconstruct the full in-plane grid, and displace each in-plane frequency out of the plane onto the curved Ewald sphere surface. The result can be passed as the frequency_slice argument of cryojax.ndimage.sample_fft_slice.

Example

import cryojax.ndimage as im

# A half (rfft) central slice from `make_frequency_slice`
frequency_slice = im.make_frequency_slice((dim, dim), fftshifted=True)

frequency_surface = im.ewald_sphere_from_slice(
    frequency_slice, voxel_size, wavelength
)
surface = im.sample_fft_slice(sampling_fft, frequency_surface)

Arguments:

  • frequency_slice: The half (rfft) central-slice coordinate system of shape (1, dim, dim // 2 + 1, 3), as returned by cryojax.ndimage.make_frequency_slice. This will typically be a rotated slice.
  • voxel_size: The voxel size, in units of length.
  • wavelength: The electron wavelength, in units of length.

Returns:

The Ewald sphere surface coordinates of shape (1, dim, dim, 3).