Skip to content

Commit 161dc62

Browse files
authored
handle conveniences (#1390)
2 parents 08614a0 + 1145c42 commit 161dc62

5 files changed

Lines changed: 116 additions & 105 deletions

File tree

src/Registration/pReg/Reg.py

Lines changed: 27 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@
3030
from sirf import SIRF
3131
from sirf.config import SIRF_HAS_SPM
3232
from sirf.SIRF import ContiguousError
33-
from sirf.Utilities import Handle, cpp_int_dtype, error, format_numpy_array_for_setter
33+
from sirf.Utilities import HANDLE, Handle, cpp_int_dtype, error, format_numpy_array_for_setter
3434

3535
if sys.version_info[0] >= 3 and sys.version_info[1] >= 4:
3636
ABC = abc.ABC
@@ -145,9 +145,9 @@ def __init__(self, src=None):
145145
self.handle = None
146146
self.name = 'NiftiImageData'
147147
if src is None:
148-
self.handle = Handle(None, -1).cReg_newObject(self.name)
148+
self.handle = HANDLE.cReg_newObject(self.name)
149149
elif isinstance(src, str):
150-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src)
150+
self.handle = HANDLE.cReg_objectFromFile(self.name, src)
151151
elif isinstance(src, SIRF.ImageData):
152152
# src is ImageData
153153
dim = src.dimensions()
@@ -508,9 +508,9 @@ def __init__(self, src=None):
508508
self.handle = None
509509
self.name = 'NiftiImageData3D'
510510
if src is None:
511-
self.handle = Handle(None, -1).cReg_newObject(self.name)
511+
self.handle = HANDLE.cReg_newObject(self.name)
512512
elif isinstance(src, str):
513-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src)
513+
self.handle = HANDLE.cReg_objectFromFile(self.name, src)
514514
elif isinstance(src, SIRF.ImageData):
515515
# src is ImageData
516516
self.handle = src.handle.cReg_NiftiImageData_from_SIRFImageData(1)
@@ -529,12 +529,11 @@ def __init__(self, src1=None, src2=None, src3=None):
529529
self.handle = None
530530
self.name = 'NiftiImageData3DTensor'
531531
if src1 is None:
532-
self.handle = Handle(None, -1).cReg_newObject(self.name)
532+
self.handle = HANDLE.cReg_newObject(self.name)
533533
elif isinstance(src1, str):
534-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src1)
534+
self.handle = HANDLE.cReg_objectFromFile(self.name, src1)
535535
elif all(isinstance(i, NiftiImageData3D) for i in (src1, src2, src3)):
536-
self.handle = Handle(None, -1).cReg_NiftiImageData3DTensor_construct_from_3_components(
537-
self.name, src1, src2, src3)
536+
self.handle = HANDLE.cReg_NiftiImageData3DTensor_construct_from_3_components(self.name, src1, src2, src3)
538537
else:
539538
raise error('Wrong source in NiftiImageData3DTensor constructor')
540539

@@ -574,12 +573,11 @@ def __init__(self, src1=None, src2=None, src3=None):
574573
self.handle = None
575574
self.name = 'NiftiImageData3DDisplacement'
576575
if src1 is None:
577-
self.handle = Handle(None, -1).cReg_newObject(self.name)
576+
self.handle = HANDLE.cReg_newObject(self.name)
578577
elif isinstance(src1, str):
579-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src1)
578+
self.handle = HANDLE.cReg_objectFromFile(self.name, src1)
580579
elif all(isinstance(i, NiftiImageData3D) for i in (src1, src2, src3)):
581-
self.handle = Handle(None, -1).cReg_NiftiImageData3DTensor_construct_from_3_components(
582-
self.name, src1, src2, src3)
580+
self.handle = HANDLE.cReg_NiftiImageData3DTensor_construct_from_3_components(self.name, src1, src2, src3)
583581
elif isinstance(src1, NiftiImageData3DDeformation):
584582
self.handle = src1.handle.cReg_NiftiImageData3DDisplacement_create_from_def()
585583
else:
@@ -600,12 +598,11 @@ def __init__(self, src1=None, src2=None, src3=None):
600598
self.handle = None
601599
self.name = 'NiftiImageData3DDeformation'
602600
if src1 is None:
603-
self.handle = Handle(None, -1).cReg_newObject(self.name)
601+
self.handle = HANDLE.cReg_newObject(self.name)
604602
elif isinstance(src1, str):
605-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src1)
603+
self.handle = HANDLE.cReg_objectFromFile(self.name, src1)
606604
elif all(isinstance(i, NiftiImageData3D) for i in (src1, src2, src3)):
607-
self.handle = Handle(None, -1).cReg_NiftiImageData3DTensor_construct_from_3_components(
608-
self.name, src1, src2, src3)
605+
self.handle = HANDLE.cReg_NiftiImageData3DTensor_construct_from_3_components(self.name, src1, src2, src3)
609606
elif isinstance(src1, NiftiImageData3DDisplacement):
610607
self.handle = src1.handle.cReg_NiftiImageData3DDeformation_create_from_disp()
611608
else:
@@ -790,7 +787,7 @@ def __init__(self):
790787
"""init."""
791788
super().__init__()
792789
self.name = 'NiftyAladinSym'
793-
self.handle = Handle(None, -1).cReg_newObject(self.name)
790+
self.handle = HANDLE.cReg_newObject(self.name)
794791

