Simple forward model#
Now we have some crystallographic functions and we can handle the detector geometry, we can perform a basic forward model of a single crystal to reassure ourselves that this wasn’t all for nothing!
[1]:
import anri.utils
anri.utils.setup() # before any JAX computation: don't grab most of the GPU memory at start-up
import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
from jax.scipy.spatial.transform import Rotation as jR
from matplotlib import pyplot as plt
import Dans_Diffraction
import numpy as np
import anri.crystal, anri.geom, anri.diffract, anri.io
import time
start = time.time()
Anri, all fundamental functions and transforms are written for single vectors.[2]:
# easy example: many hkls, single B matrix, so we vmap over hkls only, giving us [0, None]
omega_solns_both_vec = jax.vmap(anri.diffract.omega_solns_both, in_axes=[0, None])
sample_to_lab_vec = jax.vmap(anri.geom.sample_to_lab, in_axes=[0, 0, None, None, None, None])
q_lab_to_k_out_vec = jax.vmap(anri.diffract.q_lab_to_k_out, in_axes=[0, None])
raytrace_to_det_vec = jax.vmap(anri.geom.raytrace_to_det, in_axes=[0, None, None, None, None])
q_lab_to_tth_eta_vec = jax.vmap(anri.diffract.q_lab_to_tth_eta, in_axes=[0, None])
Crystallography#
Let’s take a simple Fe CIF file
[3]:
xtl = Dans_Diffraction.Crystal("../../../tests/data/cif/Fe.cif")
lpars = anri.crystal.lattice_parameters(xtl)
sg = anri.crystal.space_group(xtl)
B = anri.crystal.B_matrix(lpars)
We generate some hkls:
[4]:
dsmax = 2.0
wavelength = 0.3
refl = anri.crystal.reflections(lpars, sg, wavelength, dsmax)
F2 = anri.crystal.structure_factors(xtl, refl["hkl"], wavelength)
strong = F2 > 0.01 # drop the reflections the atom positions extinguish
hkls = refl["hkl"][strong]
ring, ring_ds = anri.crystal.rings(refl["ds"][strong])
/tmp/ipykernel_5291/4037507460.py:4: UserWarning: No isotropic thermal factors (U_iso or B_iso) for Fe: intensities have no Debye-Waller attenuation.
F2 = anri.crystal.structure_factors(xtl, refl["hkl"], wavelength)
[5]:
hkls[ring == 0] # the first ring
[5]:
array([[-1, -1, 0],
[-1, 0, -1],
[-1, 0, 1],
[-1, 1, 0],
[ 0, -1, -1],
[ 0, -1, 1],
[ 0, 1, -1],
[ 0, 1, 1],
[ 1, -1, 0],
[ 1, 0, -1],
[ 1, 0, 1],
[ 1, 1, 0]])
Let’s generate a random orientation.
[6]:
key = jax.random.key(time.time_ns())
random_euler = jax.random.uniform(key, shape=(3,), minval=-90.0, maxval=90.0)
U = jR.from_euler('XYZ', random_euler, degrees=True).as_matrix()
U
[6]:
Array([[ 0.24436488, 0.16582116, 0.95539999],
[ 0.40701435, 0.87673515, -0.25627094],
[-0.8801279 , 0.45148512, 0.14675169]], dtype=float64)
Now we can generate some scattering vectors in the sample frame:
[7]:
UB = U @ B
q_sample = (UB @ hkls.T).T
q_sample.shape
[7]:
(380, 3)
Ewald condition#
Now we can determine the omega angles required to diffract:
[8]:
chi = 0.0
wedge = 0.0
dty = 0.0
y0 = 0.0
# define incoming wavevector in the lab frame
k_in_lab = jnp.array([1., 0, 0])
k_in_lab_norm = anri.diffract.scale_norm_k(k_in_lab, wavelength)
# map it into the sample frame
k_in_sample_norm = anri.geom.lab_to_sample(k_in_lab_norm, 0.0, wedge, chi, dty, y0)
# both Friedel solutions at once: omega1 for etasign +1, omega2 for etasign -1.
# A solution exists for both or for neither, so they share one validity mask.
omega1, omega2, valid = omega_solns_both_vec(q_sample, k_in_sample_norm)
omega = jnp.concatenate([omega1, omega2])
valid = jnp.concatenate([valid, valid])
q_sample = jnp.concatenate([q_sample, q_sample])
omega_valid = omega[valid]
q_sample_valid = q_sample[valid]
Into the lab frame#
With the omega angles determined, we can rotate q_sample into the lab frame:
[9]:
q_lab = sample_to_lab_vec(q_sample_valid, omega_valid, wedge, chi, dty, y0)
[10]:
fig, ax = plt.subplots(figsize=(8,8))
ax.scatter(q_lab[:, 1], q_lab[:, 2])
ax.set_aspect(1)
ax.set(xlabel='Lab Y', ylabel='Lab Z')
plt.show()
Into the detector#
Now we can forward-project them into the detector!
Let’s describe the detector, using ImageD11’s parameter names:
[11]:
# detector geometry, with ImageD11's parameter names
pars = {
"y_center": 1000.0,
"z_center": 1000.0,
"y_size": 75.0,
"z_size": 75.0,
"tilt_x": 0.0,
"tilt_y": 0.0,
"tilt_z": 0.0,
"distance": 180e3,
"o11": 1,
"o12": 0,
"o21": 0,
"o22": 1,
}
anri.io.detector_from_pars gives us the detector’s pixel steps (slow and fast) and the position of pixel (0, 0) in the lab frame:
[12]:
det = anri.io.detector_from_pars(pars)
s_step_lab, f_step_lab, det_origin_lab = det["s_step_lab"], det["f_step_lab"], det["det_origin_lab"]
Now we can map into detector space:
[13]:
origin_lab = jnp.array([0., 0, 0])
# get outgoing scattering vector
k_out = q_lab_to_k_out_vec(q_lab, k_in_lab_norm)
# ray-trace it into the detector
sc, fc = raytrace_to_det_vec(k_out, origin_lab, s_step_lab, f_step_lab, det_origin_lab)
Results#
[14]:
fig, ax = plt.subplots(figsize=(8,8))
ax.scatter(fc, sc)
ax.set_aspect(1)
ax.set(xlabel='Detector fast', ylabel='Detector slow')
# set some sensible detector limits
ax.set_xlim(0, 2048)
ax.set_ylim(0, 2048)
plt.show()
[15]:
tth, eta = q_lab_to_tth_eta_vec(q_lab, wavelength)
[16]:
fig, ax = plt.subplots(figsize=(8,6))
ax.scatter(tth, eta, label='Peaks')
ax.vlines(np.degrees(2 * np.arcsin(ring_ds * wavelength / 2)), -25, 25, color='red', label='Unit cell')
ax.set(xlabel=r'$2\theta$', ylabel=r'$\eta$')
ax.legend(loc='upper right')
plt.show()
Index the forward-simulated peaks with ImageD11#
[17]:
import ImageD11.columnfile, ImageD11.parameters, ImageD11.unitcell, ImageD11.indexing, ImageD11.grain
# make an ImageD11 unitcell from our structure
uc = ImageD11.unitcell.unitcell(lpars, sg)
# prepare a minimal columnfile - just add detector positions and omega angles
cf_obs = ImageD11.columnfile.columnfile(new=True)
cf_obs.nrows = fc.shape[0]
cf_obs.addcolumn(fc, 'fc')
cf_obs.addcolumn(sc, 'sc')
cf_obs.addcolumn(omega_valid, 'omega')
# prepare parameters object to hold our experiment state
id11_pars = ImageD11.parameters.parameters()
# detector: the same dict we gave anri
for key, value in pars.items():
id11_pars.set(key, value)
# beam
id11_pars.set('wavelength', wavelength)
# diffractometer
id11_pars.set('chi', chi)
id11_pars.set('wedge', wedge)
id11_pars.set('t_x', 0)
id11_pars.set('t_y', 0)
id11_pars.set('t_z', 0)
id11_pars.set('omegasign', 1)
cf_obs.parameters = id11_pars
print(id11_pars.get_parameters())
{'y_center': 1000.0, 'z_center': 1000.0, 'y_size': 75.0, 'z_size': 75.0, 'tilt_x': 0.0, 'tilt_y': 0.0, 'tilt_z': 0.0, 'distance': 180000.0, 'o11': 1, 'o12': 0, 'o21': 0, 'o22': 1, 'wavelength': 0.3, 'chi': 0.0, 'wedge': 0.0, 't_x': 0, 't_y': 0, 't_z': 0, 'omegasign': 1}
Now we compute the peak geometry with ImageD11:
[18]:
print(cf_obs.titles)
cf_obs.updateGeometry()
print(cf_obs.titles)
['fc', 'sc', 'omega']
['fc', 'sc', 'omega', 'xl', 'yl', 'zl', 'tth', 'eta', 'ds', 'gx', 'gy', 'gz']
Sanity check - we should compute the same g-vectors in the sample frame (q_sample):
[19]:
jnp.abs(jnp.stack([cf_obs.gx, cf_obs.gy, cf_obs.gz], axis=1) - q_sample_valid).max()
[19]:
Array(1.55431223e-15, dtype=float64)
Now let’s set up our indexer and run it:
[20]:
ImageD11.indexing.loglevel = 3
idx = ImageD11.indexing.indexer_from_colfile_and_ucell(cf_obs, uc)
idx.ds_tol = 0.005
idx.assigntorings()
idx.hkl_tol = 0.01
idx.score_all_pairs()
[21]:
id11_ubi = idx.ubis[0]
id11_grain = ImageD11.grain.grain(id11_ubi)
# B matrices should be very similar:
print(jnp.abs(id11_grain.B - B).max())
# the U matrices that we get back should be the same under symmetry:
dU = id11_grain.U.T @ U
print(dU)
4.996003610813204e-16
[[-1.00000000e+00 2.95969171e-17 5.51026946e-16]
[-2.05280006e-17 -1.00000000e+00 1.97920961e-16]
[ 2.61334724e-16 -4.43309402e-17 1.00000000e+00]]
An even simpler check is whether our UBI from Anri indexes the g-vectors from ImageD11:
[22]:
from ImageD11.cImageD11 import score
gve_id11 = jnp.stack([cf_obs.gx, cf_obs.gy, cf_obs.gz], axis=1)
score_result = score(ubi=jnp.linalg.inv(U @ B), gv=gve_id11, tol=0.01)
print(f'Score result: {score_result}, Peaks in dataset: {cf_obs.nrows}')
Score result: 744, Peaks in dataset: 744
created an array from object
created an array from object
[23]:
end = time.time()
print(f'Took {(end - start):.1f} seconds')
Took 3.6 seconds