mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-10-01 14:22:11 +02:00
Fix a bug with oversampling in the measurement functions
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user