From b8bdcfeb44114286296122877c8f478662272214 Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Mon, 1 Apr 2019 16:17:52 -0400 Subject: [PATCH] minor tweak to allow for addition of padding --- CDTools/tools/initializers.py | 5 +++-- tests/tools/test_initializers.py | 13 ++++++++++++- 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/CDTools/tools/initializers.py b/CDTools/tools/initializers.py index 6434f12..eb4dfec 100644 --- a/CDTools/tools/initializers.py +++ b/CDTools/tools/initializers.py @@ -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 diff --git a/tests/tools/test_initializers.py b/tests/tools/test_initializers.py index 214f270..5cb6ba3 100644 --- a/tests/tools/test_initializers.py +++ b/tests/tools/test_initializers.py @@ -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])