Fix a bug with oversampling in the measurement functions

This commit is contained in:
Abe Levitan
2020-02-11 15:20:55 -05:00
parent 5cac96ff24
commit d8486306f7
2 changed files with 43 additions and 4 deletions
+9 -3
View File
@@ -47,7 +47,10 @@ def intensity(wavefield, detector_slice=None, epsilon=1e-7, saturation=None, ove
# Now we apply oversampling
if oversampling != 1:
output = avg_pool2d(output, oversampling)
if wavefield.dim() == 3:
output = avg_pool2d(output.unsqueeze(0), oversampling)[0]
else:
output = avg_pool2d(output, oversampling)
# Then we grab the detector slice
if detector_slice is not None:
@@ -97,8 +100,11 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=Non
# Now we apply oversampling
if oversampling != 1:
output = avg_pool2d(output, oversampling)
if wavefields.dim() == 4:
output = avg_pool2d(output.unsqueeze(0), oversampling)[0]
else:
output = avg_pool2d(output, oversampling)
# Then we grab the detector slice
if detector_slice is not None:
if wavefields.dim() == 4:
+34 -1
View File
@@ -28,7 +28,25 @@ def test_intensity():
t.tensor(np_result[0][det_slice]))
# With oversampling on
np_oversampling_result = (np_result[:,::2,::2] + \
np_result[:,1::2,::2] + \
np_result[:,::2,1::2] + \
np_result[:,1::2,1::2]) / 4
# With multiple fields
assert t.allclose(measurements.intensity(wavefields,epsilon=epsilon, oversampling=2),
t.tensor(np_oversampling_result,))
# With a single field
assert t.allclose(measurements.intensity(wavefields[0],epsilon=epsilon, oversampling=2),
t.tensor(np_oversampling_result[0],))
def test_incoherent_sum():
# With no explicit slice given
wavefields = t.rand((5,4,10,10,2))
epsilon=1e-6
np_result = np.sum(np.abs(cmath.torch_to_complex(wavefields))**2,axis=0) + epsilon
@@ -38,7 +56,8 @@ def test_incoherent_sum():
assert t.allclose(measurements.incoherent_sum(wavefields[:,0],epsilon=epsilon),
t.tensor(np_result[0]))
# With a slice given
det_slice = np.s_[3:,5:8]
assert t.allclose(measurements.incoherent_sum(wavefields,det_slice,epsilon=epsilon),
t.tensor(np_result[(np.s_[:],)+det_slice]))
@@ -46,6 +65,20 @@ def test_incoherent_sum():
assert t.allclose(measurements.incoherent_sum(wavefields[:,0],det_slice,epsilon=epsilon),
t.tensor(np_result[0][det_slice]))
# With oversampling on
np_oversampling_result = (np_result[:,::2,::2] + \
np_result[:,1::2,::2] + \
np_result[:,::2,1::2] + \
np_result[:,1::2,1::2]) / 4
# With multiple fields
assert t.allclose(measurements.incoherent_sum(wavefields,epsilon=epsilon, oversampling=2),
t.tensor(np_oversampling_result,))
# With a single field
assert t.allclose(measurements.incoherent_sum(wavefields[:,0],epsilon=epsilon, oversampling=2),
t.tensor(np_oversampling_result[0],))
def test_quadratic_background():
# test with intensity