Skip to content

Add multidimensional drift correction and strip-based alignment - EDS, 4DSTEM prior to arXiv - #283

Open
bobleesj wants to merge 16 commits into
electronmicroscopy:driftfrom
bobleesj:drift-paper
Open

bobleesj wants to merge 16 commits into
electronmicroscopy:driftfrom
bobleesj:drift-paper

Conversation

@bobleesj

@bobleesj bobleesj commented Aug 25, 2026 •

Copy link
Copy Markdown
Collaborator

What problem this PR addreseses

Support 3D, 4DSTEM drift correction.

PRs broken down as shown below:

Before

src/quantem/imaging/
├── drift.py              1,832 lines
└── drift_utils.py          664 lines

After

├── correction.py         orchestration
├── preprocess.py        input preparation
├── apply.py              dataset propagation
├── diagnostics.py        numerical diagnostics
├── report.py             structured summary
├── io.py                 EMD metadata and pairing
├── fourdstem.py           4D-STEM collection logic
├── plot.py               visual diagnostics
└── core/
    ├── affine.py
    ├── strip.py
    ├── nonrigid.py
    ├── knots.py
    └── warping.py

Public API change

existing usage proposed usage why
from_data(...) from_images(...) 2D images - different from from 4dstem, 3d dataset
align_translation() Same translation is align
align_affine() correct_affine() affine is correcting knots
align_nonrigid() correct_nonrigid() non-rigid is correcting knots
generate_corrected() corrected() concise, pythonic
number_knots, step num_knots, max_drift_rate clarifies the model count and physical search meaning

Affine speed:

metric scipy full, cpu (2016 paper) torch full, cuda torch automatic pyramid, cuda
algorithm time 2439.956 s 5.788 s 0.441 s
speedup vs scipy 1× 421.5× 5536.0×
drift rate, row (px/line) -0.073000 -0.073000 -0.072461
drift rate, col (px/line) 0.139000 0.139000 0.139648

The pyramid downsampling strategy enables real-time drift correction on any personal laptops.

What should the reviewer(s) do

Tutorials are available here: https://github.com/bobleesj/quantem-tutorials/tree/drift-tutorial/tutorials/imaging/drift
.

Screenshot 2026-08-26 at 10 31 08 AM Screenshot 2026-08-26 at 10 42 46 AM

bobleesj and others added 12 commits August 24, 2026 21:25
refactor: split drift into a package without behavior changes
* feat: add drift registration diagnostics

* feat: add structured drift reports

* feat: add static drift diagnostic plots

* test: validate drift diagnostics on frozen results

* fix: keep drift diagnostics consistent with fitted state
* feat: add reference-based spectrum-image correction

* fix: reject ambiguous spectrum-image reference matches
* feat: add reference-based spectrum-image correction

* fix: reject ambiguous spectrum-image reference matches

* feat: propagate drift across 4D-STEM scan axes

* test: protect 4D-STEM detector coordinates

* test: compare full detector signature field

* test: verify 4D-STEM coverage and output residency

* fix: stream native 4D-STEM scan correction

* test: compare native detector frames across backends

* test: set float32 backend parity tolerance
* feat: add reference-based spectrum-image correction

* fix: reject ambiguous spectrum-image reference matches

* feat: propagate drift across 4D-STEM scan axes

* test: protect 4D-STEM detector coordinates

* test: compare full detector signature field

* test: verify 4D-STEM coverage and output residency

* fix: stream native 4D-STEM scan correction

* test: compare native detector frames across backends

* test: set float32 backend parity tolerance

* feat: preserve calibration during drift downsampling

* test: lock downsampled coordinate provenance

* perf: add average-pooled affine pyramid search

* test: compare native and pyramid affine rates

* feat: add strip-wise residual correction primitives

* fix: make nonrigid regularization portable across torch backends

* feat: propagate residual fields to calibrated element maps

* test: lock residual-correction publication contracts

* test: recover known strip residual shifts

* feat: complete automatic affine manuscript parity

* feat: complete multi-knot manuscript parity

* test: verify complete manuscript workflow parity

* fix: keep 4D-STEM integration usable without optional GPU package

* refactor: adopt the scientist-facing Drift API

* test: verify manuscript output frame contracts

* fix: preserve natural corrected output frames

* feat: expose manual translation alignment
@bobleesj

Copy link
Copy Markdown
Collaborator Author

@wwmills ready for review.

