From 12620800c02f6e7585f4147ab43bfecc56d7a3df Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 3 Sep 2026 12:00:38 +0200 Subject: [PATCH] feat: modify define_function_handling_dict_from_class to inject function to module directly --- src/spikeinterface/core/core_tools.py | 5 +++-- src/spikeinterface/core/segmentutils.py | 2 +- src/spikeinterface/extractors/neoextractors/mcsraw.py | 2 +- src/spikeinterface/preprocessing/astype.py | 2 +- .../preprocessing/average_across_direction.py | 2 +- src/spikeinterface/preprocessing/clip.py | 6 ++---- src/spikeinterface/preprocessing/common_reference.py | 4 +--- src/spikeinterface/preprocessing/decimate.py | 2 +- .../deepinterpolation/deepinterpolation.py | 4 +--- src/spikeinterface/preprocessing/depth_order.py | 2 +- src/spikeinterface/preprocessing/detect_artifacts.py | 2 +- .../preprocessing/detect_bad_channels.py | 2 +- .../preprocessing/directional_derivative.py | 4 +--- src/spikeinterface/preprocessing/filter.py | 8 ++++---- src/spikeinterface/preprocessing/filter_gaussian.py | 2 +- .../preprocessing/highpass_spatial_filter.py | 4 +--- .../preprocessing/interpolate_bad_channels.py | 6 ++---- src/spikeinterface/preprocessing/normalize_scale.py | 10 ++++------ src/spikeinterface/preprocessing/phase_shift.py | 2 +- src/spikeinterface/preprocessing/rectify.py | 2 +- src/spikeinterface/preprocessing/remove_artifacts.py | 4 +--- src/spikeinterface/preprocessing/resample.py | 2 +- src/spikeinterface/preprocessing/scale.py | 2 +- src/spikeinterface/preprocessing/silence_periods.py | 4 +--- src/spikeinterface/preprocessing/unsigned_to_signed.py | 4 +--- src/spikeinterface/preprocessing/whiten.py | 2 +- src/spikeinterface/preprocessing/zero_channel_pad.py | 6 ++---- .../sortingcomponents/motion/motion_interpolation.py | 4 +--- 28 files changed, 39 insertions(+), 62 deletions(-) diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index ddc52e283f..4d18c8b914 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -56,7 +56,8 @@ def source_class_or_dict_of_sources_classes(*args, **kwargs): source_class_or_dict_of_sources_classes.__doc__ = source_class.__doc__ source_class_or_dict_of_sources_classes.__name__ = name - return source_class_or_dict_of_sources_classes + # This is a trick to make the function available in the global namespace of the caller module + sys._getframe(1).f_globals[name] = source_class_or_dict_of_sources_classes # Generic typing needed to help propagate typing @@ -67,7 +68,7 @@ def source_class_or_dict_of_sources_classes(*args, **kwargs): def define_function_from_class(source_class: Callable[P, T], name: str) -> Callable[P, T]: - "Wrapper to change the name of a class" + "Wrapper to inject source_class into the caller's module namespace under name." return source_class diff --git a/src/spikeinterface/core/segmentutils.py b/src/spikeinterface/core/segmentutils.py index 8bc50efc23..90218a55bb 100644 --- a/src/spikeinterface/core/segmentutils.py +++ b/src/spikeinterface/core/segmentutils.py @@ -83,7 +83,7 @@ def get_traces(self, *args, **kwargs): return self.parent_segment.get_traces(*args, **kwargs) -append_recordings = define_function_from_class(source_class=AppendSegmentRecording, name="append_segment_recording") +append_recordings = define_function_from_class(source_class=AppendSegmentRecording, name="append_recordings") class ConcatenateSegmentRecording(BaseRecording): diff --git a/src/spikeinterface/extractors/neoextractors/mcsraw.py b/src/spikeinterface/extractors/neoextractors/mcsraw.py index 1da68a2c10..06bf9e3324 100644 --- a/src/spikeinterface/extractors/neoextractors/mcsraw.py +++ b/src/spikeinterface/extractors/neoextractors/mcsraw.py @@ -60,4 +60,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_mcsraw = define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_maxwell_event") +read_mcsraw = define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_mcsraw") diff --git a/src/spikeinterface/preprocessing/astype.py b/src/spikeinterface/preprocessing/astype.py index 26d45dd711..8c7eaf488e 100644 --- a/src/spikeinterface/preprocessing/astype.py +++ b/src/spikeinterface/preprocessing/astype.py @@ -78,4 +78,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -astype = define_function_handling_dict_from_class(source_class=AstypeRecording, name="astype") +define_function_handling_dict_from_class(source_class=AstypeRecording, name="astype") diff --git a/src/spikeinterface/preprocessing/average_across_direction.py b/src/spikeinterface/preprocessing/average_across_direction.py index 23e0a1d5ae..43d46ea254 100644 --- a/src/spikeinterface/preprocessing/average_across_direction.py +++ b/src/spikeinterface/preprocessing/average_across_direction.py @@ -137,7 +137,7 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -average_across_direction = define_function_handling_dict_from_class( +define_function_handling_dict_from_class( source_class=AverageAcrossDirectionRecording, name="average_across_direction", ) diff --git a/src/spikeinterface/preprocessing/clip.py b/src/spikeinterface/preprocessing/clip.py index dd7676fd23..b0fb5f5e4e 100644 --- a/src/spikeinterface/preprocessing/clip.py +++ b/src/spikeinterface/preprocessing/clip.py @@ -167,7 +167,5 @@ def get_traces(self, start_frame, end_frame, channel_indices): return traces -clip = define_function_handling_dict_from_class(source_class=ClipRecording, name="clip") -blank_saturation = define_function_handling_dict_from_class( - source_class=BlankSaturationRecording, name="blank_saturation" -) +define_function_handling_dict_from_class(source_class=ClipRecording, name="clip") +define_function_handling_dict_from_class(source_class=BlankSaturationRecording, name="blank_saturation") diff --git a/src/spikeinterface/preprocessing/common_reference.py b/src/spikeinterface/preprocessing/common_reference.py index 95e5e22c4c..7c37c480da 100644 --- a/src/spikeinterface/preprocessing/common_reference.py +++ b/src/spikeinterface/preprocessing/common_reference.py @@ -325,6 +325,4 @@ def slice_groups(self, channel_indices): return zip(group_indices, selected_channels, group_channels) -common_reference = define_function_handling_dict_from_class( - source_class=CommonReferenceRecording, name="common_reference" -) +define_function_handling_dict_from_class(source_class=CommonReferenceRecording, name="common_reference") diff --git a/src/spikeinterface/preprocessing/decimate.py b/src/spikeinterface/preprocessing/decimate.py index 716cc46f70..8dce87a924 100644 --- a/src/spikeinterface/preprocessing/decimate.py +++ b/src/spikeinterface/preprocessing/decimate.py @@ -134,4 +134,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): ].astype(self._dtype) -decimate = define_function_handling_dict_from_class(source_class=DecimateRecording, name="decimate") +define_function_handling_dict_from_class(source_class=DecimateRecording, name="decimate") diff --git a/src/spikeinterface/preprocessing/deepinterpolation/deepinterpolation.py b/src/spikeinterface/preprocessing/deepinterpolation/deepinterpolation.py index e324048c5a..fe58750608 100644 --- a/src/spikeinterface/preprocessing/deepinterpolation/deepinterpolation.py +++ b/src/spikeinterface/preprocessing/deepinterpolation/deepinterpolation.py @@ -191,6 +191,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -deepinterpolate = define_function_handling_dict_from_class( - source_class=DeepInterpolatedRecording, name="deepinterpolate" -) +define_function_handling_dict_from_class(source_class=DeepInterpolatedRecording, name="deepinterpolate") diff --git a/src/spikeinterface/preprocessing/depth_order.py b/src/spikeinterface/preprocessing/depth_order.py index ca9caa9521..a0c07e1faf 100644 --- a/src/spikeinterface/preprocessing/depth_order.py +++ b/src/spikeinterface/preprocessing/depth_order.py @@ -41,4 +41,4 @@ def __init__(self, parent_recording, channel_ids=None, dimensions=("x", "y"), fl ) -depth_order = define_function_handling_dict_from_class(source_class=DepthOrderRecording, name="depth_order") +define_function_handling_dict_from_class(source_class=DepthOrderRecording, name="depth_order") diff --git a/src/spikeinterface/preprocessing/detect_artifacts.py b/src/spikeinterface/preprocessing/detect_artifacts.py index 95f286a85a..8ccbc29005 100644 --- a/src/spikeinterface/preprocessing/detect_artifacts.py +++ b/src/spikeinterface/preprocessing/detect_artifacts.py @@ -785,7 +785,7 @@ def __init__( # function for API -detect_and_remove_artifacts = define_function_handling_dict_from_class( +define_function_handling_dict_from_class( source_class=DetectAndRemoveArtifactsRecording, name="detect_and_remove_artifacts" ) diff --git a/src/spikeinterface/preprocessing/detect_bad_channels.py b/src/spikeinterface/preprocessing/detect_bad_channels.py index b255c04a9e..0907b35d52 100644 --- a/src/spikeinterface/preprocessing/detect_bad_channels.py +++ b/src/spikeinterface/preprocessing/detect_bad_channels.py @@ -137,7 +137,7 @@ def __init__( DetectAndRemoveBadChannelsRecording.__doc__ = DetectAndRemoveBadChannelsRecording.__doc__.format( _bad_channel_detection_kwargs_doc ) -detect_and_remove_bad_channels = define_function_handling_dict_from_class( +define_function_handling_dict_from_class( source_class=DetectAndRemoveBadChannelsRecording, name="detect_and_remove_bad_channels" ) diff --git a/src/spikeinterface/preprocessing/directional_derivative.py b/src/spikeinterface/preprocessing/directional_derivative.py index f302708055..61133a8bca 100644 --- a/src/spikeinterface/preprocessing/directional_derivative.py +++ b/src/spikeinterface/preprocessing/directional_derivative.py @@ -134,6 +134,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -directional_derivative = define_function_handling_dict_from_class( - source_class=DirectionalDerivativeRecording, name="directional_derivative" -) +define_function_handling_dict_from_class(source_class=DirectionalDerivativeRecording, name="directional_derivative") diff --git a/src/spikeinterface/preprocessing/filter.py b/src/spikeinterface/preprocessing/filter.py index 2cf3e47c93..c11520c08c 100644 --- a/src/spikeinterface/preprocessing/filter.py +++ b/src/spikeinterface/preprocessing/filter.py @@ -414,10 +414,10 @@ def __init__(self, recording, freq=3000, q=30, margin_ms="auto", dtype=None, **f # functions for API -filter = define_function_handling_dict_from_class(source_class=FilterRecording, name="filter") -bandpass_filter = define_function_handling_dict_from_class(source_class=BandpassFilterRecording, name="bandpass_filter") -notch_filter = define_function_handling_dict_from_class(source_class=NotchFilterRecording, name="notch_filter") -highpass_filter = define_function_handling_dict_from_class(source_class=HighpassFilterRecording, name="highpass_filter") +define_function_handling_dict_from_class(source_class=FilterRecording, name="filter") +define_function_handling_dict_from_class(source_class=BandpassFilterRecording, name="bandpass_filter") +define_function_handling_dict_from_class(source_class=NotchFilterRecording, name="notch_filter") +define_function_handling_dict_from_class(source_class=HighpassFilterRecording, name="highpass_filter") def causal_filter( diff --git a/src/spikeinterface/preprocessing/filter_gaussian.py b/src/spikeinterface/preprocessing/filter_gaussian.py index 1cf6873a7a..d19927f27a 100644 --- a/src/spikeinterface/preprocessing/filter_gaussian.py +++ b/src/spikeinterface/preprocessing/filter_gaussian.py @@ -143,4 +143,4 @@ def _create_gaussian(self, N: int, cutoff_f: float): return gaussian -gaussian_filter = define_function_handling_dict_from_class(source_class=GaussianFilterRecording, name="gaussian_filter") +define_function_handling_dict_from_class(source_class=GaussianFilterRecording, name="gaussian_filter") diff --git a/src/spikeinterface/preprocessing/highpass_spatial_filter.py b/src/spikeinterface/preprocessing/highpass_spatial_filter.py index 979bb00908..01a767e0bc 100644 --- a/src/spikeinterface/preprocessing/highpass_spatial_filter.py +++ b/src/spikeinterface/preprocessing/highpass_spatial_filter.py @@ -272,9 +272,7 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -highpass_spatial_filter = define_function_handling_dict_from_class( - source_class=HighpassSpatialFilterRecording, name="highpass_spatial_filter" -) +define_function_handling_dict_from_class(source_class=HighpassSpatialFilterRecording, name="highpass_spatial_filter") # ----------------------------------------------------------------------------------------------- diff --git a/src/spikeinterface/preprocessing/interpolate_bad_channels.py b/src/spikeinterface/preprocessing/interpolate_bad_channels.py index f2ae77f3d7..1e3aa5b833 100644 --- a/src/spikeinterface/preprocessing/interpolate_bad_channels.py +++ b/src/spikeinterface/preprocessing/interpolate_bad_channels.py @@ -134,7 +134,7 @@ def __init__( DetectAndInterpolateBadChannelsRecording.__doc__ = DetectAndInterpolateBadChannelsRecording.__doc__.format( _bad_channel_detection_kwargs_doc ) -detect_and_interpolate_bad_channels = define_function_handling_dict_from_class( +define_function_handling_dict_from_class( source_class=DetectAndInterpolateBadChannelsRecording, name="detect_and_interpolate_bad_channels" ) @@ -170,6 +170,4 @@ def estimate_recommended_sigma_um(recording): return mode(np.diff(np.unique(y_sorted)), keepdims=False)[0] -interpolate_bad_channels = define_function_handling_dict_from_class( - source_class=InterpolateBadChannelsRecording, name="interpolate_bad_channels" -) +define_function_handling_dict_from_class(source_class=InterpolateBadChannelsRecording, name="interpolate_bad_channels") diff --git a/src/spikeinterface/preprocessing/normalize_scale.py b/src/spikeinterface/preprocessing/normalize_scale.py index ee6d6dc20c..583dd936e0 100644 --- a/src/spikeinterface/preprocessing/normalize_scale.py +++ b/src/spikeinterface/preprocessing/normalize_scale.py @@ -323,9 +323,7 @@ def __init__( # functions for API -normalize_by_quantile = define_function_handling_dict_from_class( - source_class=NormalizeByQuantileRecording, name="normalize_by_quantile" -) -scale = define_function_handling_dict_from_class(source_class=ScaleRecording, name="scale") -center = define_function_handling_dict_from_class(source_class=CenterRecording, name="center") -zscore = define_function_handling_dict_from_class(source_class=ZScoreRecording, name="zscore") +define_function_handling_dict_from_class(source_class=NormalizeByQuantileRecording, name="normalize_by_quantile") +define_function_handling_dict_from_class(source_class=ScaleRecording, name="scale") +define_function_handling_dict_from_class(source_class=CenterRecording, name="center") +define_function_handling_dict_from_class(source_class=ZScoreRecording, name="zscore") diff --git a/src/spikeinterface/preprocessing/phase_shift.py b/src/spikeinterface/preprocessing/phase_shift.py index 3000d0c143..1568a53d14 100644 --- a/src/spikeinterface/preprocessing/phase_shift.py +++ b/src/spikeinterface/preprocessing/phase_shift.py @@ -106,7 +106,7 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -phase_shift = define_function_handling_dict_from_class(source_class=PhaseShiftRecording, name="phase_shift") +define_function_handling_dict_from_class(source_class=PhaseShiftRecording, name="phase_shift") def apply_frequency_shift(signal, shift_samples, axis=0): diff --git a/src/spikeinterface/preprocessing/rectify.py b/src/spikeinterface/preprocessing/rectify.py index 7bd91a16d9..e6d6fd1059 100644 --- a/src/spikeinterface/preprocessing/rectify.py +++ b/src/spikeinterface/preprocessing/rectify.py @@ -25,4 +25,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -rectify = define_function_handling_dict_from_class(source_class=RectifyRecording, name="rectify") +define_function_handling_dict_from_class(source_class=RectifyRecording, name="rectify") diff --git a/src/spikeinterface/preprocessing/remove_artifacts.py b/src/spikeinterface/preprocessing/remove_artifacts.py index b20f1351b0..bf31096f14 100644 --- a/src/spikeinterface/preprocessing/remove_artifacts.py +++ b/src/spikeinterface/preprocessing/remove_artifacts.py @@ -449,6 +449,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -remove_artifacts = define_function_handling_dict_from_class( - source_class=RemoveArtifactsRecording, name="remove_artifacts" -) +define_function_handling_dict_from_class(source_class=RemoveArtifactsRecording, name="remove_artifacts") diff --git a/src/spikeinterface/preprocessing/resample.py b/src/spikeinterface/preprocessing/resample.py index a801c45eff..07cfa37532 100644 --- a/src/spikeinterface/preprocessing/resample.py +++ b/src/spikeinterface/preprocessing/resample.py @@ -389,7 +389,7 @@ def _get_traces_gapped(self, start_frame, end_frame, channel_indices): return result[:pos] -resample = define_function_handling_dict_from_class(source_class=ResampleRecording, name="resample") +define_function_handling_dict_from_class(source_class=ResampleRecording, name="resample") # Some helpers to do checks diff --git a/src/spikeinterface/preprocessing/scale.py b/src/spikeinterface/preprocessing/scale.py index 701987facb..3412666d0f 100644 --- a/src/spikeinterface/preprocessing/scale.py +++ b/src/spikeinterface/preprocessing/scale.py @@ -62,7 +62,7 @@ def __init__(self, recording): self.set_channel_offsets(offsets=0.0) -scale_to_physical_units = define_function_handling_dict_from_class(ScaleToPhysicalUnits, name="scale_to_physical_units") +define_function_handling_dict_from_class(ScaleToPhysicalUnits, name="scale_to_physical_units") def scale_to_uV(recording: BasePreprocessor) -> BasePreprocessor: diff --git a/src/spikeinterface/preprocessing/silence_periods.py b/src/spikeinterface/preprocessing/silence_periods.py index 430348bbb3..cdfdc28033 100644 --- a/src/spikeinterface/preprocessing/silence_periods.py +++ b/src/spikeinterface/preprocessing/silence_periods.py @@ -247,6 +247,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -silence_periods = define_function_handling_dict_from_class( - source_class=SilencedPeriodsRecording, name="silence_periods" -) +define_function_handling_dict_from_class(source_class=SilencedPeriodsRecording, name="silence_periods") diff --git a/src/spikeinterface/preprocessing/unsigned_to_signed.py b/src/spikeinterface/preprocessing/unsigned_to_signed.py index ae1ce12281..eb57ee3cff 100644 --- a/src/spikeinterface/preprocessing/unsigned_to_signed.py +++ b/src/spikeinterface/preprocessing/unsigned_to_signed.py @@ -65,6 +65,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -unsigned_to_signed = define_function_handling_dict_from_class( - source_class=UnsignedToSignedRecording, name="unsigned_to_signed" -) +define_function_handling_dict_from_class(source_class=UnsignedToSignedRecording, name="unsigned_to_signed") diff --git a/src/spikeinterface/preprocessing/whiten.py b/src/spikeinterface/preprocessing/whiten.py index d3fd699442..503949a941 100644 --- a/src/spikeinterface/preprocessing/whiten.py +++ b/src/spikeinterface/preprocessing/whiten.py @@ -288,4 +288,4 @@ def compute_sklearn_covariance_matrix(data, regularize_kwargs): # function for API -whiten = define_function_handling_dict_from_class(source_class=WhitenRecording, name="whiten") +define_function_handling_dict_from_class(source_class=WhitenRecording, name="whiten") diff --git a/src/spikeinterface/preprocessing/zero_channel_pad.py b/src/spikeinterface/preprocessing/zero_channel_pad.py index 639e5ccced..3da622440f 100644 --- a/src/spikeinterface/preprocessing/zero_channel_pad.py +++ b/src/spikeinterface/preprocessing/zero_channel_pad.py @@ -193,7 +193,5 @@ def get_traces(self, start_frame, end_frame, channel_indices): # function for API -zero_channel_pad = define_function_handling_dict_from_class( - source_class=ZeroChannelPaddedRecording, name="zero_channel_pad" -) -pad_traces = define_function_handling_dict_from_class(source_class=TracePaddedRecording, name="pad_traces") +define_function_handling_dict_from_class(source_class=ZeroChannelPaddedRecording, name="zero_channel_pad") +define_function_handling_dict_from_class(source_class=TracePaddedRecording, name="pad_traces") diff --git a/src/spikeinterface/sortingcomponents/motion/motion_interpolation.py b/src/spikeinterface/sortingcomponents/motion/motion_interpolation.py index a144940853..be3a870911 100644 --- a/src/spikeinterface/sortingcomponents/motion/motion_interpolation.py +++ b/src/spikeinterface/sortingcomponents/motion/motion_interpolation.py @@ -523,6 +523,4 @@ def get_traces(self, start_frame, end_frame, channel_indices): return traces -interpolate_motion = define_function_handling_dict_from_class( - source_class=InterpolateMotionRecording, name="interpolate_motion" -) +define_function_handling_dict_from_class(source_class=InterpolateMotionRecording, name="interpolate_motion")