mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-10-11 18:40:25 +02:00
Add tools to calculate subpixel shifts. Possible improvements still on the table
This commit is contained in:
@@ -71,8 +71,57 @@ def find_subpixel_shift(im1, im2, search_around=(0,0), resolution=10):
|
||||
search_around (array_like) : Default (0,0), the shift to search in the vicinity of
|
||||
resolution (int): Default is 10, the resolution to calculate to in units of 1/n
|
||||
"""
|
||||
pass
|
||||
#
|
||||
# Here's my approach, perhaps it's a little unconventional. I will first
|
||||
# calculate the phase correlation function as found in ____ (cite a paper
|
||||
# defining it). This is strongly peaked, so I can take a small window
|
||||
# of say, 10x10 pixels, and then do a sinc interpolation of that area
|
||||
# using an FFT with upsampling by a factor of resolution in reciprocal
|
||||
# space
|
||||
#
|
||||
# If last dimension is not 2, then convert to a complex tensor now
|
||||
if im1.shape[-1] != 2:
|
||||
im1 = t.stack((im1,t.zeros_like(im1)),dim=-1)
|
||||
if im2.shape[-1] != 2:
|
||||
im2 = t.stack((im2,t.zeros_like(im2)),dim=-1)
|
||||
|
||||
|
||||
cor_fft = cmath.cmult(t.fft(im1,2),cmath.cconj(t.fft(im2,2)))
|
||||
|
||||
# Not sure if this is more or less stable than just the correlation
|
||||
# maximum - requires some testing
|
||||
cor = t.ifft(cor_fft / cmath.cabs(cor_fft)[:,:,None],2)
|
||||
|
||||
# Now, I need to shift the array to pull out a contiguous window
|
||||
# around the correlation maximum
|
||||
try:
|
||||
search_around = search_around.cpu()
|
||||
except:
|
||||
search_around = t.tensor(search_around)
|
||||
|
||||
window_size = 15
|
||||
shift_zero = tuple(-search_around + t.tensor([window_size,window_size]))
|
||||
cor_window = t.roll(cor, shift_zero, dims=(0,1))[:2*window_size,:2*window_size]
|
||||
|
||||
# Now we upsample this window
|
||||
cor_window_fft = cmath.fftshift(t.fft(cor_window,2))
|
||||
upsampled = t.zeros(tuple(t.tensor(cor_window_fft.shape)[:-1] * resolution) + (2,),
|
||||
dtype=cor.dtype,device=cor.device)
|
||||
|
||||
upsampled[:2*window_size,:2*window_size] = cor_window_fft
|
||||
upsampled = t.roll(upsampled,(-window_size,-window_size),dims=(0,1))
|
||||
upsampled = t.roll(cmath.cabssq(t.ifft(upsampled, 2)),(-window_size*resolution,-window_size*resolution), dims=(0,1))
|
||||
|
||||
|
||||
# And we extract the shift from the window
|
||||
sh = t.tensor(upsampled.shape).to(device=upsampled.device)
|
||||
cormax = t.tensor([t.argmax(upsampled) // sh[1],
|
||||
t.argmax(upsampled) % sh[1]]).to(device=upsampled.device)
|
||||
subpixel_shift = ((cormax + sh // 2) % sh - sh//2).to(dtype=upsampled.dtype)
|
||||
|
||||
return search_around.to(device=upsampled.device, dtype=upsampled.dtype) + \
|
||||
subpixel_shift / resolution
|
||||
|
||||
|
||||
def find_pixel_shift(im1, im2):
|
||||
"""Calculates the integer pixel shift between two images by maximizing the autocorrelation
|
||||
@@ -94,9 +143,13 @@ def find_pixel_shift(im1, im2):
|
||||
if im2.shape[-1] != 2:
|
||||
im2 = t.stack((im2,t.zeros_like(im2)),dim=-1)
|
||||
|
||||
|
||||
cor = cmath.cabs(t.ifft(cmath.cmult(t.fft(im1,2),
|
||||
cmath.cconj(t.fft(im2,2))),2))
|
||||
|
||||
cor_fft = cmath.cmult(t.fft(im1,2),cmath.cconj(t.fft(im2,2)))
|
||||
|
||||
# Not sure if this is more or less stable than just the correlation
|
||||
# maximum - requires some testing
|
||||
cor = cmath.cabs(t.ifft(cor_fft / cmath.cabs(cor_fft)[:,:,None],2))
|
||||
#cor = cmath.cabs(t.ifft(cor_fft,2))
|
||||
|
||||
sh = t.tensor(cor.shape).to(device=im1.device)
|
||||
cormax = t.tensor([t.argmax(cor) // sh[1],
|
||||
|
||||
@@ -218,7 +218,7 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10):
|
||||
j = t.arange(probe.shape[1]) - probe.shape[1]//2
|
||||
I,J = t.meshgrid(i,j)
|
||||
I = 2 * np.pi * I.to(t.float32) / probe.shape[0]
|
||||
J = 2 * np.pi * I.to(t.float32) / probe.shape[1]
|
||||
J = 2 * np.pi * J.to(t.float32) / probe.shape[1]
|
||||
I = I.to(dtype=probe.dtype,device=probe.device)
|
||||
J = J.to(dtype=probe.dtype,device=probe.device)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import pytest
|
||||
import numpy as np
|
||||
import torch as t
|
||||
|
||||
from CDTools.tools import image_processing, cmath
|
||||
from CDTools.tools import image_processing, cmath, initializers, interactions
|
||||
from scipy import ndimage
|
||||
|
||||
|
||||
@@ -63,8 +63,30 @@ def test_find_pixel_shift():
|
||||
|
||||
|
||||
def test_find_subpixel_shift():
|
||||
pass
|
||||
# We can do this by creating a test probe and a test object
|
||||
test_probe = t.rand((70,70,2))
|
||||
test_obj = t.ones((300,300,2))
|
||||
|
||||
shift = t.tensor((0.8,0.75))
|
||||
|
||||
im = interactions.ptycho_2D_sinc(test_probe, test_obj, shift)
|
||||
|
||||
retrieved_shift = image_processing.find_subpixel_shift(im, test_probe, search_around=(0,0), resolution=50)
|
||||
# tolerance of 0.03 on this measurement
|
||||
assert t.all(t.abs(shift - retrieved_shift) < 0.03)
|
||||
|
||||
|
||||
def test_find_shift():
|
||||
pass
|
||||
|
||||
# We can do this by creating a test probe and a test object
|
||||
test_probe = t.rand((200,200,2))
|
||||
test_obj = t.ones((300,300,2))
|
||||
|
||||
shift = t.tensor((0.8,0.75))
|
||||
|
||||
im = interactions.ptycho_2D_sinc(test_probe, test_obj, shift)[:-40,:-6]
|
||||
|
||||
retrieved_shift = image_processing.find_shift(im, test_probe[40:,6:], resolution=50)
|
||||
# tolerance of 0.03 on this measurement
|
||||
assert t.all(t.abs(shift + t.Tensor((40,6)) - retrieved_shift) < 0.03)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user