minor tweak to allow for addition of padding

This commit is contained in:
Abe Levitan
2019-04-01 16:17:52 -04:00
parent 525ab7e91b
commit b8bdcfeb44
2 changed files with 15 additions and 3 deletions
+3 -2
View File
@@ -9,7 +9,7 @@ from scipy.fftpack import next_fast_len
import numpy as np
def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, opt_for_fft=True):
def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, opt_for_fft=True, padding=0):
"""Returns an exit wave basis and a detector slice for the given detector geometry
It takes in the parameters for a given detector - the basis defining
@@ -26,6 +26,7 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None,
distance (float) : The sample-detector distance, in m
center (torch.Tensor) : If defined, the location of the zero frequency pixel
opt_for_fft (bool) : Default is true, whether to increase detector size to improve fft performance
padding (int) : Default is 0, an extra border to allow for subpixel shifting later
Returns:
torch.Tensor : The exit wave basis
@@ -42,7 +43,7 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None,
# This is a bit opaque but was worth doing accurately
min_left = center * 2
min_right = (det_shape - center) * 2 - 1
full_shape = t.max(min_left,min_right).to(t.int32)
full_shape = t.max(min_left,min_right).to(t.int32) + 2 * padding
if opt_for_fft:
full_shape = t.Tensor([next_fast_len(dim) for dim in full_shape]).to(t.int32)
# Then, generate a slice that pops the actual detector from the full
+12 -1
View File
@@ -19,7 +19,7 @@ def test_exit_wave_geometry():
assert t.ones(full_shape)[det_slice].shape == shape
assert t.allclose(rs_basis[0,1],t.Tensor([-8.928571428571428e-07]))
assert t.allclose(rs_basis[1,0],t.Tensor([-4.5662100456621004e-07]))
# Then test it's expanding functionality for a non-optimal array
rs_basis, full_shape, det_slice = \
initializers.exit_wave_geometry(basis, shape, wavelength,
@@ -29,6 +29,17 @@ def test_exit_wave_geometry():
assert t.ones(full_shape)[det_slice].shape == shape
assert t.allclose(rs_basis[0,1],t.Tensor([-8.333333333333333e-07]))
assert t.allclose(rs_basis[1,0],t.Tensor([-4.444444444444444e-07]))
# Then test it's padding function
rs_basis, full_shape, det_slice = \
initializers.exit_wave_geometry(basis, shape, wavelength,
distance, opt_for_fft=False,
padding=2)
exp_shape = t.Size([77,60])
assert full_shape == exp_shape
assert t.ones(full_shape)[det_slice].shape == shape
# Finally test it off-center, without expanding
center = t.Tensor([20,42])