Skip to content

Commit 1262080

Browse files
committed
feat: modify define_function_handling_dict_from_class to inject function to module directly
1 parent 32bcb0d commit 1262080

28 files changed

Lines changed: 39 additions & 62 deletions

src/spikeinterface/core/core_tools.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,8 @@ def source_class_or_dict_of_sources_classes(*args, **kwargs):
5656
source_class_or_dict_of_sources_classes.__doc__ = source_class.__doc__
5757
source_class_or_dict_of_sources_classes.__name__ = name
5858

59-
return source_class_or_dict_of_sources_classes
59+
# This is a trick to make the function available in the global namespace of the caller module
60+
sys._getframe(1).f_globals[name] = source_class_or_dict_of_sources_classes
6061

6162

6263
# Generic typing needed to help propagate typing
@@ -67,7 +68,7 @@ def source_class_or_dict_of_sources_classes(*args, **kwargs):
6768

6869

6970
def define_function_from_class(source_class: Callable[P, T], name: str) -> Callable[P, T]:
70-
"Wrapper to change the name of a class"
71+
"Wrapper to inject source_class into the caller's module namespace under name."
7172

7273
return source_class
7374

src/spikeinterface/core/segmentutils.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,7 @@ def get_traces(self, *args, **kwargs):
8383
return self.parent_segment.get_traces(*args, **kwargs)
8484

8585

86-
append_recordings = define_function_from_class(source_class=AppendSegmentRecording, name="append_segment_recording")
86+
append_recordings = define_function_from_class(source_class=AppendSegmentRecording, name="append_recordings")
8787

8888

8989
class ConcatenateSegmentRecording(BaseRecording):

src/spikeinterface/extractors/neoextractors/mcsraw.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,4 +60,4 @@ def map_to_neo_kwargs(cls, file_path):
6060
return neo_kwargs
6161

6262

63-
read_mcsraw = define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_maxwell_event")
63+
read_mcsraw = define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_mcsraw")

src/spikeinterface/preprocessing/astype.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -78,4 +78,4 @@ def get_traces(self, start_frame, end_frame, channel_indices):
7878

7979

8080
# function for API
81-
astype = define_function_handling_dict_from_class(source_class=AstypeRecording, name="astype")
81+
define_function_handling_dict_from_class(source_class=AstypeRecording, name="astype")

src/spikeinterface/preprocessing/average_across_direction.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -137,7 +137,7 @@ def get_traces(self, start_frame, end_frame, channel_indices):
137137

138138

139139
# function for API
140-
average_across_direction = define_function_handling_dict_from_class(
140+
define_function_handling_dict_from_class(
141141
source_class=AverageAcrossDirectionRecording,
142142
name="average_across_direction",
143143
)

src/spikeinterface/preprocessing/clip.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -167,7 +167,5 @@ def get_traces(self, start_frame, end_frame, channel_indices):
167167
return traces
168168

169169

170-
clip = define_function_handling_dict_from_class(source_class=ClipRecording, name="clip")
171-
blank_saturation = define_function_handling_dict_from_class(
172-
source_class=BlankSaturationRecording, name="blank_saturation"
173-
)
170+
define_function_handling_dict_from_class(source_class=ClipRecording, name="clip")
171+
define_function_handling_dict_from_class(source_class=BlankSaturationRecording, name="blank_saturation")

src/spikeinterface/preprocessing/common_reference.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -325,6 +325,4 @@ def slice_groups(self, channel_indices):
325325
return zip(group_indices, selected_channels, group_channels)
326326

327327

328-
common_reference = define_function_handling_dict_from_class(
329-
source_class=CommonReferenceRecording, name="common_reference"
330-
)
328+
define_function_handling_dict_from_class(source_class=CommonReferenceRecording, name="common_reference")

src/spikeinterface/preprocessing/decimate.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -134,4 +134,4 @@ def get_traces(self, start_frame, end_frame, channel_indices):
134134
].astype(self._dtype)
135135

136136

137-
decimate = define_function_handling_dict_from_class(source_class=DecimateRecording, name="decimate")
137+
define_function_handling_dict_from_class(source_class=DecimateRecording, name="decimate")

src/spikeinterface/preprocessing/deepinterpolation/deepinterpolation.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -191,6 +191,4 @@ def get_traces(self, start_frame, end_frame, channel_indices):
191191

192192

193193
# function for API
194-
deepinterpolate = define_function_handling_dict_from_class(
195-
source_class=DeepInterpolatedRecording, name="deepinterpolate"
196-
)
194+
define_function_handling_dict_from_class(source_class=DeepInterpolatedRecording, name="deepinterpolate")

src/spikeinterface/preprocessing/depth_order.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,4 +41,4 @@ def __init__(self, parent_recording, channel_ids=None, dimensions=("x", "y"), fl
4141
)
4242

4343

44-
depth_order = define_function_handling_dict_from_class(source_class=DepthOrderRecording, name="depth_order")
44+
define_function_handling_dict_from_class(source_class=DepthOrderRecording, name="depth_order")

0 commit comments

Comments
 (0)