diff --git a/CDTools/tools/measurements.py b/CDTools/tools/measurements.py index d3e94a1..31f7c2f 100644 --- a/CDTools/tools/measurements.py +++ b/CDTools/tools/measurements.py @@ -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: diff --git a/tests/tools/test_measurements.py b/tests/tools/test_measurements.py index 4b75c82..36375c4 100644 --- a/tests/tools/test_measurements.py +++ b/tests/tools/test_measurements.py @@ -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