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, withndim = 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, withndim = len(shape).grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.outputs_rfftfreqs: Return a frequency grid for use withjax.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, withndim = 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, withndim = len(shape).grid_spacing: The grid spacing (i.e. pixel/voxel size), in units of length.outputs_rfftfreqs: Return a frequency grid for use withjax.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 withjax.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 withjax.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: IfTrue, 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
¤
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. Eithernearestorlinear.squared: IfFalse, the whitening filter is the inverse square root of the image power. IfTrue, 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. Setpixel_sizeifoffsetis given in Angstroms, and leave as1.0ifoffsetis 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 bycryojax.ndimage.make_frequency_grid. If not provided, generate on-the-fly.pixel_size: The pixel size of thefrequency_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. IfNone, 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)
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)
¤
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
and in 3D this is
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.
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 thefrequenciespassed 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: TheCallablewrapped into aAbstractFourierOperator. Has signatureout = fn(frequencies, *args, **kwargs)args: Passed tofn.kwargs: Passed tofn.
__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
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)
¤
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 ofdownsample_factordownsample_factor: A scale factor at which to downsampleimage_or_volumeby. Must be a value greater than1.operation: A function such asoperation = lambda x, y: f(x, y), wherexandyare JAX arrays. See [equinox.nn.Pool] (https://docs.kidger.site/equinox/api/nn/pool/#equinox.nn.Pool) for documentation.center_correct: IfTrue, apply a phase shift in the fourier domain to correct the array center after downsampling. Applies only to evendownsample_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 downsampleimage_or_volumeby. Must be a value greater than1.outputs_real_space: IfFalse, theimage_or_volumeis 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: IfTrue, also return the downsample factor actually resolved on each axis, as a tuple the same length asimage_or_volume.ndim. Equal todownsample_factoron every axis whendownsample_factoris 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 lengthinput.ndim, one coordinate array per axis, in array-axis order and in index units (soinput[0, 0]sits at coordinate(0, 0)). Each must be broadcastable to the same shape.order: The order of the interpolation kernel:1for linear,3for cubic B-spline. Cubic is more accurate, at the cost of reading a4^ndimrather than a2^ndimneighborhood per query point.mode: How to extrapolate beyond the edges ofinput, using JAX's out-of-bounds indexing modes, e.g."fill"or"clip".cval: The value returned for out-of-bounds coordinates whenmodeis"fill". Ignored for other modes.unroll: IfTrue(the default), gather the interpolation taps one at a time, which keeps memory use small and predictable at large batch sizes. IfFalse, gather them all at once, which is often substantially faster fororder=3but 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 bycryojax.ndimage.prepare_sampling_fft. Shape(dim, dim // 2 + 1)or(dim, dim, dim // 2 + 1).frequencies: A sequence of lengthinput.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 incryojax.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:1for linear,3for 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: Seecryojax.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,), where0corresponds to the real-space center at grid indexn // 2(for both even and oddn).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 asx,y.n_spread: The width (number of grid points, per axis) of the kernel used to spread each point. Controls speed / accuracy tradeoff: largern_spreadis more accurate but slower. Must be chosen relative tovarianceand the pixel/voxel size — too small truncates the Gaussian and silently biases the result. Seecryojax.ndimage.variance_to_nspreadto pick a value for a givenvariance. Must not exceed the smallest dimension ofshape(otherwise a single point's kernel support would wrap around the grid more than once, aliasing the result).use_erf: IfTrue(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). IfFalse, 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 survivesjax.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 moderateM, with the crossover shifting by architecture) -- benchmark your own workload if this matters.
True/Falseapplies 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 theCRYOJAX_ENABLE_PALLASenvironment variable (Falseif 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 theCRYOJAX_PALLAS_BLOCK_SIZEenvironment variable if you've benchmarked a better value for your own (GPU,M,n_spread). - The pure-JAX forward pass usually wins outright (the Pallas
forward kernel needs an atomic scatter, which contends under
realistic atom density). Leave
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,), where0corresponds to the real-space center at grid indexn // 2(for both even and oddn).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 asx,y,z.n_spread: The width (number of grid points, per axis) of the kernel used to spread each point. Controls speed / accuracy tradeoff: largern_spreadis more accurate but slower. Must be chosen relative tovarianceand the pixel/voxel size — too small truncates the Gaussian and silently biases the result. Seecryojax.ndimage.variance_to_nspreadto pick a value for a givenvariance. Must not exceed the smallest dimension ofshape(otherwise a single point's kernel support would wrap around the grid more than once, aliasing the result).use_erf: IfTrue(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). IfFalse, 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 survivesjax.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 moderateM, with the crossover shifting by architecture) -- benchmark your own workload if this matters.
True/Falseapplies 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 theCRYOJAX_ENABLE_PALLASenvironment variable (Falseif 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 theCRYOJAX_PALLAS_BLOCK_SIZEenvironment variable if you've benchmarked a better value for your own (GPU,M,n_spread). - The pure-JAX forward pass usually wins outright (the Pallas
forward kernel needs an atomic scatter, which contends under
realistic atom density). Leave
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 tospread_gaussians_2d/spread_gaussians_3d.pixel_size: The pixel/voxel size of the gridvariancewill 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 tocryojax.ndimage.sample_fft_slice. Either"linear"(the default), or"cubic", which is substantially more accurate at the cost of reading a4^3rather than a2^3neighborhood per query point.pad_scale: Scale factor at which to Fourier-padreal_voxel_gridbefore 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. bycryojax.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 bycryojax.ndimage.make_frequency_slice, or a full Ewald sphere surface of shape(1, dim, dim, 3), as returned bycryojax.ndimage.ewald_sphere_from_slice.interp: The interpolation method, either"linear"or"cubic". This must match theinterpthatsampling_fftwas prepared with --- seecryojax.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: Seecryojax.ndimage.map_coordinates. Forinterp="cubic",unroll=Falseis 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 bycryojax.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).