Skip to content

Commit ca4fe7c

Browse files
authored
Merge pull request #737 from DiamondLightSource/frames_average
Frames/projection averaging feature
2 parents 6da081f + cfcffd9 commit ca4fe7c

10 files changed

Lines changed: 240 additions & 53 deletions

File tree

‎httomo/method_wrappers/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,11 @@
77
# import all other wrappers to make sure they are available to the factory function
88
# (add imports here when createing new wrappers)
99
import httomo.method_wrappers.datareducer
10+
import httomo.method_wrappers.sino360_to_180
1011
import httomo.method_wrappers.dezinging
1112
import httomo.method_wrappers.distortion_correction
1213
import httomo.method_wrappers.seam_blender
14+
import httomo.method_wrappers.average_frames
1315
import httomo.method_wrappers.images
1416
import httomo.method_wrappers.reconstruction
1517
import httomo.method_wrappers.rotation
Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
from httomo.method_wrappers.generic import GenericMethodWrapper
2+
from httomo.block_interfaces import T
3+
import numpy as np
4+
5+
6+
class AverageFramesWrapper(GenericMethodWrapper):
7+
"""
8+
Wrapper for frames/projection averaging.
9+
"""
10+
11+
@classmethod
12+
def should_select_this_class(cls, module_path: str, method_name: str) -> bool:
13+
return "average_projection_frames" in method_name
14+
15+
def _preprocess_data(self, block: T) -> T:
16+
# when the angular preview is getting changed by averaging the angles should be changed accordingly
17+
config_params = self._config_params
18+
k = config_params["projection_averaging_factor"]
19+
20+
n_proj = block.data.shape[0] # original data angular size
21+
n_full = n_proj // k
22+
remainder = n_proj % k
23+
24+
n_out = n_full + (remainder > 0)
25+
26+
averaged_angles = np.empty(n_out, dtype=block.angles.dtype)
27+
28+
if n_full:
29+
averaged_angles[:n_full] = (
30+
block.angles_radians[: n_full * k].reshape(n_full, k).mean(axis=1)
31+
)
32+
33+
if remainder:
34+
averaged_angles[-1] = block.angles_radians[n_full * k :].mean()
35+
36+
block.angles_radians = averaged_angles
37+
return block

‎httomo/method_wrappers/reconstruction.py‎

Lines changed: 2 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -12,24 +12,17 @@
1212

1313

1414
class ReconstructionWrapper(GenericMethodWrapper):
15-
"""Wraps reconstruction functions, limiting the length of the angles array
16-
before calling the method."""
15+
"""Wraps reconstruction functions."""
1716

1817
@classmethod
1918
def should_select_this_class(cls, module_path: str, method_name: str) -> bool:
2019
return module_path.endswith(".algorithm")
2120

2221
def _preprocess_data(self, block: T) -> T:
23-
# this is essential for the angles cutting below to be valid
2422
assert (
2523
self.pattern == Pattern.sinogram
2624
), "reconstruction methods must be sinogram"
27-
28-
# for 360 degrees data the angular dimension will be truncated while angles are not.
29-
# Truncating angles if the angular dimension has got a different size
30-
datashape0 = block.data.shape[0]
31-
if datashape0 != len(block.angles_radians):
32-
block.angles_radians = block.angles_radians[0:datashape0]
25+
assert len(block.angles_radians) == block.data.shape[0]
3326
self._input_shape = block.data.shape
3427
return super()._preprocess_data(block)
3528

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
from httomo.method_wrappers.generic import GenericMethodWrapper
2+
from httomo.block_interfaces import T
3+
4+
5+
class Sino360to180Wrapper(GenericMethodWrapper):
6+
"""
7+
Wrapper to perform extended FoV (360degrees) data conversion to a standard 180 degrees data.
8+
The wrapper is responsible for changing the angles after the data has changed.
9+
"""
10+
11+
@classmethod
12+
def should_select_this_class(cls, module_path: str, method_name: str) -> bool:
13+
return "sino_360_to_180" in method_name
14+
15+
def _postprocess_data(self, block: T) -> T:
16+
# for 360 degrees data the angular dimension is truncated so the angles should be changed in a similar fashion.
17+
block.angles_radians = block.angles_radians[0 : block.data.shape[0]]
18+
return block