795792
def get_transformation_matrix_forward(self):
796793
"""Get forward transformation matrix."""
@@ -809,7 +806,7 @@ def print_all_wrapped_methods():
809806
"""Print all wrapped methods."""
810807
print("""In C++, this class is templated. \"dataType\"
811808
corresponds to \"float\" for Matlab and python.""")
812-
Handle(None, -1).cReg_NiftyRegistration_print_all_wrapped_methods('NiftyAladinSym')
809+
HANDLE.cReg_NiftyRegistration_print_all_wrapped_methods('NiftyAladinSym')
813810

814811

815812
class NiftyF3dSym(_NiftyRegistration):
@@ -818,7 +815,7 @@ def __init__(self):
818815
"""init."""
819816
super().__init__()
820817
self.name = 'NiftyF3dSym'
821-
self.handle = Handle(None, -1).cReg_newObject(self.name)
818+
self.handle = HANDLE.cReg_newObject(self.name)
822819

823820
def set_floating_time_point(self, floating_time_point):
824821
"""Set floating time point."""
@@ -839,7 +836,7 @@ def print_all_wrapped_methods():
839836
"""Print all wrapped methods."""
840837
print("""In C++, this class is templated. \"dataType\"
841838
corresponds to \"float\" for Matlab and python.""")
842-
Handle(None, -1).cReg_NiftyRegistration_print_all_wrapped_methods('NiftyF3dSym')
839+
HANDLE.cReg_NiftyRegistration_print_all_wrapped_methods('NiftyF3dSym')
843840

844841

845842
if SIRF_HAS_SPM:
@@ -850,7 +847,7 @@ def __init__(self):
850847
"""init."""
851848
super().__init__()
852849
self.name = 'SPMRegistration'
853-
self.handle = Handle(None, -1).cReg_newObject(self.name)
850+
self.handle = HANDLE.cReg_newObject(self.name)
854851

855852
def get_transformation_matrix_forward(self, idx=0):
856853
"""Get forward transformation matrix."""
@@ -890,7 +887,7 @@ class NiftyResampler:
890887
def __init__(self):
891888
"""init."""
892889
self.name = 'NiftyResampler'
893-
self.handle = Handle(None, -1).cReg_newObject(self.name)
890+
self.handle = HANDLE.cReg_newObject(self.name)
894891
self.reference_image = None
895892
self.floating_image = None
896893

@@ -1071,7 +1068,7 @@ class ImageWeightedMean:
10711068
def __init__(self):
10721069
"""init."""
10731070
self.name = 'ImageWeightedMean'
1074-
self.handle = Handle(None, -1).cReg_newObject(self.name)
1071+
self.handle = HANDLE.cReg_newObject(self.name)
10751072

10761073
def add_image(self, image, weight):
10771074
"""Add an image and its corresponding weight.
@@ -1104,9 +1101,9 @@ def __init__(self, src1=None, src2=None):
11041101
self.handle = None
11051102
self.name = 'AffineTransformation'
11061103
if src1 is None:
1107-
self.handle = Handle(None, -1).cReg_newObject(self.name)
1104+
self.handle = HANDLE.cReg_newObject(self.name)
11081105
elif isinstance(src1, str):
1109-
self.handle = Handle(None, -1).cReg_objectFromFile(self.name, src1)
1106+
self.handle = HANDLE.cReg_objectFromFile(self.name, src1)
11101107
elif isinstance(src1, numpy.ndarray) and src2 is None:
11111108
src1 = format_numpy_array_for_setter(src1)
11121109
if src1.shape != (4, 4):
@@ -1116,15 +1113,15 @@ def __init__(self, src1=None, src2=None):
11161113
for i in range(4):
11171114
for j in range(4):
11181115
trans[i, j] = src1[j, i]
1119-
self.handle = Handle(None, -1).cReg_AffineTransformation_construct_from_TM(trans)
1116+
self.handle = HANDLE.cReg_AffineTransformation_construct_from_TM(trans)
11201117
elif isinstance(src1, numpy.ndarray) and src2 is not None and \
11211118
isinstance(src2, Quaternion):
11221119
src1 = format_numpy_array_for_setter(src1)
11231120
self.handle = src2.handle.cReg_AffineTransformation_construct_from_trans_and_quaternion(src1)
11241121
elif isinstance(src1, numpy.ndarray) and isinstance(src2, numpy.ndarray):
11251122
src1 = format_numpy_array_for_setter(src1)
11261123
src2 = format_numpy_array_for_setter(src2)
1127-
self.handle = Handle(None, -1).cReg_AffineTransformation_construct_from_trans_and_euler(src1, src2)
1124+
self.handle = HANDLE.cReg_AffineTransformation_construct_from_trans_and_euler(src1, src2)
11281125
else:
11291126
raise error("""AffineTransformation accepts no args, filename,
11301127
4x4 array or translation with quaternion.""")
@@ -1191,7 +1188,7 @@ def get_quaternion(self):
11911188
def get_identity():
11921189
"""Get identity matrix."""
11931190
mat = AffineTransformation()
1194-
mat.handle = Handle(None, -1).cReg_AffineTransformation_get_identity()
1191+
mat.handle = HANDLE.cReg_AffineTransformation_get_identity()
11951192
return mat
11961193

11971194
@staticmethod
@@ -1221,7 +1218,7 @@ def __init__(self, src=None):
12211218
array is wrong size.""")
12221219
if src.dtype is not numpy.float32:
12231220
src = src.astype(numpy.float32)
1224-
self.handle = Handle(None, -1).cReg_Quaternion_construct_from_array(src)
1221+
self.handle = HANDLE.cReg_Quaternion_construct_from_array(src)
12251222
elif isinstance(src, AffineTransformation):
12261223
self.handle = src.handle.cReg_Quaternion_construct_from_AffineTransformation()
12271224
else:

src/common/SIRF.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@
3535

3636
import deprecation
3737

38-
from sirf.Utilities import Handle, assert_validities, assert_validity, cpp_int_dtype
38+
from sirf.Utilities import HANDLE, assert_validities, assert_validity, cpp_int_dtype
3939

4040
if sys.version_info[0] >= 3 and sys.version_info[1] >= 4:
4141
ABC = abc.ABC
@@ -679,7 +679,7 @@ def __ne__(self, other):
679679
return not (self == other)
680680

681681
def read(self, file, engine, verb):
682-
self.handle = Handle(None, -1).cSIRF_readImageData(file, engine, verb)
682+
self.handle = HANDLE.cSIRF_readImageData(file, engine, verb)
683683

684684
def fill(self, image):
685685
self.handle.cSIRF_fillImageFromImage(image)
@@ -713,7 +713,7 @@ class DataHandleVector:
713713
"""
714714
def __init__(self):
715715
self.name = 'DataHandleVector'
716-
self.handle = Handle(None, -1).cSIRF_newObject(self.name)
716+
self.handle = HANDLE.cSIRF_newObject(self.name)
717717

718718
def push_back(self, handle):
719719
"""Push back new data handle."""

src/common/Utilities.py

Lines changed: 17 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
'''Utilities used by all engines
22
'''
3+
import functools
34
import inspect
45
import os
56
import re
@@ -34,13 +35,19 @@
3435

3536

3637
class Handle:
38+
"""
39+
Convenience class for:
40+
- calling SWIG backend functions
41+
+ routing to correct backend based on function name prefix
42+
+ auto-unwrapping function arguments from:
43+
* numpy data arrays (via `.ctypes.data`)
44+
* handle pointers (via `._handle`)
45+
- casting return types
46+
- pointer destrunction on deletion
47+
"""
3748
def __init__(self, handle, check_stack: int | None = None):
3849
self._handle = handle
3950
if (check_stack is None or check_stack >= 0) and pyiutil.executionStatus(handle) != 0:
40-
check_stack = inspect.stack()[1 if check_stack is None else check_stack]
41-
#print('\nFile: %s' % check_stack[1])
42-
#print('Line: %d' % check_stack[2])
43-
#print('check_status found the following message sent from the engine:')
4451
msg = pyiutil.executionError(handle)
4552
file = pyiutil.executionErrorFile(handle)
4653
line = pyiutil.executionErrorLine(handle)
@@ -91,9 +98,11 @@ def __getattr__(self, name):
9198
func = getattr(backend, name)
9299
match name:
93100
case 'cSIRF_axpby' | 'cSIRF_axpbyAlt' | 'cReg_AffineTransformation_construct_from_trans_and_quaternion': # 2nd arg=self
101+
@functools.wraps(func)
94102
def wrapped(one, *args):
95103
return Handle(func(self._toarg(one), self._handle, *map(self._toarg, args)))
96104
case _:
105+
@functools.wraps(func)
97106
def wrapped(*args):
98107
if self.valid: # 1st arg=self
99108
return Handle(func(self._handle, *map(self._toarg, args)))
@@ -113,6 +122,9 @@ def _toarg(cls, obj):
113122
return obj
114123

115124

125+
HANDLE = Handle(None, -1) # for calling backend functions which do not require a handle argument
126+
127+
116128
def cpp_int_bits():
117129
"""Returns the number of bits in a C++ integer."""
118130
return pyiutil.intBits()
@@ -152,7 +164,7 @@ def examples_data_path(data_type):
152164
Returns the path to PET/MR/Registration data used by SIRF/examples demos.
153165
data_type: either 'PET' or 'MR' or 'Registration'
154166
'''
155-
path = Handle(None, -1).cSIRF_examples_data_path(data_type)
167+
path = HANDLE.cSIRF_examples_data_path(data_type)
156168
return str(path)
157169

158170

0 commit comments

Comments
 (0)