Source code for abtem.magnetism.iam

from __future__ import annotations

from abc import abstractmethod
from typing import TYPE_CHECKING, Callable, Optional, Sequence

import dask.array as da
import numpy as np
from ase import Atoms, units
from ase.data import chemical_symbols
from numba import jit  # type: ignore
from scipy.integrate import trapezoid  # type: ignore
from scipy.interpolate import interp1d  # type: ignore
from scipy.optimize import brentq  # type: ignore

from abtem.core.axes import AxisMetadata, OrdinalAxis, RealSpaceAxis, ThicknessAxis
from abtem.core.energy import energy2sigma
from abtem.core.grid import coordinate_grid
from abtem.inelastic.phonons import BaseFrozenPhonons
from abtem.integrals import cutoff_taper
from abtem.magnetism.parametrizations import LyonParametrization
from abtem.potentials.iam import (
    BaseField,
    FieldArray,
    _FieldBuilderFromAtoms,
)

if TYPE_CHECKING:
    from abtem.potentials.iam import PotentialArray

CUTOFF = 4.25


[docs] def radial_prefactor_a(r: np.ndarray, parameters: np.ndarray) -> Callable: r = r[:, None] a = parameters[None, :, 0] b = parameters[None, :, 1] ni = (np.arange(0, 5) / 2 + 3)[None] a = a / (r**ni + b) a = a.sum(-1) a = a * cutoff_taper(r[:, 0], np.max(r), 0.85) func = interp1d(r[:, 0], a, fill_value=0.0, bounds_error=False) return func
[docs] def radial_prefactor_b1(r: np.ndarray, parameters: np.ndarray) -> Callable: r = r[:, None] a = parameters[None, :, 0] b = parameters[None, :, 1] ni = (np.arange(0, 5) / 2 + 3)[None] b1 = a * ni * r ** (ni - 2) / (r**ni + b) ** 2 b1 = b1.sum(-1) b1 = b1 * cutoff_taper(r[:, 0], np.max(r), 0.85) func = interp1d(r[:, 0], b1, fill_value=0.0, bounds_error=False) return func
[docs] def radial_prefactor_b2(r: np.ndarray, parameters: np.ndarray) -> Callable: r = r[:, None] a = parameters[None, :, 0] b = parameters[None, :, 1] ni = (np.arange(0, 5) / 2 + 3)[None] b2 = a * (2 * b - (ni - 2) * r**ni) / (r**ni + b) ** 2 b2 = b2.sum(-1) b2 = b2 * cutoff_taper(r[:, 0], np.max(r), 0.85) func = interp1d(r[:, 0], b2, fill_value=0.0, bounds_error=False) return func
[docs] def unit_vector_from_angles(theta: np.ndarray, phi: np.ndarray) -> np.ndarray: R = np.sin(theta) m = np.array([R * np.cos(phi), R * np.sin(phi), np.sqrt(1 - R**2)]).T return m
[docs] def atomic_vector_potential_3d( extent: tuple[float, float, float], gpts: tuple[int, int, int], origin: tuple[float, float, float], magnetic_moment: np.ndarray, parameters: np.ndarray, cutoff: float, ) -> np.ndarray: x, y, z = coordinate_grid(extent, gpts, origin, endpoint=False) parameters = np.array(parameters) r = np.sqrt(x**2 + y**2 + z**2) r_interp = np.linspace(0, cutoff, 200) a = radial_prefactor_a(r_interp, parameters) r_vec = np.stack([x, y, z], axis=0) m_cross_r = np.cross(magnetic_moment, r_vec, axis=0) field = a(r)[None] * m_cross_r return field
[docs] def atomic_magnetic_field_3d( extent: tuple[float, float, float], gpts: tuple[int, int, int], origin: tuple[float, float, float], magnetic_moment: np.ndarray, parameters: np.ndarray, cutoff: float, ) -> np.ndarray: magnetic_moment = np.array(magnetic_moment) parameters = np.array(parameters) x, y, z = coordinate_grid(extent, gpts, origin, endpoint=False) r = np.sqrt(x**2 + y**2 + z**2) r_vec = np.stack([x, y, z]) r_interp = np.linspace(0, cutoff, 100) b1 = radial_prefactor_b1(r_interp, parameters) b2 = radial_prefactor_b2(r_interp, parameters) mr = np.sum(r_vec * magnetic_moment[:, None, None, None], axis=0) B = ( b1(r)[None] * r_vec * mr[None] + b2(r)[None] * magnetic_moment[:, None, None, None] ) return B
def _superpose_field_3d( atoms: Atoms, gpts: tuple[int, int, int], atom_field_func: Callable, parameters: Optional[dict] = None, cutoff: Optional[float] = None, ) -> np.ndarray: array = np.zeros((3,) + gpts) if cutoff is None: cutoff = 6.0 if parameters is None: parameters = LyonParametrization().parameters for position, symbol, magnetic_moment in zip( atoms.positions, atoms.symbols, atoms.get_array("magnetic_moments") ): extent = atoms.cell.array.diagonal() array += atom_field_func( extent=extent, gpts=gpts, origin=position, magnetic_moment=magnetic_moment, parameters=parameters[symbol], cutoff=cutoff, ) return array
[docs] def magnetic_field_3d(atoms: Atoms, gpts: tuple[int, int, int], cutoff: float = 6.0): return _superpose_field_3d(atoms, gpts, atomic_magnetic_field_3d, cutoff=cutoff)
[docs] def vector_potential_3d(atoms: Atoms, gpts: tuple[int, int, int], cutoff: float = 6.0): return _superpose_field_3d(atoms, gpts, atomic_vector_potential_3d, cutoff=cutoff)
[docs] def radial_cutoff(func: Callable, tolerance: float = 1e-3): return brentq(lambda x: func(x) - tolerance, a=1e-3, b=1e3)
[docs] def index_mask(indices, shape): mask = (indices[:, 0] >= 0) * (indices[:, 0] < shape[0]) for i, n in enumerate(shape[1:], start=1): mask *= (indices[:, i] >= 0) * (indices[:, i] < n) return mask
[docs] def rotate_points_2d(points, phi): R = np.array([[np.cos(phi), np.sin(phi)], [-np.sin(phi), np.cos(phi)]]) points = R.dot(points.T).T return points
[docs] def cartesian2polar_3d(v: np.ndarray) -> tuple[float, float, float]: r = float(np.linalg.norm(v)) theta = np.arccos(v[2] / r) xy_magnitude = np.linalg.norm(v[:2]) if xy_magnitude > 0.0: phi = np.sign(v[1]) * np.arccos(v[0] / xy_magnitude) else: phi = 0.0 return r, theta, phi
[docs] def symmetric_arange(cutoff: float, sampling: float) -> np.ndarray: cutoff = np.ceil(cutoff / sampling) * sampling values = np.arange(0, cutoff + sampling / 2, sampling) return np.concatenate([-values[::-1][:-1], values])
[docs] @jit(nopython=True, fastmath=True, nogil=True, cache=True) def bilinear_weighted_sum( array: np.ndarray, x: int, y: int, wx0: float, wx1: float, wy0: float, wy1: float ) -> float: return ( array[x, y] * wx0 * wy0 + array[x + 1, y] * wx1 * wy0 + array[x, y + 1] * wx0 * wy1 + array[x + 1, y + 1] * wx1 * wy1 )
[docs] @jit(nopython=True, fastmath=True, nogil=True, cache=True) def interpolate(array_out, array_in, position, sampling_out, sampling_in): nx = array_in.shape[1] ny = array_in.shape[2] scale_x = sampling_out[0] / sampling_in[0] scale_y = sampling_out[1] / sampling_in[1] region_x = int(np.floor(nx // 2 / scale_x)) region_y = int(np.floor(ny // 2 / scale_y)) left = max(int(round(position[0] / sampling_out[0])) - region_x, 0) right = min( int(round(position[0] / sampling_out[0])) + region_x, array_out.shape[1] ) bottom = max(int(round(position[1] / sampling_out[1])) - region_y, 0) top = min(int(round(position[1] / sampling_out[1])) + region_y, array_out.shape[2]) shift_x = np.float32(position[0] / sampling_in[0] - nx // 2) shift_y = np.float32(position[1] / sampling_in[1] - ny // 2) for i in range(left, right): x = np.float32(i * scale_x) - shift_x xf = np.floor(x) wx1 = x - xf wx0 = np.float32(1) - wx1 for j in range(bottom, top): y = np.float32(j * scale_y) - shift_y yf = np.floor(y) wy1 = y - yf wy0 = np.float32(1) - wy1 array_out[0, i, j] += bilinear_weighted_sum( array_in[0], int(xf), int(yf), wx0, wx1, wy0, wy1 ) array_out[1, i, j] += bilinear_weighted_sum( array_in[1], int(xf), int(yf), wx0, wx1, wy0, wy1 ) array_out[2, i, j] += bilinear_weighted_sum( array_in[2], int(xf), int(yf), wx0, wx1, wy0, wy1 )
[docs] @jit(nopython=True, fastmath=True, nogil=True, cache=True) def interpolate_quasi_dipole_field_projections( magnetic_field, sampling, positions, magnetic_moments, slice_limits, integral_limits, integral_sampling, tables, B, ): # B is a caller-allocated (3, tables.shape[2], tables.shape[3]) scratch # buffer, fully overwritten every iteration below: some numba/numpy # pairings fail to type numba's internal np.zeros -> np.empty lowering # inside @njit, so it can't be allocated in here. for position, magnetic_moment in zip(positions, magnetic_moments): shifted_limits = slice_limits - position[2] i = np.argmin(np.abs(integral_limits - shifted_limits[0])) j = np.argmin(np.abs(integral_limits - shifted_limits[1])) j = min(tables.shape[1] - 1, j) b1xxi = tables[0, j] - tables[0, i] b1yyi = b1xxi.T b1xyi = tables[1, j] - tables[1, i] b1xzi = tables[2, j] - tables[2, i] b1yzi = b1xzi.T b1zzi = tables[3, j] - tables[3, i] b2i = tables[4, j] - tables[4, i] B[0] = ( (b1xxi + b2i) * magnetic_moment[0] + b1xyi * magnetic_moment[1] + b1xzi * magnetic_moment[2] ) B[1] = ( b1xyi * magnetic_moment[0] + (b1yyi + b2i) * magnetic_moment[1] + b1xzi.T * magnetic_moment[2] ) B[2] = ( b1xzi * magnetic_moment[0] + b1yzi * magnetic_moment[1] + (b2i + b1zzi) * magnetic_moment[2] ) interpolate(magnetic_field, B, position, sampling, integral_sampling) return magnetic_field
[docs] @jit(nopython=True, fastmath=True, nogil=True, cache=True) def interpolate_quasi_dipole_vector_field_projections( magnetic_field, sampling, positions, magnetic_moments, slice_limits, integral_limits, integral_sampling, tables, A, ): # A is a caller-allocated (3, tables.shape[2], tables.shape[3]) scratch # buffer, fully overwritten every iteration below: some numba/numpy # pairings fail to type numba's internal np.zeros -> np.empty lowering # inside @njit, so it can't be allocated in here. for position, magnetic_moment in zip(positions, magnetic_moments): shifted_limits = slice_limits - position[2] i = np.argmin(np.abs(integral_limits - shifted_limits[0])) j = np.argmin(np.abs(integral_limits - shifted_limits[1])) j = min(tables.shape[1] - 1, j) Ix = tables[0, j] - tables[0, i] Iy = Ix.T Iz = tables[1, j] - tables[1, i] A[0] = magnetic_moment[1] * Iz - magnetic_moment[2] * Iy A[1] = magnetic_moment[2] * Ix - magnetic_moment[0] * Iz A[2] = magnetic_moment[0] * Iy - magnetic_moment[1] * Ix interpolate(magnetic_field, A, position, sampling, integral_sampling) return magnetic_field
[docs] class QuasiDipoleProjections: def __init__( self, interpolation_func, parametrization: str = "lyon", # cutoff_tolerance: float = 1e-3, cutoff: float = CUTOFF, integration_steps: float = 0.01, sampling: float = 0.1, slice_thickness: float = 0.1, ): self._parametrization = LyonParametrization() self._cutoff = cutoff self._step_size = integration_steps self._slice_thickness = slice_thickness self._sampling = sampling self._interpolation_func = interpolation_func self._tables: dict[str, np.ndarray] = {} @property def slice_thickness(self): return self._slice_thickness
[docs] def cutoff(self, symbol): return self._cutoff
def _xy_coordinates(self, symbol): cutoff = self.cutoff(symbol) return symmetric_arange(cutoff, self._sampling) def _slice_limits(self, symbol): n = np.ceil(self.cutoff(symbol) / self.slice_thickness) slice_cutoff = n * self.slice_thickness slice_limits = np.linspace(-slice_cutoff, slice_cutoff, int(n) * 2 + 1) return slice_limits @property def parametrization(self): return self._parametrization @property def finite(self): return True @property def periodic(self): return False @property def sampling(self): return self._sampling @abstractmethod def _calculate_integral_table(self, symbol): pass
[docs] def get_integral_table(self, symbol: str): try: table = self._tables[symbol] except KeyError: table = self._calculate_integral_table(symbol) self._tables[symbol] = table return table
[docs] def integrate_on_grid( self, atoms: Atoms, a: float, b: float, gpts: tuple[int, int], sampling: tuple[float, float], device: str = "cpu", ): if len(atoms) == 0: return np.zeros((3,) + gpts, dtype=np.float32) positions = atoms.positions magnetic_moments = atoms.get_array("magnetic_moments") slice_limits = np.array([a, b]) integral_sampling = (self._sampling,) * 2 array = np.zeros((3,) + gpts, dtype=np.float32) for number in np.unique(atoms.numbers): mask = atoms.numbers == number positions = atoms.positions[mask] magnetic_moments = atoms.get_array("magnetic_moments")[mask] symbol = chemical_symbols[number] if symbol not in self._parametrization.parameters: if not np.allclose(magnetic_moments, 0): raise ValueError(f"Symbol {symbol} is not in the parametrization.") continue integral_limits = self._slice_limits(symbol) tables = self.get_integral_table(symbol) scratch = np.zeros((3, tables.shape[2], tables.shape[3]), dtype=tables.dtype) self._interpolation_func( array, sampling, positions, magnetic_moments, slice_limits, integral_limits, integral_sampling, tables, scratch, ) return array
[docs] class QuasiDipoleMagneticFieldProjections(QuasiDipoleProjections): def __init__( self, parametrization: str = "lyon", # cutoff_tolerance: float = 1e-3, cutoff: float = CUTOFF, integration_steps: float = 0.01, sampling: float = 0.1, slice_thickness: float = 0.1, ): super().__init__( interpolate_quasi_dipole_field_projections, parametrization=parametrization, cutoff=cutoff, integration_steps=integration_steps, sampling=sampling, slice_thickness=slice_thickness, ) def _calculate_integral_table(self, symbol: str) -> np.ndarray: r = np.linspace(0, self.cutoff(symbol), 100) parameters = np.array(self.parametrization.parameters[symbol]) b1_radial = radial_prefactor_b1(r, parameters) b2_radial = radial_prefactor_b2(r, parameters) x = self._xy_coordinates(symbol) slice_limits = self._slice_limits(symbol) shape = (5, len(slice_limits), *(len(x),) * 2) tables = np.zeros(shape, dtype=np.float32) for i, (a, b) in enumerate(zip(slice_limits[:-1], slice_limits[1:]), start=1): n = int(np.round((b - a) / self._step_size)) + 1 z = np.linspace(a, b, n) r = np.sqrt( x[:, None, None] ** 2 + x[None, :, None] ** 2 + z[None, None] ** 2 ) tables[0, i] = tables[0, i - 1] + trapezoid( b1_radial(r) * x[:, None, None] ** 2, x=z, axis=-1 ) tables[1, i] = tables[1, i - 1] + trapezoid( b1_radial(r) * x[:, None, None] * x[None, :, None], x=z, axis=-1 ) tables[2, i] = tables[2, i - 1] + trapezoid( b1_radial(r) * x[:, None, None] * z[None, None, :], x=z, axis=-1 ) tables[3, i] = tables[3, i - 1] + trapezoid( b1_radial(r) * z[None, None, :] ** 2, x=z, axis=-1 ) tables[4, i] = tables[4, i - 1] + trapezoid(b2_radial(r), x=z, axis=-1) return tables
[docs] class QuasiDipoleVectorPotentialProjections(QuasiDipoleProjections): def __init__( self, parametrization: str = "lyon", # cutoff_tolerance: float = 1e-3, cutoff: float = CUTOFF, integration_steps: float = 0.01, sampling: float = 0.1, slice_thickness: float = 0.1, ): super().__init__( interpolate_quasi_dipole_vector_field_projections, parametrization=parametrization, cutoff=cutoff, integration_steps=integration_steps, sampling=sampling, slice_thickness=slice_thickness, ) def _calculate_integral_table(self, symbol): r = np.linspace(0, self.cutoff(symbol), 100) parameters = np.array(self.parametrization.parameters[symbol]) a_radial = radial_prefactor_a(r, parameters) x = self._xy_coordinates(symbol) slice_limits = self._slice_limits(symbol) shape = (2, len(slice_limits), *(len(x),) * 2) tables = np.zeros(shape, dtype=np.float32) for i, (a, b) in enumerate(zip(slice_limits[:-1], slice_limits[1:]), start=1): n = int(np.round((b - a) / self._step_size)) + 1 z = np.linspace(a, b, n) r = np.sqrt( x[:, None, None] ** 2 + x[None, :, None] ** 2 + z[None, None] ** 2 ) Ix = trapezoid(a_radial(r) * x[:, None, None], x=z, axis=-1) tables[0, i] = tables[0, i - 1] + Ix Iz = trapezoid(a_radial(r) * z[None, None], x=z, axis=-1) tables[1, i] = tables[1, i - 1] + Iz return tables
[docs] class BaseMagneticField(BaseField): @property def base_shape(self): """Shape of the base axes of the potential.""" return ( self.num_slices, 3, ) + self.gpts @property def base_axes_metadata(self): """List of AxisMetadata for the base axes.""" return [ ThicknessAxis( label="z", values=tuple(np.cumsum(self.slice_thickness)), units="Å" ), OrdinalAxis( values=("Bx", "By", "Bz"), ), RealSpaceAxis( label="x", sampling=self.sampling[0], units="Å", endpoint=False ), RealSpaceAxis( label="y", sampling=self.sampling[1], units="Å", endpoint=False ), ]
[docs] class BaseVectorPotential(BaseField): @property def base_shape(self): """Shape of the base axes of the potential.""" return ( self.num_slices, 3, ) + self.gpts @property def base_axes_metadata(self): """List of AxisMetadata for the base axes.""" return [ ThicknessAxis( label="z", values=tuple(np.cumsum(self.slice_thickness)), units="Å" ), OrdinalAxis( values=("Ax", "Ay", "Az"), ), RealSpaceAxis( label="x", sampling=self.sampling[0], units="Å", endpoint=False ), RealSpaceAxis( label="y", sampling=self.sampling[1], units="Å", endpoint=False ), ]
[docs] class MagneticFieldArray(BaseMagneticField, FieldArray): _base_dims = 4 def __init__( self, array: np.ndarray | da.core.Array, slice_thickness: float | Sequence[float], extent: Optional[float | tuple[float, float]] = None, sampling: Optional[float | tuple[float, float]] = None, exit_planes: Optional[int | tuple[int, ...]] = None, ensemble_axes_metadata: Optional[list[AxisMetadata]] = None, metadata: Optional[dict] = None, ): if metadata is None: metadata = {} metadata = {"label": "magnetic field", "units": "T", **metadata} super().__init__( array=array, slice_thickness=slice_thickness, extent=extent, sampling=sampling, exit_planes=exit_planes, ensemble_axes_metadata=ensemble_axes_metadata, metadata=metadata, )
[docs] @classmethod def from_array_and_metadata( cls, array: np.ndarray | da.core.Array, axes_metadata: list[AxisMetadata], metadata: dict, ) -> MagneticFieldArray: raise NotImplementedError
[docs] class VectorPotentialArray(BaseVectorPotential, FieldArray): _base_dims = 4 def __init__( self, array: np.ndarray | da.core.Array, slice_thickness: float | Sequence[float], extent: Optional[float | tuple[float, float]] = None, sampling: Optional[float | tuple[float, float]] = None, exit_planes: Optional[int | tuple[int, ...]] = None, ensemble_axes_metadata: Optional[list[AxisMetadata]] = None, metadata: Optional[dict] = None, ): if metadata is None: metadata = {} metadata = {"label": "vector potential", "units": "ÅT", **metadata} super().__init__( array=array, slice_thickness=slice_thickness, extent=extent, sampling=sampling, exit_planes=exit_planes, ensemble_axes_metadata=ensemble_axes_metadata, metadata=metadata, )
[docs] @classmethod def from_array_and_metadata( cls, array: np.ndarray | da.core.Array, axes_metadata: list[AxisMetadata], metadata: dict, ) -> VectorPotentialArray: raise NotImplementedError
[docs] def adjust_coulomb_potential(self, potential_array: PotentialArray, energy: float): # kg * s−2 * A−1 * Å * Å # kg * m2 * s−1 # A * s # A * s * kg-1 * m-2 * s # kg-1 * s2 * A * m-2 e_over_hbar = units._e / (units._hplanck / (2 * np.pi)) * 1e-10 unit_conversion = e_over_hbar / energy2sigma(energy) * 1e-10 adjusted_potential = potential_array.copy() adjusted_potential.array[:] -= self.array[:, 2] * unit_conversion return adjusted_potential
[docs] class MagneticField(_FieldBuilderFromAtoms, BaseMagneticField): _exclude_from_copy = ("parametrization",) def __init__( self, atoms: Atoms | BaseFrozenPhonons, gpts: Optional[int | tuple[int, int]] = None, sampling: Optional[float | tuple[float, float]] = None, slice_thickness: float | tuple[float, ...] = 1, parametrization: str = "lyon", exit_planes: Optional[int | tuple[int, ...]] = None, plane: str | tuple[tuple[float, float, float], tuple[float, float, float]] = "xy", origin: tuple[float, float, float] = (0.0, 0.0, 0.0), box: Optional[tuple[float, float, float]] = None, periodic: bool = True, integrator=None, device: Optional[str] = None, ): if integrator is None: integrator = QuasiDipoleMagneticFieldProjections( parametrization=parametrization ) super().__init__( atoms=atoms, array_object=MagneticFieldArray, gpts=gpts, sampling=sampling, slice_thickness=slice_thickness, exit_planes=exit_planes, device=device, plane=plane, origin=origin, box=box, periodic=periodic, integrator=integrator, )
[docs] class VectorPotential(_FieldBuilderFromAtoms, BaseMagneticField): _exclude_from_copy = ("parametrization",) def __init__( self, atoms: Atoms | BaseFrozenPhonons, gpts: Optional[int | tuple[int, int]] = None, sampling: Optional[float | tuple[float, float]] = None, slice_thickness: float | tuple[float, ...] = 1, parametrization: str = "lyon", exit_planes: Optional[int | tuple[int, ...]] = None, plane: str | tuple[tuple[float, float, float], tuple[float, float, float]] = "xy", origin: tuple[float, float, float] = (0.0, 0.0, 0.0), box: Optional[tuple[float, float, float]] = None, periodic: bool = True, integrator=None, device: Optional[str] = None, ): if integrator is None: integrator = QuasiDipoleVectorPotentialProjections( parametrization=parametrization ) super().__init__( atoms=atoms, array_object=VectorPotentialArray, gpts=gpts, sampling=sampling, slice_thickness=slice_thickness, exit_planes=exit_planes, device=device, plane=plane, origin=origin, box=box, periodic=periodic, integrator=integrator, )