Modeling cryo-EM volumes¤
There are many different volume representations of biological structures for cryo-EM, including atomic models, voxel maps, and neural network representations. Further, there are many ways to generate these volumes, such as from protein generative modeling and molecular dynamics. The optimal implementation to use depends on the user's needs. Therefore, CryoJAX supports a variety of these representations as well as a modeling interface for usage downstream. This page discusses how to use this interface and documents the volumes included in the library.
Core base classes¤
cryojax.simulator.AbstractVolumeParametrization
cryojax.simulator.AbstractVolumeParametrization
¤
Abstract interface for a parametrization of a volume. Specifically, the cryo-EM image formation process typically starts with a scattering potential. "Volumes" and "scattering potentials" in cryoJAX are synonymous.
Info
In, cryojax, potentials should be built in units of inverse length squared,
\([L]^{-2}\). This rescaled potential is defined to be
where \(V\) is the electrostatic potential energy, \(\mathbf{r}\) is a positional coordinate, \(m_0\) is the electron rest mass, and \(e\) is the electron charge.
For a single atom, this rescaled potential has the advantage that under usual scattering approximations (i.e. the first-born approximation), the fourier transform of this quantity is closely related to tabulated electron scattering factors. In particular, for a single atom with scattering factor \(f^{(e)}(\mathbf{q})\) and scattering vector \(\mathbf{q}\), its rescaled potential is equal to
where \(\boldsymbol{\xi} = 2 \mathbf{q}\) is the wave vector coordinate and \(\mathcal{F}^{-1}\) is the inverse fourier transform operator in the convention
The rescaled potential \(U\) gives the following time-independent schrodinger equation for the scattering problem,
where \(k\) is the incident wavenumber of the electron beam.
References:
- For the definition of the rescaled potential, see Chapter 69, Page 2003, Equation 69.6 from Hawkes, Peter W., and Erwin Kasper. Principles of Electron Optics, Volume 4: Advanced Wave Optics. Academic Press, 2022.
- To work out the correspondence between the rescaled potential and the electron scattering factors, see the supplementary information from Vulović, Miloš, et al. "Image formation modeling in cryo-electron microscopy." Journal of structural biology 183.1 (2013): 19-32.
to_representation(rng_key: PRNGKeyArray | None = None) -> cryojax.simulator.AbstractVolumeRepresentation
¤
Core interface for computing the
cryojax.simulator.AbstractVolumeRepresentation for imaging.
Users looking to create custom volumes often won't implement this function
directly, but rather will implement the
a cryojax.simulator.AbstractVolumeRepresentation subclass.
Implementing a cryojax.simulator.AbstractVolumeParametrization
is useful when there is a distinction between how exactly to parametrize
the volume for analysis and how to represent it for imaging.
Arguments:
rng_key: An optional RNG key for including noise / stochastic elements to volume simulation.
cryojax.simulator.AbstractVolumeRepresentation
cryojax.simulator.AbstractVolumeRepresentation(cryojax.simulator.AbstractVolumeParametrization)
¤
Abstract interface for the representation of a volume, such as atomic coordinates, voxels, or a neural network.
Volume representations contain information of coordinates and may be
passed to cryojax.simulator.AbstractVolumeIntegrator
classes for imaging.
rotate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Rotate the coordinate system of the volume.
Volume representations¤
Atom-based volumes¤
cryojax.simulator.AbstractAtomVolume
cryojax.simulator.AbstractAtomVolume(cryojax.simulator.AbstractVolumeRepresentation)
¤
Abstract interface for a volume represented as a point-cloud.
translate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
cryojax.simulator.GaussianMixtureVolume(cryojax.simulator.AbstractAtomVolume)
¤
A representation of a volume as a mixture of gaussians, with multiple gaussians used per position.
The convention of allowing multiple gaussians per position
follows "Robust Parameterization of Elastic and Absorptive
Electron Atomic Scattering Factors" by Peng et al. (1996). The
\(a\) and \(b\) parameters in this work correspond to
amplitudes = a and variances = b / 8\pi^2.
Info
Use the following to load a GaussianMixtureVolume
from these tabulated electron scattering factors.
from cryojax.constants import PengScatteringFactorParameters
from cryojax.io import read_atoms_from_pdb
from cryojax.simulator import GaussianMixtureVolume
# Load positions of atoms and one-hot encoded atom names
atom_positions, atom_types = read_atoms_from_pdb(...)
parameters = PengScatteringFactorParameters(atom_types)
potential = GaussianMixtureVolume.from_tabulated_parameters(
atom_positions, parameters
)
__init__(positions: Float[NDArrayLike, 'M 3'], amplitudes: float | Float[NDArrayLike, ''] | Float[NDArrayLike, 'M'] | Float[NDArrayLike, 'M K'], variances: float | Float[NDArrayLike, ''] | Float[NDArrayLike, 'M'] | Float[NDArrayLike, 'M K'])
¤
Arguments:
positions: The coordinates of the gaussians in units of angstroms.amplitudes: The amplitude for each gaussian. To simulate in physical units of a scattering potential, this should have units of angstroms.variances: The variance for each gaussian. This has units of angstroms squared.
from_tabulated_parameters(atom_positions: Float[NDArrayLike, 'n_atoms 3'], parameters: cryojax.constants.PengScatteringFactorParameters, extra_b_factors: cryojax.jax_util.FloatLike | Float[NDArrayLike, 'n_atoms'] | None = None) -> typing.Self
classmethod
¤
Initialize a GaussianMixtureVolume from tabulated electron
scattering factor parameters (Peng et al. 1996). This treats
the scattering potential as a mixture of five gaussians
per atom.
References:
- Peng, L-M. "Electron atomic scattering factors and scattering potentials of crystals." Micron 30.6 (1999): 625-648.
- Peng, L-M., et al. "Robust parameterization of elastic and absorptive electron atomic scattering factors." Acta Crystallographica Section A: Foundations of Crystallography 52.2 (1996): 257-276.
Arguments:
atom_positions: The coordinates of the atoms in units of angstroms.parameters: A pytree for the scattering factor parameters from Peng et al. (1996).extra_b_factors: Additional per-atom B-factors that are added to the values inscattering_parameters.b.
to_representation(rng_key: PRNGKeyArray | None = None) -> typing.Self
¤
Since this class is itself an
AbstractVolumeRepresentation, this function maps to the identity.
Arguments:
rng_key: Not used in this implementation.
rotate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new potential with rotated positions.
translate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new potential with rotated positions.
cryojax.simulator.GaussianFourierVolume(cryojax.simulator.AbstractAtomVolume)
¤
A representation of a volume that accepts an array of
atom positions and an electron scattering factor for these
atoms, projected/rendered via non-uniform FFTs (see
cryojax.simulator.GaussianFourierProjection/
cryojax.simulator.GaussianFourierRenderFn).
A Gaussian at each atom
import cryojax.simulator as cxs
import cryojax.ndimage as im
positions = ... # load atom positions
b_factor = ... # ... and a B-factor
volume = cxs.GaussianFourierVolume(
positions=positions, kernel_fns=im.FourierGaussian(b_factor=b_factor)
)
The arguments positions and kernel_fns may also be
pytrees of arrays and scattering factors, where each tree leaf represents
a different atom type.
Multiple atom types
import cryojax.simulator as cxs
import cryojax.ndimage as im
positions_1, positions_2 = ...
b_factor_1, b_factor_2 = ...
volume = cxs.GaussianFourierVolume(
positions=(positions_1, positions_2),
kernel_fns=(im.FourierGaussian(b_factor=b_factor_1), im.FourierGaussian(b_factor=b_factor_2))
)
See cryojax.simulator.GaussianFourierVolume.from_tabulated_parameters for
loading a volume from tabulated electron scattering factors.
__init__(positions: PyTree[Float[NDArrayLike, '_ 3'], 'T'], kernel_fns: PyTree[cryojax.ndimage.FourierGaussian, 'T'])
¤
Arguments:
positions: A pytree of atom positions.kernel_fns: A pytree of functions with the same tree structure aspositions, where each leaf is acryojax.ndimage.FourierGaussianrepresenting the atom type's scattering factor. These may have amplitudes and b-factors with a batch dimension to simulate form factors. To use a different amplitude for each atom position, usecryojax.simulator.GaussianMixtureVolumeinstead.
from_tabulated_parameters(positions_by_element: tuple[Float[NDArrayLike, '_ 3'], ...], parameters: cryojax.constants.PengScatteringFactorParameters, *, b_factor_by_element: cryojax.jax_util.FloatLike | tuple[cryojax.jax_util.FloatLike, ...] | None = None) -> typing.Self
classmethod
¤
to_representation(rng_key: PRNGKeyArray | None = None) -> typing.Self
¤
Since this class is itself an
AbstractVolumeRepresentation, this function maps to the identity.
Arguments:
rng_key: Not used in this implementation.
rotate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new potential with rotated positions.
translate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new potential with translated positions.
Voxel-based volumes¤
Fourier-space¤
cryojax.simulator.AbstractVoxelVolume
cryojax.simulator.AbstractVoxelVolume(cryojax.simulator.AbstractVolumeRepresentation)
¤
Abstract interface for a volume represented with voxels.
Info
If you are using a volume in a voxel representation
pass, the voxel size must be passed as the
pixel_size argument, e.g.
import cryojax.simulator as cxs
from cryojax.io import read_array_from_mrc
real_voxel_grid, voxel_size = read_array_from_mrc("example.mrc")
volume = cxs.FourierVoxelGridVolume.from_real_voxel_grid(real_voxel_grid)
...
config = cxs.BasicImageConfig(shape, pixel_size=voxel_size, ...)
If this is not done, the resulting image will be incorrect and not rescaled to the specified to the different pixel size.
Fourier-space conventions
The fourier_voxel_grid and frequency_slice arguments to
FourierVoxelGridVolume.__init__ should be loaded with the zero frequency
component in the center of the box.
cryojax.simulator.FourierVoxelGridVolume(cryojax.simulator.AbstractVoxelVolume)
¤
A volume representation for a 3D voxel grid in fourier-space.
Note
Prefer the class-method constructor from_real_voxel_grid over direct
instantiation. This prepares values for interpolation; only use __init__
assumes if more control is desired.
shape
property
¤
The cubic shape of the volume in real-space.
__init__(values: Complex[NDArrayLike, 'dim dim dim//2+1'], frequency_slice: Float[NDArrayLike, '1 dim dim//2+1 3'], interp: Literal['linear', 'cubic'] = 'linear')
¤
Arguments:
values: The cubic voxel grid in fourier space, truncated to the half-space(dim, dim, dim // 2 + 1)and already prepared for interpolation bycryojax.ndimage.prepare_sampling_fft.frequency_slice: The frequency slice coordinate system, in pixel units. This should be the output ofcryojax.ndimage.make_frequency_slice.interp: The interpolation method used for fourier slice extraction, either"linear"(the default) or"cubic". This should be the same value passed tocryojax.ndimage.prepare_sampling_fft.
from_real_voxel_grid(real_voxel_grid: Float[NDArrayLike, 'dim dim dim'], /, *, interp: Literal['linear', 'cubic'] = 'linear', pad_scale: float = 1.0) -> typing.Self
classmethod
¤
Load from a real-valued 3D voxel grid.
Arguments:
real_voxel_grid: A voxel grid in real space.interp: The interpolation method used for fourier slice extraction, either"linear"(the default) or"cubic". The corresponding interpolation kernel is deconvolved out of the voxel grid here, which is what makes slice extraction accurate --- seecryojax.ndimage.prepare_sampling_fft.pad_scale: Scale factor at which to padreal_voxel_gridbefore fourier transform. Must be a value greater than1.0.
to_representation(rng_key: PRNGKeyArray | None = None) -> typing.Self
¤
Since this class is itself an
AbstractVolumeRepresentation, this function maps to the identity.
Arguments:
rng_key: Not used in this implementation.
rotate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new volume with a rotated frequency_slice.
Real-space¤
Real-space projections
The cryojax.simulator.RealVoxelGridVolume does not have an associated
method in cryoJAX for computing projections and therefore
cannot be used with with cryojax.simulator.make_image_model.
cryojax.simulator.RealVoxelGridVolume(cryojax.simulator.AbstractVoxelVolume)
¤
A 3D voxel grid in real-space.
shape
property
¤
The shape of the voxel grid.
__init__(values: Float[NDArrayLike, 'dim dim dim'], coordinate_grid: Float[NDArrayLike, 'dim dim dim 3'])
¤
Arguments:
values: The voxel grid in real space.coordinate_grid: A coordinate grid, in pixel units.
from_real_voxel_grid(real_voxel_grid: Float[NDArrayLike, 'dim dim dim'], /, *, coordinate_grid: Float[Array, 'dim dim dim 3'] | None = None, crop_scale: float | None = None) -> typing.Self
classmethod
¤
Load a RealVoxelGridVolume from a real-valued 3D voxel grid.
Arguments:
real_voxel_grid: A voxel grid in real space.coordinate_grid: A coordinate grid, in pixel units. Built fromreal_voxel_grid's shape if not given.crop_scale: Scale factor at which to cropreal_voxel_grid. Must be a value greater than1.
to_representation(rng_key: PRNGKeyArray | None = None) -> typing.Self
¤
Since this class is itself an
AbstractVolumeRepresentation, this function maps to the identity.
Arguments:
rng_key: Not used in this implementation.
rotate_to_pose(pose: cryojax.simulator.AbstractPose) -> typing.Self
¤
Return a new volume with a rotated coordinate_grid.
Volume rendering¤
cryojax.simulator.AbstractVolumeRenderFn
cryojax.simulator.AutoVolumeRenderFn(cryojax.simulator.AbstractVolumeRenderFn)
¤
Volume rendering auto selection from cryoJAX
AbstractVolumeRenderFn implementations.
Info
Based on the cryojax.simulator.AbstractVolumeRepresentation passed
at runtime, this class chooses a default rendering function.
In particular,
| Volume representation | Rendering function |
|---|---|
cryojax.simulator.GaussianMixtureVolume |
cryojax.simulator.GaussianMixtureRenderFn |
cryojax.simulator.GaussianFourierVolume |
cryojax.simulator.GaussianFourierRenderFn |
To use advanced options for a given rendering function, see each respective class.
__init__(shape: tuple[int, int, int], voxel_size: cryojax.jax_util.FloatLike, options: dict[str, Any] = {})
¤
Arguments:
shape: The shape of the voxel grid for rendering.voxel_size: The voxel size for rendering.options: Keyword arguments passed to the resolved rendering function, e.g.GaussianMixtureRenderFn(shape, voxel_size, **options).
__call__(volume_representation: cryojax.simulator.AbstractVolumeRepresentation, *, outputs_real_space: bool = True, outputs_rfft: bool = False, fftshifted: bool = False) -> Inexact[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}'] | Complex[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}//2+1']
¤
cryojax.simulator.GaussianMixtureRenderFn(cryojax.simulator.AbstractVolumeRenderFn)
¤
Render a voxel grid from the GaussianMixtureVolume.
If GaussianMixtureVolume is instantiated from electron scattering
factors via from_tabulated_parameters, this renders an electrostatic
potential as tabulated in Peng et al. 1996. The elastic electron
scattering factors defined in this work are
where \(a_i\) is stored as GaussianMixtureVolume.amplitudes,
\(b_i / 8 \pi^2\) are the GaussianMixtureVolume.variances, and
\(\mathbf{q}\) is the scattering vector.
Under usual scattering approximations (i.e. the first-born approximation), the rescaled electrostatic potential energy \(U(\mathbf{r})\) for a given atom type is \(\mathcal{F}^{-1}[f^{(e)}(\boldsymbol{\xi} / 2)](\mathbf{r})\), which is computed analytically as
where \(\mathbf{r}'\) is the position of the atom. Including an additional B-factor (denoted by \(B\)) gives the expression for the potential \(U(\mathbf{r})\) of a single atom type and its fourier transform pair \(\tilde{U}(\boldsymbol{\xi}) \equiv \mathcal{F}[U](\boldsymbol{\xi})\),
where \(\mathbf{q} = \boldsymbol{\xi} / 2\) gives the relationship between the wave vector and the scattering vector.
In practice, for a discretization on a grid with voxel size \(\Delta r\) and grid point \(\mathbf{r}_{\ell}\), the potential is evaluated as the average value inside the voxel
where \(j\) indexes the components of the spatial coordinate vector \(\mathbf{r}\). The above expression is evaluated using the error function as
Speed up gradients with Pallas
render_fn = cxs.GaussianMixtureRenderFn(
shape, voxel_size, n_spread=7, enable_pallas={"bwd": True}
)
__init__(shape: tuple[int, int, int], voxel_size: cryojax.jax_util.FloatLike, *, n_batches: int = 1, n_spread: int | tuple[int, ...] | None = None, enable_pallas: bool | Mapping[str, bool] | None = None)
¤
Arguments:
shape: The shape of the resulting voxel grid.voxel_size: The voxel size of the resulting voxel grid.n_batches: The number of batches over groups of positions used to render the voxel grid. By default,n_batches = 1, which renders the voxel grid for all positions at once. This is useful to decrease GPU memory usage. Applies to both the dense (n_spread=None) and spreading (n_spreadset) backends.n_spread: IfNone(default), render the voxel grid with dense gaussian integrals evaluated over the whole grid. If anint, instead directly spread each gaussian onto only then_spreadnearest grid points (per dimension), trading accuracy for speed. If atupleofints (one value per gaussian component, i.e. of lengthGaussianMixtureVolume.amplitudes.shape[-1]), spread each gaussian component with its own width instead of one shared width -- useful when a volume's gaussian components have widths spanning an order of magnitude or more (e.g. X-ray/electron scattering factors written as a sum of 5 gaussians), where a singlen_spreadwould either truncate the widest components or waste computation spreading the narrowest ones too widely. Seecryojax.simulator.suggest_n_spreadto choose these values fromvolume_representation.variances.enable_pallas: Use the Pallas/Triton GPU backend instead of pure-JAX for then_spreadspreading backend (ignored ifn_spreadisNone). Pallas is JAX's framework for writing custom GPU/TPU kernels. This is most advantageous for the backward pass --{"bwd": True}. Seecryojax.ndimage.spread_gaussians_3d'senable_pallasfor the full picture.None(default) defers toCRYOJAX_ENABLE_PALLAS.
__call__(volume_representation: cryojax.simulator.GaussianMixtureVolume, *, outputs_real_space: bool = True, outputs_rfft: bool = False, fftshifted: bool = False) -> Inexact[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}'] | Complex[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}//2+1']
¤
Arguments:
volume_representation: TheGaussianMixtureVolume.outputs_real_space: IfTrue, return a voxel grid in real-space.outputs_rfft: IfTrue, return a fourier-space voxel grid transformed withcryojax.ndimage.rfftn. Otherwise, usefftn. Does nothing ifoutputs_real_space = True.fftshifted: IfTrue, return a fourier-space voxel grid with the zero frequency component in the center of the grid viajax.numpy.fft.fftshift. Otherwise, the zero frequency component is in the corner. Does nothing ifoutputs_real_space = True.
cryojax.simulator.GaussianFourierRenderFn(cryojax.simulator.AbstractVolumeRenderFn)
¤
Render a voxel grid from a GaussianFourierVolume using non-uniform FFTs
and Fourier-domain convolution. Good when kernels span at least a
couple pixels; see cryojax.simulator.GaussianMixtureRenderFn for an
alternative that directly spreads narrow kernels onto the grid.
Info
By default, the non-uniform FFT runs on a pure-JAX backend using
nufftax.
Setting the environment variable CRYOJAX_FINUFFT_BACKEND=jax-finufft switches to
jax-finufft, which
can be more computationally efficient and less memory-demanding, at the
cost of being trickier to install and having more limited integration
with multi-GPU JAX.
__init__(shape: tuple[int, int, int], voxel_size: cryojax.jax_util.FloatLike, *, sampling_mode: Literal['average', 'point'] = 'average', upsample_factor: int | float = 1.0, eps: float = 1e-06, options: dict[str, Any] = {})
¤
Arguments:
shape: The shape of the resulting voxel grid.voxel_size: The voxel size of the resulting voxel grid.sampling_mode: If'average', convolve with a box function to sample the projected volume at a pixel to be the average value of the underlying continuous function. If'point', the volume at a pixel will be point sampled.upsample_factor: How much to upsample the grid on which atoms are spread onto.eps: The precision of the underlying non-uniform FFT implementation. Seefinufftfor documentation.options: A dictionary of options for advanced usage, passed directly to the underlying non-uniform FFT implementation.
__call__(volume_representation: cryojax.simulator.GaussianFourierVolume, *, outputs_real_space: bool = True, outputs_rfft: bool = False, fftshifted: bool = False) -> Inexact[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}'] | Complex[Array, '{self.shape[0]} {self.shape[1]} {self.shape[2]}//2+1']
¤
Arguments:
volume_representation: TheGaussianFourierVolume.outputs_real_space: IfTrue, return a voxel grid in real-space.outputs_rfft: IfTrue, return a fourier-space voxel grid transformed withcryojax.ndimage.rfftn. Otherwise, usefftn. Does nothing ifoutputs_real_space = True.fftshifted: IfTrue, return a fourier-space voxel grid with the zero frequency component in the center of the grid viajax.numpy.fft.fftshift. Otherwise, the zero frequency component is in the corner. Does nothing ifoutputs_real_space = True.