If you agree with the API proposed and folder architecture, we can probably merge so that you can add your new algo as needed and refactor with upcoming PRs. This PR is to enable further scaling of features and reproduce manuscript results.

@bobleesj
bobleesj requested a review from wwmills August 26, 2026 17:41
@bobleesj
bobleesj marked this pull request as ready for review August 26, 2026 17:41
wwmills and others added 4 commits September 20, 2026 11:00
…#24)

* feat: center-out scanline drift correction, with low-pass alignment costs

Adds a scipy non-rigid stage that solves scanlines one at a time, outward
from the center of the slow axis, each row seeded from the neighbor one row
closer to the center. Fixing the center row fixes the gauge, which is where
most of the uncontrolled error in an unanchored per-row solve lives.

core/centerout.py
    solve_rows_center_out() solves one scan's rows center-out; the prediction
    for each row is both the L-BFGS-B start point and the center of a
    step_cap_px box. correct_nonrigid_center_out() is the sweep loop:
    center-out solve against the mean of the other warped scans, residual
    knot smoothing, re-warp, translation alignment. It owns its own loop
    rather than branching correct_nonrigid, whose body differs at nearly
    every step, so the torch and scipy paths coexist as two methods.

core/knots.py
    transform_row_numpy() maps one scanline's knots to canvas coordinates in
    numpy, for K=1 and K>=2. The per-row scipy cost needs one row without
    building the full (H, W) torch coordinate grid. Conventions match the
    torch path exactly: the fast axis spans num_rows - 1 / num_cols - 1, and
    num_rows is the full image height. A mismatch there would not raise --
    the row cost and the warp would optimize different problems.

core/warping.py, core/affine.py, core/nonrigid.py
    soften_and_lowpass() restricts each image to its valid-weight bounding
    box, applies a separable linear ramp, mean-subtracts inside the mask, and
    optionally Gaussian low-passes in Fourier space. Threaded through
    translate_align, cross_corr_batch, the affine grid search (where the ramp
    doubles as cost_taper) and correct_nonrigid. All defaults leave existing
    call sites unchanged.

imaging/a_sites.py
    fit_a_sites() finds peaks, indexes them against the lattice and fits a
    2-D Gaussian per site; refit_adaptive() re-fits within an n_sigma window.

Also switches a few docstrings to US spelling.

* feat: bond-cloud statistics in a_sites

Moves the nearest-neighbor bond analysis out of the figure notebook and next
to the A-site fitting that produces its input.

bond_statistics() collects the six nearest-neighbor bond vectors per site by
walking the fitted lattice indices, clips outliers at 4 sigma over three
rounds, and returns the per-direction clouds with their 2-D scatter.
cloud_ellipse() and cloud_half_extent() give the covariance ellipse geometry
and the extent the widest magnified ellipse needs.

All three are numpy-only and consume the dict fit_a_sites() returns, so the
analysis is reusable without matplotlib. Plotting stays in the notebook.

* refactor!: lowpass is a real-space sigma in pixels

soften_and_lowpass took ``lowpass`` as a reciprocal-space scale: the kernel
was exp(-0.5 * lowpass**2 * f**2) with f in cycles per pixel, so lowpass=32
blurred by sigma = 32 / (2*pi) = 5.1 px. Callers had to carry that factor in
their heads, and papers reporting "a 32 pixel Gaussian" were describing a
5 px one.

The kernel is now exp(-0.5 * (2*pi*sigma)**2 * f**2), so ``lowpass`` is the
real-space sigma directly. BREAKING: existing callers must divide their old
value by 2*pi -- lowpass=32 becomes lowpass=5.1.

* fix: apply the low-pass settings to the affine's closing registration

correct_affine ends each search pass by re-registering the pair with
warp_and_translate, and that call received neither the low-pass, the taper,
nor the subpixel mode used by the grid search. The final knot positions --
and the _knots_after_affine checkpoint taken from them -- therefore came from
a plain DFT alignment on unfiltered images, so callers needed a second
align_translation afterwards to get the registration their settings asked
for, and plot_combined(stage="affine") disagreed with the corrected output.

Also collapses lowpass_ramp into cost_taper. The two tapered different
copies of the same images -- one for estimating the shift, one for scoring --
and no caller set them independently. correct_nonrigid and
correct_nonrigid_center_out take cost_taper for the same reason.

---------

Co-authored-by: wwmills <226772206+wwmills@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants