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()
In Anri, all fundamental functions and transforms are written for single vectors.
This was written to significantly simplify the functions themselves, keeping them easy to understand.
Additonally, when forward simulating many grains or voxels, you will likely have more complicated array shapes, so it makes sense to leave the broadcasting to the user or another part of the program for now.
I’m currently still grappling with the best way to expose vmapped functions in the API, so for now I will manually declare them here:
[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()
../_images/tutorials_forward_model_simple_19_0.png

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()
../_images/tutorials_forward_model_simple_28_0.png
[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()
../_images/tutorials_forward_model_simple_30_0.png

Index the forward-simulated peaks with ImageD11#

As a sanity check, we should be able to index the peaks with ImageD11 and recover the UBI.
Let’s prepare the columnfile for 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