{ "cells": [ { "cell_type": "markdown", "id": "60fd0d2e", "metadata": {}, "source": [ "# Rendering a scanning 3DXRD dataset\n", "\n", "The renderer turns a grain map into the detector frames that a scanning 3DXRD experiment would record.\n", "It renders one dty row at a time, and `anri.io.simulate_sparse` runs it for every row of a scan, writing ImageD11's sparse pixel format. ImageD11 can then process the simulated data like real data.\n", "\n", "This notebook goes through what the renderer does, one step at a time, on a small phantom:\n", "\n", "1. **Select** which peaks can reach the row. A peak is one (map entry, hkl, Friedel branch) combination.\n", "2. **Model** each peak as a 3D Gaussian in (slow, fast, $\\omega$): the centroid comes from the forward model, and the covariance from propagating the spreads of the beam through it.\n", "3. **Integrate** each Gaussian over a small window of detector cells (frames × pixels).\n", "4. **Weight** each peak by its intensity factors: density, $|F|^2$, Lorentz, polarisation, and how much of the voxel the beam illuminates in each frame.\n", "5. **Merge** all peaks into one list of sparse pixels." ] }, { "cell_type": "code", "execution_count": null, "id": "f35f148c", "metadata": {}, "outputs": [], "source": [ "import anri.utils\n", "anri.utils.setup() # before any JAX computation: don't grab most of the GPU memory at start-up\n", "\n", "import os\n", "import tempfile\n", "import time\n", "\n", "import jax\n", "import jax.numpy as jnp\n", "import numpy as np\n", "from matplotlib import pyplot as plt\n", "from ImageD11.sinograms.tensor_map import TensorMap\n", "\n", "import Dans_Diffraction\n", "\n", "import anri.crystal, anri.fwd, anri.io\n", "\n", "start = time.time()" ] }, { "cell_type": "markdown", "id": "6acd4cd9", "metadata": {}, "source": [ "## The map\n", "\n", "The renderer takes a map as a flat list of **entries**, each with a position in the sample frame, a UBI matrix and a density.\n", "Nothing requires one entry per voxel: a voxel can have several entries, e.g. for two orientations.\n", "The renderer handles one phase at a time, so a multi-phase map is rendered phase by phase.\n", "\n", "Our phantom is one slice of a deformed $\\alpha$-quartz grain, with intragranular misorientation and strain, from [flyxdm](https://github.com/AxelHenningsson/flyxdm) (see `tests/data/phantoms/quartz_flyxdm/README.md`).\n", "It is stored as an ImageD11 `TensorMap`, and `anri.io.entries_from_tensormap` turns it into entries:" ] }, { "cell_type": "code", "execution_count": null, "id": "634edc30", "metadata": {}, "outputs": [], "source": [ "tmap = TensorMap.from_h5(\"../../../tests/data/phantoms/quartz_flyxdm/quartz_flyxdm_tmap.h5\")\n", "entries = anri.io.entries_from_tensormap(tmap)\n", "print({k: v.shape for k, v in entries.items()})\n", "\n", "fig, ax = plt.subplots(figsize=(5, 5))\n", "sc = ax.scatter(entries[\"pos\"][:, 0], entries[\"pos\"][:, 1], s=2, c=entries[\"ubi\"][:, 0, 0])\n", "fig.colorbar(sc, ax=ax, shrink=0.8, label=\"UBI[0, 0]\")\n", "ax.set(aspect=1, xlabel=\"sample x (µm)\", ylabel=\"sample y (µm)\", title=\"Map entries\")\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "2be7f5ee", "metadata": {}, "source": [ "## Crystallography\n", "\n", "Each peak needs its hkl and its structure factor $|F|^2$.\n", "`anri.crystal.reflections` lists every hkl separately (not merged into rings), and `anri.crystal.structure_factors` gives $|F|^2$ for each." ] }, { "cell_type": "code", "execution_count": null, "id": "643ca097", "metadata": {}, "outputs": [], "source": [ "wavelength = 0.2845704100778472 # angstrom, as in the flyxdm simulation\n", "xtl = Dans_Diffraction.Crystal(\"../../../tests/data/cif/SiO2.cif\")\n", "refl = anri.crystal.reflections(anri.crystal.lattice_parameters(xtl), anri.crystal.space_group(xtl), wavelength, 1.0)\n", "F2 = anri.crystal.structure_factors(xtl, refl[\"hkl\"], wavelength)\n", "strong = F2 > 0.01 # drop the reflections the atom positions extinguish\n", "hkls = refl[\"hkl\"][strong].astype(float)\n", "F2 = F2[strong]\n", "print(len(hkls), \"hkls\")" ] }, { "cell_type": "markdown", "id": "91692b1d", "metadata": {}, "source": [ "## Geometry\n", "\n", "`anri.io.geom_from_pars` builds the geometry from ImageD11 parameters (detector, beam and goniometer).\n", "Note that ImageD11's wedge has the opposite sign to anri's.\n", "On top of those, the renderer needs the things that set the peak shapes and weights:\n", "\n", "| parameter | meaning |\n", "|---|---|\n", "| `sig_wavelength` | standard deviation of the wavelength (Å) |\n", "| `sig_ky`, `sig_kz` | standard deviations of the horizontal and vertical beam divergence (radians) |\n", "| `sig_beam`, `width_beam` | the beam's profile across it, horizontally: a flat top of `width_beam` (default 0, a Gaussian) blurred by a Gaussian of `sig_beam` (same units as dty, here µm) |\n", "| `sig_beam_v`, `width_beam_v` | the same vertically; only needed for 3D maps of cube voxels (`voxel_3d=True`), as a 2D map's voxels are columns that the whole vertical profile crosses |\n", "| `voxel_size` | side length of the (square) voxels, same units as dty |\n", "| `pol_factor` | degree of horizontal polarisation, 1 for a fully horizontally polarised beam |\n", "| `sig_psf` | standard deviation of the detector point spread (pixels) |" ] }, { "cell_type": "code", "execution_count": null, "id": "558bc37b", "metadata": {}, "outputs": [], "source": [ "pars = {\n", " \"y_center\": 1049.9, \"y_size\": 75.0, \"tilt_y\": 0.0,\n", " \"z_center\": 1116.5, \"z_size\": 75.0, \"tilt_z\": 0.0, \"tilt_x\": 0.0,\n", " \"distance\": 150e3, \"o11\": -1, \"o12\": 0, \"o21\": 0, \"o22\": -1,\n", " \"wavelength\": wavelength, \"wedge\": 0.0, \"chi\": 0.0,\n", "}\n", "det_shape = (2162, 2068) # (slow, fast) pixels\n", "\n", "geom = anri.io.geom_from_pars(\n", " pars, y0=0.0, sig_wavelength=wavelength * 1e-4, sig_ky=1e-4, sig_kz=1e-4,\n", " sig_beam=0.5, voxel_size=float(tmap.steps[1]), sig_psf=0.3,\n", ")" ] }, { "cell_type": "markdown", "id": "8debe435", "metadata": {}, "source": [ "## The row\n", "\n", "A dty **row** is one ImageD11 scan: a series of frames, each with its own $\\omega$ and dty.\n", "`anri.fwd.make_row` sorts the frames by $\\omega$ and puts each frame's edges halfway between neighbouring $\\omega$ values, so frames don't have to be evenly spaced or recorded in order.\n", "Here we take one row at dty = 0: 180° of rotation in 0.25° steps." ] }, { "cell_type": "code", "execution_count": null, "id": "0f7dd295", "metadata": {}, "outputs": [], "source": [ "ostep = 0.25\n", "omega = np.arange(0.0, 180.0, ostep) + ostep / 2 # frame centres\n", "row = anri.fwd.make_row(omega, dty=np.zeros_like(omega))\n", "print(len(omega), \"frames, omega from\", row[\"omega_min\"], \"to\", row[\"omega_max\"])" ] }, { "cell_type": "markdown", "id": "638f030e", "metadata": {}, "source": [ "## Step 1: which peaks reach this row?\n", "\n", "Every (entry, hkl, branch) is a candidate peak. The branch picks one of the two $\\omega$ solutions of the Ewald condition (`etasign` = +1 or -1, ImageD11's omega1 and omega2).\n", "\n", "For each candidate, `anri.fwd.select_peaks` computes the centroid (slow, fast, $\\omega$) with the voxel at its real position for the row's dty, and keeps the peak if its centroid is within a margin of the row (on the detector, inside the $\\omega$ range) and its voxel is within reach of the beam at that $\\omega$: no further from the beam's centre line than half its width + $4\\sigma$ + one voxel.\n", "`anri.fwd.render_row` does this for you. We call it directly here to look at the result:" ] }, { "cell_type": "code", "execution_count": null, "id": "c0c458a9", "metadata": {}, "outputs": [], "source": [ "window = (3, 7, 7) # (frames, slow, fast) cells per peak: render_row's default\n", "margin = jnp.array([\n", " window[1] // 2 + 1, # slow (pixels)\n", " window[2] // 2 + 1, # fast (pixels)\n", " (window[0] // 2 + 1) * ostep, # omega (degrees)\n", " geom[\"width_beam\"] / 2 + 4 * geom[\"sig_beam\"] + geom[\"voxel_size\"], # distance from the beam, across it\n", " geom[\"width_beam_v\"] / 2 + 4 * geom[\"sig_beam_v\"] + geom[\"voxel_size\"], # and vertically (cubes only)\n", "])\n", "ent = {k: jnp.asarray(v) for k, v in entries.items()}\n", "geom_j = jax.tree.map(jnp.asarray, geom)\n", "row_j = jax.tree.map(jnp.asarray, row)\n", "\n", "mask = anri.fwd.select_peaks(ent[\"ubi\"], ent[\"pos\"], jnp.asarray(hkls), geom_j, row_j, margin, det_shape)\n", "e, h, b = np.nonzero(np.asarray(mask)) # entry, hkl and branch of each selected peak\n", "print(f\"{mask.size} candidate peaks, {e.size} reach this row\")" ] }, { "cell_type": "markdown", "id": "f86b8016", "metadata": {}, "source": [ "Most candidates never reach this row: at the $\\omega$ where they diffract, the beam is passing through a different part of the sample.\n", "\n", "## Step 2: one peak as a Gaussian\n", "\n", "A peak's centroid comes from the forward model, as in the simple forward model tutorial.\n", "Its covariance comes from propagating `sig_wavelength`, `sig_ky` and `sig_kz` through the same model with a Jacobian, as in the forward model covariance tutorial.\n", "The spread of the origin over the voxel is not propagated. The voxel's extent and the beam profile enter through the dty weight instead (step 4).\n", "The detector point spread is then added to the slow and fast variances, plus a tiny floor so that a peak with no broadening at all still renders.\n", "\n", "The renderer puts each voxel at its real position for the row's dty. Our row is at dty = $y_0$ = 0, where that is just the voxel's sample position rotated by $\\omega$, so the box-beam forward model, `anri.fwd.get_centroid_box`, gives the renderer's centroid, and `anri.fwd.propagate_cov_box` its covariance.\n", "\n", "Let's take a peak from the voxel nearest the centre of the map that is fairly broad in $\\omega$:" ] }, { "cell_type": "code", "execution_count": null, "id": "d7ef9959", "metadata": {}, "outputs": [], "source": [ "g = geom_j\n", "beam_and_gonio = (g[\"wavelength\"], g[\"k_in_lab\"], 0.0, 0.0, g[\"wedge\"], g[\"chi\"]) # ky = kz = 0: no divergence offset\n", "detector = (g[\"s_step_lab\"], g[\"f_step_lab\"], g[\"det_origin_lab\"])\n", "cov_in = anri.fwd.get_cov_in(jnp.zeros(3), g[\"sig_wavelength\"], g[\"sig_ky\"], g[\"sig_kz\"]) # origin spread 0, see above\n", "\n", "def peak_centroid_and_cov(i):\n", " ubi, pos, hkl, etasign = ent[\"ubi\"][e[i]], ent[\"pos\"][e[i]], jnp.asarray(hkls[h[i]]), 1.0 - 2.0 * b[i]\n", " centroid, valid = anri.fwd.get_centroid_box(ubi, pos, hkl, etasign, *beam_and_gonio, *detector)\n", " cov = anri.fwd.propagate_cov_box(ubi, pos, hkl, etasign, *beam_and_gonio, *detector, cov_in)\n", " return np.asarray(centroid), np.asarray(cov) # (slow, fast, omega) and its 3 x 3 covariance\n", "\n", "centre = np.argmin(np.linalg.norm(entries[\"pos\"], axis=1))\n", "best, best_sd = None, 0.0\n", "for i in np.flatnonzero(e == centre):\n", " c, cov = peak_centroid_and_cov(i)\n", " sd_omega = np.sqrt(cov[2, 2])\n", " inside = 100 < c[0] < det_shape[0] - 100 and 100 < c[1] < det_shape[1] - 100 and 5 < c[2] < 175\n", " if inside and best_sd < sd_omega < 0.2:\n", " best, best_sd = i, sd_omega\n", "i = best\n", "centroid, cov = peak_centroid_and_cov(i)\n", "\n", "# (slow, fast, omega) covariance, plus the point spread on slow and fast.\n", "# The renderer's floor (1e-4 px², 1e-8 deg²) is far smaller than these variances, so we leave it out here.\n", "cov3 = cov[:3, :3] + np.diag([geom[\"sig_psf\"] ** 2, geom[\"sig_psf\"] ** 2, 0.0])\n", "ss, ff, oo = np.diag(cov3)\n", "sf, so, fo = cov3[0, 1], cov3[0, 2], cov3[1, 2]\n", "\n", "print(\"hkl\", hkls[h[i]], \"etasign\", 1 - 2 * b[i])\n", "print(\"centroid (slow, fast, omega):\", centroid.round(3))\n", "print(\"standard deviations (px, px, deg):\", np.sqrt(np.diag(cov3)).round(3))\n", "print(\"correlations slow-fast, slow-omega, fast-omega:\",\n", " (np.array([sf, so, fo]) / np.sqrt([ss * ff, ss * oo, ff * oo])).round(2))" ] }, { "cell_type": "markdown", "id": "adb335b9", "metadata": {}, "source": [ "## Step 3: integrating over a window\n", "\n", "Each peak gets a fixed-size window of cells (frames, slow pixels, fast pixels), centred on the cell that holds its centroid. The default is (3, 7, 7).\n", "Because every peak has the same window size, batches of peaks have fixed array shapes, so they can be compiled once, vectorised and split over devices.\n", "\n", "Integrating a correlated 3D Gaussian over boxes has no closed form, so the renderer factorises it by conditioning:\n", "\n", "1. **$\\omega$:** the mass of the Gaussian in each frame, and the mean and variance of $\\omega$ within that frame (a truncated normal).\n", "2. **slow given $\\omega$:** a 1D Gaussian, whose mean shifts with that within-frame mean of $\\omega$.\n", "3. **fast given slow and $\\omega$:** likewise, using the within-pixel mean of slow.\n", "\n", "Conditioning on the within-cell means, rather than the cell centres, keeps peaks that are narrower than a frame or a pixel at the right position. It also keeps the correlations between slow, fast and $\\omega$: the wavelength spread, for example, smears a spot radially along its Debye-Scherrer ring.\n", "\n", "Let's render our peak with the default window and with two more frames.\n", "The `captured` output is the fraction of the Gaussian that falls inside the window:" ] }, { "cell_type": "code", "execution_count": null, "id": "3928d4b1", "metadata": {}, "outputs": [], "source": [ "one = (e[i : i + 1], h[i : i + 1], b[i : i + 1], ent, jnp.asarray(hkls), jnp.asarray(F2), geom_j, row_j)\n", "for w in [(3, 7, 7), (5, 7, 7)]:\n", " frame, pixel, value, captured = anri.fwd.render_peaks(*one, w, det_shape)\n", " print(f\"window {w}: captured {float(captured[0]):.4f}\")" ] }, { "cell_type": "markdown", "id": "7a397c67", "metadata": {}, "source": [ "This peak is broad in $\\omega$, so three frames lose a few percent of it. `anri.fwd.render_row` returns `captured` for every peak, so you can check this and enlarge the window if needed.\n", "\n", "To check this, we simulate the peak directly: draw the wavelength and beam divergence from their spreads, push every sample through the forward model, add the point spread, and histogram the (slow, fast, $\\omega$) positions into the same cells. This tests the linearised covariance as well as the integration over the window.\n", "Within a row, the per-peak and per-frame factors (step 4) are the same in every cell of a window, so the rendered values divided by their total, times `captured`, are the rendered fraction of the peak in each cell:" ] }, { "cell_type": "code", "execution_count": null, "id": "516a8064", "metadata": {}, "outputs": [], "source": [ "w = (5, 7, 7)\n", "frame, pixel, value, captured = anri.fwd.render_peaks(*one, w, det_shape)\n", "value = np.asarray(value).reshape(w)\n", "pixel = np.asarray(pixel).reshape(w)\n", "frames_in_file = np.asarray(frame).reshape(w)[:, 0, 0]\n", "rendered = value / value.sum() * float(captured[0]) # fraction of the peak in each cell\n", "\n", "# the window's cells: pixel (0, 0) of the window, and the frame edges in sorted-omega order\n", "s0, f0 = divmod(int(pixel[0, 0, 0]), det_shape[1])\n", "jo = np.searchsorted(row[\"omega_edges\"], centroid[2]) - 1 - w[0] // 2\n", "edges = [s0 - 0.5 + np.arange(w[1] + 1), f0 - 0.5 + np.arange(w[2] + 1), np.asarray(row[\"omega_edges\"])[jo : jo + w[0] + 1]]\n", "\n", "# Monte Carlo: beam spreads through the forward model, then the point spread\n", "n = 1_000_000\n", "rng = np.random.default_rng(0)\n", "spreads = rng.standard_normal((n, 3)) * [geom[\"sig_wavelength\"], geom[\"sig_ky\"], geom[\"sig_kz\"]]\n", "ubi, pos, hkl, etasign = ent[\"ubi\"][e[i]], ent[\"pos\"][e[i]], jnp.asarray(hkls[h[i]]), 1.0 - 2.0 * b[i]\n", "\n", "def sample(wavelength, ky, kz):\n", " return anri.fwd.get_centroid_box(ubi, pos, hkl, etasign, wavelength, g[\"k_in_lab\"], ky, kz, *beam_and_gonio[4:], *detector)[0]\n", "\n", "samples = np.array(jax.jit(jax.vmap(sample))(geom[\"wavelength\"] + spreads[:, 0], spreads[:, 1], spreads[:, 2]))\n", "samples[:, :2] += geom[\"sig_psf\"] * rng.standard_normal((n, 2))\n", "mc = np.moveaxis(np.histogramdd(samples, bins=edges)[0], 2, 0) / n # (frame, slow, fast) like the window\n", "\n", "vmax = max(rendered.max(), mc.max())\n", "fig, axs = plt.subplots(2, w[0], figsize=(2.2 * w[0], 4.8), constrained_layout=True)\n", "for k in range(w[0]):\n", " for ax, img, name in ((axs[0, k], rendered[k], \"rendered\"), (axs[1, k], mc[k], \"Monte Carlo\")):\n", " ax.imshow(img, vmin=0, vmax=vmax)\n", " ax.set(title=f\"frame {frames_in_file[k]}\\n{name}\" if name == \"rendered\" else name, xticks=[], yticks=[])\n", "plt.show()\n", "\n", "print(f\"largest difference in any cell: {np.abs(rendered - mc).max():.4f} of the peak (Monte Carlo noise ~{np.sqrt(mc.max() / n):.4f})\")\n", "print(f\"inside the window: rendered {float(captured[0]):.4f}, Monte Carlo {mc.sum():.4f}\")" ] }, { "cell_type": "markdown", "id": "196d317d", "metadata": {}, "source": [ "They agree to within the Monte Carlo noise.\n", "\n", "`anri.fwd.check_render` runs this check for a random sample of peaks from your own map and geometry, and reports, for each, the largest difference in any cell and the fraction inside the window both ways:" ] }, { "cell_type": "code", "execution_count": null, "id": "aa373b24", "metadata": {}, "outputs": [], "source": [ "check = anri.fwd.check_render(entries, hkls, geom, row, det_shape, window=window, n_peaks=100, n_samples=100_000)\n", "err = check[\"max_cell_error\"]\n", "print(f\"{len(err)} peaks: largest cell difference, median {np.median(err):.4f}, 99th percentile {np.percentile(err, 99):.4f}, max {err.max():.4f} of the peak\")\n", "print(f\"largest difference in the fraction inside the window: {np.abs(check['captured'] - check['captured_mc']).max():.4f}\")" ] }, { "cell_type": "markdown", "id": "405a3e61", "metadata": {}, "source": [ "With 100 000 samples per peak, the Monte Carlo noise in a bright cell is about 0.002 of the peak, so these differences are mostly noise.\n", "\n", "## Step 4: intensity\n", "\n", "The value of a peak in a cell is\n", "\n", "$$\n", "\\text{density} \\times |F|^2 \\times L \\times P \\times w_\\text{beam}(\\text{frame}) \\times T(\\text{frame}) \\times \\text{(fraction of the Gaussian in the cell)}\n", "$$\n", "\n", "- $L$ is the Lorentz factor for rotation about the $\\omega$ axis, $1 / |\\hat{\\omega} \\cdot (\\hat{k}_\\text{in} \\times \\hat{k}_\\text{out})|$. For a vertical axis this is $1 / (\\sin 2\\theta\\, |\\sin\\eta|)$, the inverse of ImageD11's `lf`.\n", "- $P$ is the polarisation factor, which matches ImageD11's `polarization`.\n", "- $T$ is an optional transmission per frame, passed to `make_row`.\n", "- $w_\\text{beam}$, from `anri.fwd.beam_weight`, is how much of the voxel the beam illuminates in that frame (its dty moves the voxel across the beam). For a square voxel of side $a$, rotated by $\\omega$, the length of beam path through it is a trapezoid as a function of the distance across the beam, of area $a^2$. The renderer convolves it with the beam's profile across it, a flat top blurred by a Gaussian, in closed form. The profile integrates to 1, so a wider beam spreads the same flux. For cubes (3D maps), the vertical profile is integrated over the cube's height too." ] }, { "cell_type": "code", "execution_count": null, "id": "9c57c807", "metadata": {}, "outputs": [], "source": [ "distance = np.linspace(-2.5, 2.5, 501)\n", "voxel_at = jnp.stack([jnp.zeros_like(distance), distance, jnp.zeros_like(distance)], axis=1) # across the beam (lab y)\n", "weight = jax.vmap(anri.fwd.beam_weight, in_axes=(0, None, None))\n", "beams = [\n", " (\"narrow beam, $\\\\sigma$ = 0.05 µm\", 0.05, 0.0),\n", " (f\"this notebook's beam, $\\\\sigma$ = {geom['sig_beam']} µm\", geom[\"sig_beam\"], 0.0),\n", " (\"flat top 3 µm wide, $\\\\sigma$ = 0.05 µm\", 0.05, 3.0),\n", "]\n", "fig, axs = plt.subplots(1, 3, figsize=(14, 3.5), constrained_layout=True)\n", "for ax, (title, sig, width) in zip(axs, beams):\n", " beam = dict(geom_j, sig_beam=sig, width_beam=width)\n", " for om in (0.0, 30.0, 45.0):\n", " ax.plot(distance, weight(voxel_at, om, beam), label=f\"$\\\\omega$ = {om:.0f}°\")\n", " ax.set(title=title, xlabel=\"voxel's distance from the beam's centre (µm)\")\n", "axs[0].set(ylabel=\"$w_\\\\mathrm{beam}$\")\n", "axs[0].legend()\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "9264aaf2", "metadata": {}, "source": [ "With a narrow beam, you can see the voxel's trapezoid (a rectangle at $\\omega$ = 0°, a triangle at 45°). A beam wider than the voxel smooths that out, and a flat-top beam wider than the voxel lights it fully over most of its width. Integrated over the distance, the weight is always the voxel's area.\n", "\n", "Pencil, line and box beams are all set this way: narrow or wide, Gaussian or flat-top, in each direction.\n", "\n", "## Step 5: a whole row\n", "\n", "`anri.fwd.render_row` does all of the above for every peak of the row:\n", "\n", "- it selects peaks in chunks of entries,\n", "- renders them in fixed-size batches (`batch` peaks at once, split over the devices of `anri.utils.mesh`),\n", "- drops contributions below `min_value`,\n", "- and sorts the rest by (frame, pixel), summing duplicates where peaks overlap.\n", "\n", "It returns sparse pixels, `frame` (index in file order), `pixel` (slow × n_fast + fast) and `value`, plus some statistics." ] }, { "cell_type": "code", "execution_count": null, "id": "bfe0c1d8", "metadata": {}, "outputs": [], "source": [ "t0 = time.time()\n", "frame, pixel, value, stats = anri.fwd.render_row(entries, hkls, F2, geom, row, det_shape, window=window)\n", "print(f\"{stats['n_peaks']} peaks -> {frame.size} sparse pixels in {time.time() - t0:.1f} s\")\n", "low = stats[\"captured\"] < 0.99\n", "print(f\"{low.sum()} peaks ({100 * low.mean():.1f}%) have less than 99% inside their window\")" ] }, { "cell_type": "markdown", "id": "971032b3", "metadata": {}, "source": [ "Most of those peaks are either at the ends of the scanned $\\omega$ range, so part of them is outside the scan, or broad in $\\omega$, like the peak above.\n", "\n", "Summing the row over all frames gives a detector image.\n", "The left panel shows the whole detector. The spots are only a few pixels across, so each is thickened with a maximum filter to make it visible at this scale.\n", "The right panel zooms into the red box and shows the actual pixels.\n", "Each spot sums the peaks of all the voxels along the beam path, so variations of orientation and strain within the grain spread it out:" ] }, { "cell_type": "code", "execution_count": null, "id": "717a76da", "metadata": {}, "outputs": [], "source": [ "from scipy.ndimage import maximum_filter\n", "\n", "image = np.bincount(pixel, weights=value, minlength=det_shape[0] * det_shape[1]).reshape(det_shape)\n", "s, f = int(centroid[0]), int(centroid[1])\n", "\n", "fig, axs = plt.subplots(1, 2, figsize=(11, 5.5), constrained_layout=True)\n", "axs[0].imshow(np.log10(1 + maximum_filter(image, size=9)), cmap=\"gray_r\")\n", "axs[0].add_patch(plt.Rectangle((f - 40, s - 40), 80, 80, fill=False, color=\"r\"))\n", "axs[0].set(title=\"dty = 0, all frames (spots enlarged)\", xlabel=\"fast\", ylabel=\"slow\")\n", "axs[1].imshow(np.log10(1 + image[s - 40 : s + 40, f - 40 : f + 40]), cmap=\"gray_r\", extent=(f - 40.5, f + 39.5, s + 39.5, s - 40.5))\n", "axs[1].set(title=\"Red box, actual pixels\", xlabel=\"fast\", ylabel=\"slow\")\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "f3406561", "metadata": {}, "source": [ "The spots stop well inside the detector because we only generated hkls up to $d^* = 1.0$ Å$^{-1}$ (`dsmax` above).\n", "At this wavelength that is $2\\theta$ = 16.4°, 587 pixels from the beam centre, while the detector reaches $d^* \\approx 1.6$ at its edges and 2.3 in its corners.\n", "A larger `dsmax` fills more of the detector, at the cost of more hkls to render.\n", "\n", "As a check that the intensities land where the forward model says they should, we mark the centroid (slow, fast) of every selected peak on top of them:" ] }, { "cell_type": "code", "execution_count": null, "id": "a61c48eb", "metadata": {}, "outputs": [], "source": [ "def centroid_of(entry, hkl_index, branch):\n", " ubi, pos, hkl = ent[\"ubi\"][entry], ent[\"pos\"][entry], jnp.asarray(hkls)[hkl_index]\n", " return anri.fwd.get_centroid_box(ubi, pos, hkl, 1.0 - 2.0 * branch, *beam_and_gonio, *detector)[0]\n", "\n", "centroids = np.asarray(jax.jit(jax.vmap(centroid_of))(e, h, b)) # (slow, fast, omega) of every selected peak\n", "\n", "def show(ax, half, marker_size):\n", " # rendered intensities (actual pixels) within `half` pixels of our peak, and the centroids of the peaks there\n", " ax.imshow(np.log10(1 + image[s - half : s + half, f - half : f + half]), cmap=\"gray_r\",\n", " extent=(f - half - 0.5, f + half - 0.5, s + half - 0.5, s - half - 0.5))\n", " near = (np.abs(centroids[:, 0] - s) < half) & (np.abs(centroids[:, 1] - f) < half)\n", " ax.scatter(centroids[near, 1], centroids[near, 0], s=marker_size, c=\"r\", marker=\".\", linewidths=0)\n", " ax.set(title=f\"{2 * half} x {2 * half} pixels: {near.sum()} peak centroids\", xlabel=\"fast\", ylabel=\"slow\")\n", "\n", "fig, axs = plt.subplots(1, 2, figsize=(11, 5.5), constrained_layout=True)\n", "show(axs[0], 150, 0.5)\n", "show(axs[1], 40, 2)\n", "plt.show()" ] }, { "cell_type": "markdown", "id": "ab440b53", "metadata": {}, "source": [ "In both regions, every spot sits on the centroids of its peaks, and there are no centroids without a spot." ] }, { "cell_type": "markdown", "id": "ea2391f0", "metadata": {}, "source": [ "## A whole scan\n", "\n", "`anri.io.simulate_sparse` runs `render_row` for each row of a scan and writes ImageD11's sparse pixel format, one group per row (`\"1.1\"`, `\"2.1\"`, ...). Values are rounded to counts, and only counts above `cut` are kept.\n", "`anri.io.write_pars`, `anri.io.write_zero_distortion`, `anri.io.write_dataset` and `anri.io.write_peaks_table` then make the parameter, DataSet and peaks table files that ImageD11's S3DXRD pipeline expects. `tests/unit/io/test_imaged11.py` runs that whole chain.\n", "\n", "Here we simulate the three rows nearest the centre of the sample, into a temporary folder:" ] }, { "cell_type": "code", "execution_count": null, "id": "4bd96371", "metadata": {}, "outputs": [], "source": [ "omega_grid, dty_grid = anri.io.motor_grid((0.0, 180.0), ostep, (-geom[\"voxel_size\"], geom[\"voxel_size\"]), geom[\"voxel_size\"])\n", "print(\"motor grid (rows, frames):\", omega_grid.shape)\n", "\n", "from ImageD11.sparseframe import SparseScan\n", "\n", "with tempfile.TemporaryDirectory() as tmp:\n", " path = os.path.join(tmp, \"sparse.h5\")\n", " scan_stats = anri.io.simulate_sparse(path, entries, hkls, F2, geom, omega_grid, dty_grid, det_shape)\n", " print(\"pixels written per row:\", scan_stats[\"n_pixels\"])\n", " scan = SparseScan(path, \"2.1\") # the middle row, read back by ImageD11\n", " print(\"ImageD11 reads\", scan.shape, \"with\", scan.nnz.sum(), \"non-zero pixels\")" ] }, { "cell_type": "markdown", "id": "7c1c19fc", "metadata": {}, "source": [ "## Settings you may need to change\n", "\n", "- **The beam and the voxels**: `geom_from_pars` takes the beam's profile across it (`sig_beam`, `width_beam` horizontally, `sig_beam_v`, `width_beam_v` vertically) and `voxel_3d` for maps of cubes. A pencil beam is narrow both ways, a line beam wide one way, a box beam (e.g. for DCT) wider than the sample both ways; `k_in_lab` sets its direction.\n", "- **`window`** (default (3, 7, 7) frames × pixels × pixels): peaks broader than their window lose their tails. Check `stats[\"captured\"]` from `render_row`, or run `check_render`. Peaks that are broad in $\\omega$, near $\\eta$ = 0° or 180°, need more frames.\n", "- **The intensity scale, `min_value` and `cut`**: rendered values are density × $|F|^2$ × $L$ × $P$ × $w_\\text{beam}$ × (fraction of the peak), so their scale is set by the densities you give the map. `render_row` drops contributions below `min_value` (default $10^{-3}$) before merging peaks, and `simulate_sparse` rounds to integer counts and keeps those above `cut` (default 1). Choose densities that give realistic counts, or faint peaks disappear.\n", "- **`batch`**: how many peaks are rendered at once. `anri.fwd.guess_batch_size` picks the largest that fits in a fraction of the free memory (default 25%).\n", "- **`anri.utils.setup(n_cpu=4)`**: the number of XLA CPU devices that the work is split over.\n", "\n", "## Speed, memory and devices\n", "\n", "- **Call `anri.utils.setup()` first**, before JAX computes anything. It sets the number of XLA CPU devices that rows are split over, stops JAX reserving most of the GPU memory at start-up, and, with jaxlib 0.11 or newer, turns off XLA:CPU fusions that miscompile the renderer.\n", "- **`batch`** is the number of peaks rendered at once, over all devices. Each peak produces one value per window cell (3 × 7 × 7 = 147 by default), so memory grows with `batch` × window cells: `anri.fwd.guess_batch_size` measures it for your problem.\n", "- **Precision:** this notebook runs in float32, JAX's default. With float64 (`jax.config.update(\"jax_enable_x64\", True)`) every number printed here agrees to within a few parts in $10^4$, at twice the memory.\n", "- **Gradients:** pixel values are differentiable with respect to the map (UBI, position, density) and the geometry. Which pixels a peak touches is not: the window is placed from rounded centroids." ] }, { "cell_type": "code", "execution_count": null, "id": "849fdd38", "metadata": {}, "outputs": [], "source": [ "print(f\"Took {time.time() - start:.0f} seconds\")" ] } ], "metadata": { "language_info": { "name": "python" } }, "nbformat": 4, "nbformat_minor": 5 }