‎httomo/runner/method_wrapper.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -194,6 +194,7 @@ def calculate_max_slices(
194194
data_dtype: np.dtype,
195195
slicing_dim: int,
196196
non_slice_dims_shape: Tuple[int, int],
197+
angles: np.ndarray,
197198
available_memory: int,
198199
) -> Tuple[int, int]:
199200
"""If it runs on GPU, determine the maximum number of slices that can fit in the

‎tests/conftest.py‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -229,6 +229,11 @@ def FBP3d_tomobar_denoising():
229229
return "docs/source/pipelines_full/FBP3d_tomobar_denoising.yaml"
230230

231231

232+
@pytest.fixture
233+
def angles_averaging():
234+
return "docs/source/pipelines_full/angles_averaging.yaml"
235+
236+
232237
@pytest.fixture
233238
def FISTA3d_tomobar():
234239
return "docs/source/pipelines_full/FISTA3d_tomobar.yaml"
@@ -339,6 +344,12 @@ def FBP2d_astra_i12_119647_npz():
339344
return np.load("tests/test_data/raw_data/i12/FBP2d_astra_i12_119647.npz")
340345

341346

347+
@pytest.fixture
348+
def angle_average_LPrec_i12_119647_npz():
349+
# 10 slices numpy array
350+
return np.load("tests/test_data/raw_data/i12/angle_average_LPrec_i12_119647.npz")
351+
352+
342353
@pytest.fixture
343354
def FBP3d_tomobar_TVdenoising_i13_177906_npz():
344355
# 10 slices numpy array
Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,48 @@
1+
from httomo.method_wrappers import make_method_wrapper
2+
from httomo.method_wrappers.sino360_to_180 import Sino360to180Wrapper
3+
from httomo.runner.auxiliary_data import AuxiliaryData
4+
from httomo.runner.dataset import DataSetBlock
5+
from ..testing_utils import make_mock_preview_config, make_mock_repo
6+
from httomo_backends.methods_database.query import Pattern
7+
8+
import numpy as np
9+
from mpi4py import MPI
10+
from pytest_mock import MockerFixture
11+
12+
13+
def test_sino_360_to_180(mocker: MockerFixture):
14+
GLOBAL_SHAPE = (10, 20, 30)
15+
GLOBAL_SHAPE_MOD = (5, 20, 30)
16+
17+
class FakeModule:
18+
def sino_360_to_180_tester(data):
19+
np.testing.assert_array_equal(data, 1)
20+
return data
21+
22+
mocker.patch(
23+
"httomo.method_wrappers.generic.import_module", return_value=FakeModule
24+
)
25+
wrp = make_method_wrapper(
26+
make_mock_repo(mocker, pattern=Pattern.sinogram),
27+
"mocked_module_path.morph",
28+
"sino_360_to_180_tester",
29+
MPI.COMM_WORLD,
30+
make_mock_preview_config(mocker),
31+
)
32+
assert isinstance(wrp, Sino360to180Wrapper)
33+
34+
aux_data = AuxiliaryData(angles=2.0 * np.ones(GLOBAL_SHAPE[0], dtype=np.float32))
35+
data = np.ones(
36+
GLOBAL_SHAPE_MOD, dtype=np.float32
37+
) # assuming the data already averaged here by factor of 2
38+
input = DataSetBlock(
39+
data[:, 0:3, :],
40+
slicing_dim=1,
41+
aux_data=aux_data,
42+
chunk_shape=GLOBAL_SHAPE_MOD,
43+
global_shape=GLOBAL_SHAPE_MOD,
44+
)
45+
46+
wrp.execute(input)
47+
48+
assert aux_data.get_angles().shape[0] == 5
Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
from httomo.method_wrappers import make_method_wrapper
2+
from httomo.method_wrappers.average_frames import AverageFramesWrapper
3+
from httomo.runner.auxiliary_data import AuxiliaryData
4+
from httomo.runner.dataset import DataSetBlock
5+
from ..testing_utils import make_mock_preview_config, make_mock_repo
6+
from httomo_backends.methods_database.query import Pattern
7+
8+
import numpy as np
9+
from mpi4py import MPI
10+
from pytest_mock import MockerFixture
11+
12+
13+
def test_angle_averaging(mocker: MockerFixture):
14+
GLOBAL_SHAPE = (10, 20, 30)
15+
16+
class FakeModule:
17+
def average_projection_frames_tester(data, projection_averaging_factor):
18+
np.testing.assert_array_equal(data, 1)
19+
return data
20+
21+
mocker.patch(
22+
"httomo.method_wrappers.generic.import_module", return_value=FakeModule
23+
)
24+
wrp = make_method_wrapper(
25+
make_mock_repo(mocker, pattern=Pattern.sinogram),
26+
"mocked_module_path.morph",
27+
"average_projection_frames_tester",
28+
MPI.COMM_WORLD,
29+
make_mock_preview_config(mocker),
30+
projection_averaging_factor=2,
31+
)
32+
assert isinstance(wrp, AverageFramesWrapper)
33+
34+
aux_data = AuxiliaryData(
35+
angles=2.0 * np.ones(GLOBAL_SHAPE[0] + 10, dtype=np.float32)
36+
)
37+
data = np.ones(
38+
GLOBAL_SHAPE, dtype=np.float32
39+
) # assuming the data already averaged here by factor of 2
40+
input = DataSetBlock(
41+
data[:, 0:3, :],
42+
slicing_dim=1,
43+
aux_data=aux_data,
44+
chunk_shape=GLOBAL_SHAPE,
45+
global_shape=GLOBAL_SHAPE,
46+
)
47+
48+
wrp.execute(input)
49+
50+
assert aux_data.get_angles().shape[0] == 5

‎tests/method_wrappers/test_reconstruction.py‎

Lines changed: 0 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -1,57 +1,13 @@
1-
import math
21
from httomo.method_wrappers import make_method_wrapper
3-
from httomo.method_wrappers.reconstruction import ReconstructionWrapper
42
from httomo.runner.auxiliary_data import AuxiliaryData
53
from httomo.runner.dataset import DataSetBlock
64
from ..testing_utils import make_mock_preview_config, make_mock_repo
75

8-
from httomo_backends.methods_database.query import Pattern
9-
106
import numpy as np
117
from mpi4py import MPI
128
from pytest_mock import MockerFixture
139

1410

15-
def test_recon_handles_reconstruction_angle_reshape(mocker: MockerFixture):
16-
GLOBAL_SHAPE = (10, 20, 30)
17-
18-
class FakeModule:
19-
# we give the angles a different name on purpose
20-
def recon_tester(data, theta):
21-
np.testing.assert_array_equal(data, 1)
22-
np.testing.assert_array_equal(theta, 2)
23-
assert data.shape[0] == len(theta)
24-
return data
25-
26-
mocker.patch(
27-
"httomo.method_wrappers.generic.import_module", return_value=FakeModule
28-
)
29-
wrp = make_method_wrapper(
30-
make_mock_repo(mocker, pattern=Pattern.sinogram),
31-
"mocked_module_path.algorithm",
32-
"recon_tester",
33-
MPI.COMM_WORLD,
34-
make_mock_preview_config(mocker),
35-
)
36-
assert isinstance(wrp, ReconstructionWrapper)
37-
38-
aux_data = AuxiliaryData(
39-
angles=2.0 * np.ones(GLOBAL_SHAPE[0] + 10, dtype=np.float32)
40-
)
41-
data = np.ones(GLOBAL_SHAPE, dtype=np.float32)
42-
input = DataSetBlock(
43-
data[:, 0:3, :],
44-
slicing_dim=1,
45-
aux_data=aux_data,
46-
chunk_shape=GLOBAL_SHAPE,
47-
global_shape=GLOBAL_SHAPE,
48-
)
49-
50-
wrp.execute(input)
51-
52-
assert aux_data.get_angles().shape[0] == GLOBAL_SHAPE[0]
53-
54-
5511
def test_recon_handles_reconstruction_axisswap(mocker: MockerFixture):
5612
class FakeModule:
5713
def recon_tester(data, theta):

‎tests/test_parallel_pipeline_big.py‎

Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -176,6 +176,77 @@ def test_pipe_parallel_FBP3d_tomobar_k11_38730_in_memory_preview(
176176
assert res_norm < 1e-6
177177

178178

179+
# ########################################################################
180+
@pytest.mark.full_data_parallel
181+
def test_angles_averaging_LPRec_i12_119647_preview(
182+
get_files: Callable,
183+
cmd_mpirun,
184+
i12_119647,
185+
angles_averaging,
186+
angle_average_LPrec_i12_119647_npz,
187+
output_folder,
188+
):
189+
190+
change_value_parameters_method_pipeline(
191+
angles_averaging,
192+
method=[
193+
"standard_tomo",
194+
"average_projection_frames",
195+
],
196+
key=[
197+
"preview",
198+
"projection_averaging_factor",
199+
],
200+
value=[
201+
{"detector_y": {"start": 800, "stop": 1200}},
202+
12,
203+
],
204+
)
205+
206+
cmd_mpirun.insert(9, i12_119647)
207+
cmd_mpirun.insert(10, angles_averaging)
208+
cmd_mpirun.insert(11, output_folder)
209+
210+
process = Popen(
211+
cmd_mpirun, env=os.environ, shell=False, stdin=PIPE, stdout=PIPE, stderr=PIPE
212+
)
213+
output, error = process.communicate()
214+
print(output)
215+
216+
files = get_files(output_folder)
217+
218+
#: check the generated reconstruction (hdf5 file)
219+
h5_files = list(filter(lambda x: ".h5" in x, files))
220+
assert len(h5_files) == 1
221+
222+
# load the pre-saved numpy array for comparison bellow
223+
data_gt = angle_average_LPrec_i12_119647_npz["data"]
224+
axis_slice = angle_average_LPrec_i12_119647_npz["axis_slice"]
225+
slices, sizeX, sizeY = np.shape(data_gt)
226+
227+
step = axis_slice // (slices + 2)
228+
# store for the result
229+
data_result = np.zeros((slices, sizeX, sizeY), dtype=np.float32)
230+
231+
path_to_data = "data/"
232+
h5_file_name = "LPRec3d_tomobar"
233+
for file_to_open in h5_files:
234+
if h5_file_name in file_to_open:
235+
h5f = h5py.File(file_to_open, "r")
236+
index_prog = step
237+
for i in range(slices):
238+
data_result[i, :, :] = h5f[path_to_data][:, index_prog, :]
239+
index_prog += step
240+
h5f.close()
241+
else:
242+
message_str = f"File name with {h5_file_name} string cannot be found."
243+
raise FileNotFoundError(message_str)
244+
245+
residual_im = data_gt - data_result
246+
res_norm = np.linalg.norm(residual_im.flatten()).astype("float32")
247+
assert res_norm < 1e-6
248+
249+
179250
# ########################################################################
180251

181252

0 commit comments

Comments
 (0)