From 12620800c02f6e7585f4147ab43bfecc56d7a3df Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 3 Sep 2026 12:00:38 +0200 Subject: [PATCH 1/2] 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") From 7999ad9f206ffa9167ee2707dd9ff82a8426b29c Mon Sep 17 00:00:00 2001 From: Alessio Buccino Date: Thu, 3 Sep 2026 12:37:19 +0200 Subject: [PATCH 2/2] refactor: define_function_rom_class injects to module directly --- .../comparison/multicomparisons.py | 8 ++------ .../comparison/paircomparisons.py | 8 +++----- src/spikeinterface/core/binaryfolder.py | 2 +- .../core/binaryrecordingextractor.py | 2 +- src/spikeinterface/core/core_tools.py | 11 ++--------- src/spikeinterface/core/generate.py | 2 +- src/spikeinterface/core/npyfoldersnippets.py | 2 +- .../core/npysnippetsextractor.py | 4 ++-- src/spikeinterface/core/npzsortingextractor.py | 2 +- src/spikeinterface/core/segmentutils.py | 18 +++++++----------- src/spikeinterface/core/sortingfolder.py | 6 ++---- .../core/unitsaggregationsorting.py | 2 +- src/spikeinterface/core/zarrextractors.py | 4 ++-- src/spikeinterface/curation/curationsorting.py | 2 +- .../curation/mergeunitssorting.py | 2 +- .../curation/remove_duplicated_spikes.py | 4 +--- .../curation/splitunitsorting.py | 2 +- .../extractors/alfsortingextractor.py | 2 +- src/spikeinterface/extractors/cbin_ibl.py | 2 +- .../extractors/cellexplorersortingextractor.py | 2 +- .../extractors/combinatoextractors.py | 2 +- .../extractors/hdsortextractors.py | 2 +- .../extractors/herdingspikesextractors.py | 2 +- src/spikeinterface/extractors/iblextractors.py | 4 ++-- .../extractors/klustaextractors.py | 2 +- .../extractors/mclustextractors.py | 2 +- .../extractors/mcsh5extractors.py | 2 +- src/spikeinterface/extractors/mdaextractors.py | 4 ++-- .../extractors/neoextractors/alphaomega.py | 4 ++-- .../extractors/neoextractors/axon.py | 2 +- .../extractors/neoextractors/axona.py | 2 +- .../extractors/neoextractors/biocam.py | 2 +- .../extractors/neoextractors/blackrock.py | 6 ++---- .../extractors/neoextractors/ced.py | 2 +- .../extractors/neoextractors/edf.py | 2 +- .../extractors/neoextractors/intan.py | 6 ++---- .../extractors/neoextractors/maxwell.py | 4 ++-- .../extractors/neoextractors/mcsraw.py | 2 +- .../extractors/neoextractors/neuralynx.py | 6 ++---- .../extractors/neoextractors/neuroexplorer.py | 2 +- .../extractors/neoextractors/neuronexus.py | 2 +- .../extractors/neoextractors/neuroscope.py | 8 ++------ .../extractors/neoextractors/nix.py | 2 +- .../extractors/neoextractors/plexon.py | 4 ++-- .../extractors/neoextractors/plexon2.py | 6 +++--- .../extractors/neoextractors/spike2.py | 2 +- .../extractors/neoextractors/spikegadgets.py | 2 +- .../extractors/neoextractors/spikeglx.py | 2 +- .../extractors/neoextractors/tdt.py | 2 +- src/spikeinterface/extractors/nwbextractors.py | 6 +++--- .../extractors/phykilosortextractors.py | 4 ++-- .../extractors/shybridextractors.py | 6 ++---- .../extractors/sinapsrecordingextractors.py | 6 ++---- .../extractors/spykingcircusextractors.py | 2 +- .../extractors/tridesclousextractors.py | 2 +- .../extractors/waveclussnippetstextractors.py | 4 +--- .../extractors/waveclustextractors.py | 2 +- .../whitematterrecordingextractor.py | 2 +- .../extractors/xclustextractors.py | 2 +- .../extractors/yassextractors.py | 2 +- src/spikeinterface/generation/noise_tools.py | 4 +--- .../postprocessing/alignsorting.py | 2 +- 62 files changed, 91 insertions(+), 130 deletions(-) diff --git a/src/spikeinterface/comparison/multicomparisons.py b/src/spikeinterface/comparison/multicomparisons.py index a0feda4a52..1878f3f1fe 100644 --- a/src/spikeinterface/comparison/multicomparisons.py +++ b/src/spikeinterface/comparison/multicomparisons.py @@ -244,9 +244,7 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return spiketrain -compare_multiple_sorters = define_function_from_class( - source_class=MultiSortingComparison, name="compare_multiple_sorters" -) +define_function_from_class(source_class=MultiSortingComparison, name="compare_multiple_sorters") class MultiTemplateComparison(BaseMultiComparison, MixinTemplateComparison): @@ -331,6 +329,4 @@ def _populate_nodes(self): self.graph.add_node(node) -compare_multiple_templates = define_function_from_class( - source_class=MultiTemplateComparison, name="compare_multiple_templates" -) +define_function_from_class(source_class=MultiTemplateComparison, name="compare_multiple_templates") diff --git a/src/spikeinterface/comparison/paircomparisons.py b/src/spikeinterface/comparison/paircomparisons.py index 47e85a03de..45591d9e41 100644 --- a/src/spikeinterface/comparison/paircomparisons.py +++ b/src/spikeinterface/comparison/paircomparisons.py @@ -207,7 +207,7 @@ def get_agreement_fraction(self, unit1=None, unit2=None): return self.agreement_scores.at[unit1, unit2] -compare_two_sorters = define_function_from_class(source_class=SymmetricSortingComparison, name="compare_two_sorters") +define_function_from_class(source_class=SymmetricSortingComparison, name="compare_two_sorters") class GroundTruthComparison(BasePairSorterComparison): @@ -744,9 +744,7 @@ def count_units_categories( """ -compare_sorter_to_ground_truth = define_function_from_class( - source_class=GroundTruthComparison, name="compare_sorter_to_ground_truth" -) +define_function_from_class(source_class=GroundTruthComparison, name="compare_sorter_to_ground_truth") class TemplateComparison(BasePairComparison, MixinTemplateComparison): @@ -859,4 +857,4 @@ def _do_agreement(self): self.agreement_scores = pd.DataFrame(agreement_scores, index=self.unit_ids[0], columns=self.unit_ids[1]) -compare_templates = define_function_from_class(source_class=TemplateComparison, name="compare_templates") +define_function_from_class(source_class=TemplateComparison, name="compare_templates") diff --git a/src/spikeinterface/core/binaryfolder.py b/src/spikeinterface/core/binaryfolder.py index ee86c4dc7c..3f08757f2e 100644 --- a/src/spikeinterface/core/binaryfolder.py +++ b/src/spikeinterface/core/binaryfolder.py @@ -113,4 +113,4 @@ def get_binary_description(self): return d -read_binary_folder = define_function_from_class(source_class=BinaryFolderRecording, name="read_binary_folder") +define_function_from_class(source_class=BinaryFolderRecording, name="read_binary_folder") diff --git a/src/spikeinterface/core/binaryrecordingextractor.py b/src/spikeinterface/core/binaryrecordingextractor.py index 059526a028..bc09aad29a 100644 --- a/src/spikeinterface/core/binaryrecordingextractor.py +++ b/src/spikeinterface/core/binaryrecordingextractor.py @@ -244,4 +244,4 @@ def __del__(self): # For backward compatibility (old good time) BinDatRecordingExtractor = BinaryRecordingExtractor -read_binary = define_function_from_class(source_class=BinaryRecordingExtractor, name="read_binary") +define_function_from_class(source_class=BinaryRecordingExtractor, name="read_binary") diff --git a/src/spikeinterface/core/core_tools.py b/src/spikeinterface/core/core_tools.py index 4d18c8b914..4c778b0569 100644 --- a/src/spikeinterface/core/core_tools.py +++ b/src/spikeinterface/core/core_tools.py @@ -60,17 +60,10 @@ def source_class_or_dict_of_sources_classes(*args, **kwargs): sys._getframe(1).f_globals[name] = source_class_or_dict_of_sources_classes -# Generic typing needed to help propagate typing -# across multiple language servers -# see https://github.com/SpikeInterface/spikeinterface/issues/4319 -P = ParamSpec("P") -T = TypeVar("T") - - -def define_function_from_class(source_class: Callable[P, T], name: str) -> Callable[P, T]: +def define_function_from_class(source_class, name: str) -> None: "Wrapper to inject source_class into the caller's module namespace under name." - return source_class + sys._getframe(1).f_globals[name] = source_class def read_python(path): diff --git a/src/spikeinterface/core/generate.py b/src/spikeinterface/core/generate.py index 2c0ca2cfd5..3cd9fc3000 100644 --- a/src/spikeinterface/core/generate.py +++ b/src/spikeinterface/core/generate.py @@ -2145,7 +2145,7 @@ def get_num_samples(self) -> int: return self.num_samples -inject_templates = define_function_from_class(source_class=InjectTemplatesRecording, name="inject_templates") +define_function_from_class(source_class=InjectTemplatesRecording, name="inject_templates") ## toy example zone ## diff --git a/src/spikeinterface/core/npyfoldersnippets.py b/src/spikeinterface/core/npyfoldersnippets.py index 1e465d827a..2ae360e1ee 100644 --- a/src/spikeinterface/core/npyfoldersnippets.py +++ b/src/spikeinterface/core/npyfoldersnippets.py @@ -48,4 +48,4 @@ def __init__(self, folder_path): self._bin_kwargs = d["kwargs"] -read_npy_snippets_folder = define_function_from_class(source_class=NpyFolderSnippets, name="read_npy_snippets_folder") +define_function_from_class(source_class=NpyFolderSnippets, name="read_npy_snippets_folder") diff --git a/src/spikeinterface/core/npysnippetsextractor.py b/src/spikeinterface/core/npysnippetsextractor.py index 88c2ea9de3..246e6b7a2b 100644 --- a/src/spikeinterface/core/npysnippetsextractor.py +++ b/src/spikeinterface/core/npysnippetsextractor.py @@ -160,7 +160,7 @@ def get_frames(self, indices=None): """ if indices is None: return self._spikestimes - raise self._spikestimes[indices] + return self._spikestimes[indices] -read_npy_snippets = define_function_from_class(source_class=NpySnippetsExtractor, name="read_npy_snippets") +define_function_from_class(source_class=NpySnippetsExtractor, name="read_npy_snippets") diff --git a/src/spikeinterface/core/npzsortingextractor.py b/src/spikeinterface/core/npzsortingextractor.py index af608d3fb7..4b11135e62 100644 --- a/src/spikeinterface/core/npzsortingextractor.py +++ b/src/spikeinterface/core/npzsortingextractor.py @@ -81,4 +81,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return spike_times.astype("int64") -read_npz_sorting = define_function_from_class(source_class=NpzSortingExtractor, name="read_npz_sorting") +define_function_from_class(source_class=NpzSortingExtractor, name="read_npz_sorting") diff --git a/src/spikeinterface/core/segmentutils.py b/src/spikeinterface/core/segmentutils.py index 90218a55bb..3a104f9416 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_recordings") +define_function_from_class(source_class=AppendSegmentRecording, name="append_recordings") class ConcatenateSegmentRecording(BaseRecording): @@ -211,9 +211,7 @@ def get_traces(self, start_frame, end_frame, channel_indices): return traces -concatenate_recordings = define_function_from_class( - source_class=ConcatenateSegmentRecording, name="concatenate_recordings" -) +define_function_from_class(source_class=ConcatenateSegmentRecording, name="concatenate_recordings") class SelectSegmentRecording(BaseRecording): @@ -269,9 +267,7 @@ def split_recording(recording: BaseRecording): return recording_list -select_segment_recording = define_function_from_class( - source_class=SelectSegmentRecording, name="select_segment_recording" -) +define_function_from_class(source_class=SelectSegmentRecording, name="select_segment_recording") class AppendSegmentSorting(BaseSorting): @@ -319,7 +315,7 @@ def get_unit_spike_train(self, *args, **kwargs): return self.parent_segment.get_unit_spike_train(*args, **kwargs) -append_sortings = define_function_from_class(source_class=AppendSegmentSorting, name="append_sortings") +define_function_from_class(source_class=AppendSegmentSorting, name="append_sortings") class ConcatenateSegmentSorting(BaseSorting): @@ -511,7 +507,7 @@ def get_unit_spike_train( return spike_frames -concatenate_sortings = define_function_from_class(source_class=ConcatenateSegmentSorting, name="concatenate_sortings") +define_function_from_class(source_class=ConcatenateSegmentSorting, name="concatenate_sortings") class SplitSegmentSorting(BaseSorting): @@ -570,7 +566,7 @@ def __init__(self, parent_sorting: BaseSorting, recording_or_recording_list=None self._kwargs = {"parent_sorting": parent_sorting, "recording_or_recording_list": recording_list} -split_sorting = define_function_from_class(source_class=SplitSegmentSorting, name="split_sorting") +define_function_from_class(source_class=SplitSegmentSorting, name="split_sorting") class SelectSegmentSorting(BaseSorting): @@ -604,7 +600,7 @@ def __init__(self, sorting: BaseSorting, segment_indices: int | list[int]): self._kwargs = {"sorting": sorting, "segment_indices": [int(s) for s in segment_indices]} -select_segment_sorting = define_function_from_class(source_class=SelectSegmentSorting, name="select_segment_sorting") +define_function_from_class(source_class=SelectSegmentSorting, name="select_segment_sorting") class SelectSegmentEvent(BaseEvent): diff --git a/src/spikeinterface/core/sortingfolder.py b/src/spikeinterface/core/sortingfolder.py index ad331cf835..377b9b480a 100644 --- a/src/spikeinterface/core/sortingfolder.py +++ b/src/spikeinterface/core/sortingfolder.py @@ -126,7 +126,5 @@ def write_sorting(sorting, save_path): cached.dump(save_path / "npz.json", relative_to=save_path) -read_numpy_sorting_folder = define_function_from_class( - source_class=NumpyFolderSorting, name="read_numpy_sorting_folder" -) -read_npz_folder = define_function_from_class(source_class=NpzFolderSorting, name="read_npz_folder") +define_function_from_class(source_class=NumpyFolderSorting, name="read_numpy_sorting_folder") +define_function_from_class(source_class=NpzFolderSorting, name="read_npz_folder") diff --git a/src/spikeinterface/core/unitsaggregationsorting.py b/src/spikeinterface/core/unitsaggregationsorting.py index 3f650cd802..d25c2c110e 100644 --- a/src/spikeinterface/core/unitsaggregationsorting.py +++ b/src/spikeinterface/core/unitsaggregationsorting.py @@ -171,4 +171,4 @@ def get_unit_spike_train( return times -aggregate_units = define_function_from_class(UnitsAggregationSorting, "aggregate_units") +define_function_from_class(UnitsAggregationSorting, "aggregate_units") diff --git a/src/spikeinterface/core/zarrextractors.py b/src/spikeinterface/core/zarrextractors.py index 02ae5e9ecd..88807b4a9a 100644 --- a/src/spikeinterface/core/zarrextractors.py +++ b/src/spikeinterface/core/zarrextractors.py @@ -483,8 +483,8 @@ def write_sorting(sorting: BaseSorting, folder_path: str | Path, storage_options add_sorting_to_zarr_group(sorting, zarr_root, **kwargs) -read_zarr_recording = define_function_from_class(source_class=ZarrRecordingExtractor, name="read_zarr_recording") -read_zarr_sorting = define_function_from_class(source_class=ZarrSortingExtractor, name="read_zarr_sorting") +define_function_from_class(source_class=ZarrRecordingExtractor, name="read_zarr_recording") +define_function_from_class(source_class=ZarrSortingExtractor, name="read_zarr_sorting") def read_zarr( diff --git a/src/spikeinterface/curation/curationsorting.py b/src/spikeinterface/curation/curationsorting.py index 1e0e8fc1d0..498dfefcaa 100644 --- a/src/spikeinterface/curation/curationsorting.py +++ b/src/spikeinterface/curation/curationsorting.py @@ -291,4 +291,4 @@ def _add_new_stage(self, new_sorting, edges): self._sorting_stages_i += 1 -curation_sorting = define_function_from_class(source_class=CurationSorting, name="curation_sorting") +define_function_from_class(source_class=CurationSorting, name="curation_sorting") diff --git a/src/spikeinterface/curation/mergeunitssorting.py b/src/spikeinterface/curation/mergeunitssorting.py index 314a1a417e..16287f3646 100644 --- a/src/spikeinterface/curation/mergeunitssorting.py +++ b/src/spikeinterface/curation/mergeunitssorting.py @@ -128,7 +128,7 @@ def __init__(self, sorting, units_to_merge, new_unit_ids=None, properties_policy ) -merge_units_sorting = define_function_from_class(source_class=MergeUnitsSorting, name="merge_units_sorting") +define_function_from_class(source_class=MergeUnitsSorting, name="merge_units_sorting") class MergeUnitsSortingSegment(BaseSortingSegment): diff --git a/src/spikeinterface/curation/remove_duplicated_spikes.py b/src/spikeinterface/curation/remove_duplicated_spikes.py index f8716ba54a..195ca729e4 100644 --- a/src/spikeinterface/curation/remove_duplicated_spikes.py +++ b/src/spikeinterface/curation/remove_duplicated_spikes.py @@ -91,6 +91,4 @@ def get_unit_spike_train(self, unit_id, start_frame: int | None = None, end_fram return spike_train[start:end] -remove_duplicated_spikes = define_function_from_class( - source_class=RemoveDuplicatedSpikesSorting, name="remove_duplicated_spikes" -) +define_function_from_class(source_class=RemoveDuplicatedSpikesSorting, name="remove_duplicated_spikes") diff --git a/src/spikeinterface/curation/splitunitsorting.py b/src/spikeinterface/curation/splitunitsorting.py index 087a808b02..3b5ecc1de4 100644 --- a/src/spikeinterface/curation/splitunitsorting.py +++ b/src/spikeinterface/curation/splitunitsorting.py @@ -117,7 +117,7 @@ def __init__(self, sorting, split_unit_id, indices_list, new_unit_ids=None, prop ) -split_unit_sorting = define_function_from_class(source_class=SplitUnitSorting, name="split_unit_sorting") +define_function_from_class(source_class=SplitUnitSorting, name="split_unit_sorting") class SplitSortingUnitSegment(BaseSortingSegment): diff --git a/src/spikeinterface/extractors/alfsortingextractor.py b/src/spikeinterface/extractors/alfsortingextractor.py index ba08bca2e8..11fd291721 100644 --- a/src/spikeinterface/extractors/alfsortingextractor.py +++ b/src/spikeinterface/extractors/alfsortingextractor.py @@ -65,4 +65,4 @@ def get_unit_spike_train( return spike_frames[(spike_frames >= start_frame) & (spike_frames < end_frame)].astype("int64", copy=False) -read_alf_sorting = define_function_from_class(source_class=ALFSortingExtractor, name="read_alf_sorting") +define_function_from_class(source_class=ALFSortingExtractor, name="read_alf_sorting") diff --git a/src/spikeinterface/extractors/cbin_ibl.py b/src/spikeinterface/extractors/cbin_ibl.py index 891cbaee07..472aa9337b 100644 --- a/src/spikeinterface/extractors/cbin_ibl.py +++ b/src/spikeinterface/extractors/cbin_ibl.py @@ -136,7 +136,7 @@ def get_traces(self, start_frame, end_frame, channel_indices): return traces[:, channel_indices] -read_cbin_ibl = define_function_from_class(source_class=CompressedBinaryIblExtractor, name="read_cbin_ibl") +define_function_from_class(source_class=CompressedBinaryIblExtractor, name="read_cbin_ibl") def extract_stream_info(meta_file, meta): diff --git a/src/spikeinterface/extractors/cellexplorersortingextractor.py b/src/spikeinterface/extractors/cellexplorersortingextractor.py index 8659baded8..4356a7751d 100644 --- a/src/spikeinterface/extractors/cellexplorersortingextractor.py +++ b/src/spikeinterface/extractors/cellexplorersortingextractor.py @@ -195,4 +195,4 @@ def get_unit_spike_train( return spike_frames -read_cellexplorer = define_function_from_class(source_class=CellExplorerSortingExtractor, name="read_cellexplorer") +define_function_from_class(source_class=CellExplorerSortingExtractor, name="read_cellexplorer") diff --git a/src/spikeinterface/extractors/combinatoextractors.py b/src/spikeinterface/extractors/combinatoextractors.py index 5ecc5cdf24..5e4f98378d 100644 --- a/src/spikeinterface/extractors/combinatoextractors.py +++ b/src/spikeinterface/extractors/combinatoextractors.py @@ -100,4 +100,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_combinato = define_function_from_class(source_class=CombinatoSortingExtractor, name="read_combinato") +define_function_from_class(source_class=CombinatoSortingExtractor, name="read_combinato") diff --git a/src/spikeinterface/extractors/hdsortextractors.py b/src/spikeinterface/extractors/hdsortextractors.py index e60164fe21..bf2fa7809b 100644 --- a/src/spikeinterface/extractors/hdsortextractors.py +++ b/src/spikeinterface/extractors/hdsortextractors.py @@ -257,4 +257,4 @@ def _squeeze(arr): return arr -read_hdsort = define_function_from_class(source_class=HDSortSortingExtractor, name="read_hdsort") +define_function_from_class(source_class=HDSortSortingExtractor, name="read_hdsort") diff --git a/src/spikeinterface/extractors/herdingspikesextractors.py b/src/spikeinterface/extractors/herdingspikesextractors.py index 1450d4b8a5..7a237e603b 100644 --- a/src/spikeinterface/extractors/herdingspikesextractors.py +++ b/src/spikeinterface/extractors/herdingspikesextractors.py @@ -75,4 +75,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_herdingspikes = define_function_from_class(source_class=HerdingspikesSortingExtractor, name="read_herdingspikes") +define_function_from_class(source_class=HerdingspikesSortingExtractor, name="read_herdingspikes") diff --git a/src/spikeinterface/extractors/iblextractors.py b/src/spikeinterface/extractors/iblextractors.py index 8a57e40ec3..5a091cb945 100644 --- a/src/spikeinterface/extractors/iblextractors.py +++ b/src/spikeinterface/extractors/iblextractors.py @@ -364,5 +364,5 @@ def __init__( self._kwargs = dict(pid=pid, good_clusters_only=good_clusters_only, load_unit_properties=load_unit_properties) -read_ibl_recording = define_function_from_class(source_class=IblRecordingExtractor, name="read_ibl_recording") -read_ibl_sorting = define_function_from_class(source_class=IblSortingExtractor, name="read_ibl_sorting") +define_function_from_class(source_class=IblRecordingExtractor, name="read_ibl_recording") +define_function_from_class(source_class=IblSortingExtractor, name="read_ibl_sorting") diff --git a/src/spikeinterface/extractors/klustaextractors.py b/src/spikeinterface/extractors/klustaextractors.py index d18c8d831f..e968f6fb5d 100644 --- a/src/spikeinterface/extractors/klustaextractors.py +++ b/src/spikeinterface/extractors/klustaextractors.py @@ -150,4 +150,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_klusta = define_function_from_class(source_class=KlustaSortingExtractor, name="read_klusta") +define_function_from_class(source_class=KlustaSortingExtractor, name="read_klusta") diff --git a/src/spikeinterface/extractors/mclustextractors.py b/src/spikeinterface/extractors/mclustextractors.py index f04ca82b06..6dcbfc5382 100644 --- a/src/spikeinterface/extractors/mclustextractors.py +++ b/src/spikeinterface/extractors/mclustextractors.py @@ -101,4 +101,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_mclust = define_function_from_class(source_class=MClustSortingExtractor, name="read_mclust") +define_function_from_class(source_class=MClustSortingExtractor, name="read_mclust") diff --git a/src/spikeinterface/extractors/mcsh5extractors.py b/src/spikeinterface/extractors/mcsh5extractors.py index 3e295f33b4..1a07cbecba 100644 --- a/src/spikeinterface/extractors/mcsh5extractors.py +++ b/src/spikeinterface/extractors/mcsh5extractors.py @@ -161,4 +161,4 @@ def openMCSH5File(filename, stream_id): return mcs_info -read_mcsh5 = define_function_from_class(source_class=MCSH5RecordingExtractor, name="read_mcsh5") +define_function_from_class(source_class=MCSH5RecordingExtractor, name="read_mcsh5") diff --git a/src/spikeinterface/extractors/mdaextractors.py b/src/spikeinterface/extractors/mdaextractors.py index 12a7474df7..9e6b5c1112 100644 --- a/src/spikeinterface/extractors/mdaextractors.py +++ b/src/spikeinterface/extractors/mdaextractors.py @@ -290,8 +290,8 @@ def get_unit_spike_train( return np.rint(self._spike_times[inds]).astype(int) -read_mda_recording = define_function_from_class(source_class=MdaRecordingExtractor, name="read_mda_recording") -read_mda_sorting = define_function_from_class(source_class=MdaSortingExtractor, name="read_mda_sorting") +define_function_from_class(source_class=MdaRecordingExtractor, name="read_mda_recording") +define_function_from_class(source_class=MdaSortingExtractor, name="read_mda_sorting") def _concatenate(list): diff --git a/src/spikeinterface/extractors/neoextractors/alphaomega.py b/src/spikeinterface/extractors/neoextractors/alphaomega.py index b6aecd7eee..20b4b5fdf6 100644 --- a/src/spikeinterface/extractors/neoextractors/alphaomega.py +++ b/src/spikeinterface/extractors/neoextractors/alphaomega.py @@ -88,5 +88,5 @@ def map_to_neo_kwargs(cls, folder_path): return neo_kwargs -read_alphaomega = define_function_from_class(source_class=AlphaOmegaRecordingExtractor, name="read_alphaomega") -read_alphaomega_event = define_function_from_class(source_class=AlphaOmegaEventExtractor, name="read_alphaomega_event") +define_function_from_class(source_class=AlphaOmegaRecordingExtractor, name="read_alphaomega") +define_function_from_class(source_class=AlphaOmegaEventExtractor, name="read_alphaomega_event") diff --git a/src/spikeinterface/extractors/neoextractors/axon.py b/src/spikeinterface/extractors/neoextractors/axon.py index f225a091ae..2309f8bb8d 100644 --- a/src/spikeinterface/extractors/neoextractors/axon.py +++ b/src/spikeinterface/extractors/neoextractors/axon.py @@ -61,4 +61,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_axon = define_function_from_class(source_class=AxonRecordingExtractor, name="read_axon") +define_function_from_class(source_class=AxonRecordingExtractor, name="read_axon") diff --git a/src/spikeinterface/extractors/neoextractors/axona.py b/src/spikeinterface/extractors/neoextractors/axona.py index 40be565896..9f411ea9b4 100644 --- a/src/spikeinterface/extractors/neoextractors/axona.py +++ b/src/spikeinterface/extractors/neoextractors/axona.py @@ -45,4 +45,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_axona = define_function_from_class(source_class=AxonaRecordingExtractor, name="read_axona") +define_function_from_class(source_class=AxonaRecordingExtractor, name="read_axona") diff --git a/src/spikeinterface/extractors/neoextractors/biocam.py b/src/spikeinterface/extractors/neoextractors/biocam.py index b3ccb92cbd..60fd937df1 100644 --- a/src/spikeinterface/extractors/neoextractors/biocam.py +++ b/src/spikeinterface/extractors/neoextractors/biocam.py @@ -93,4 +93,4 @@ def map_to_neo_kwargs(cls, file_path, fill_gaps_strategy): return neo_kwargs -read_biocam = define_function_from_class(source_class=BiocamRecordingExtractor, name="read_biocam") +define_function_from_class(source_class=BiocamRecordingExtractor, name="read_biocam") diff --git a/src/spikeinterface/extractors/neoextractors/blackrock.py b/src/spikeinterface/extractors/neoextractors/blackrock.py index 08b3645bb2..8a04fb11f9 100644 --- a/src/spikeinterface/extractors/neoextractors/blackrock.py +++ b/src/spikeinterface/extractors/neoextractors/blackrock.py @@ -135,7 +135,5 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_blackrock = define_function_from_class(source_class=BlackrockRecordingExtractor, name="read_blackrock") -read_blackrock_sorting = define_function_from_class( - source_class=BlackrockSortingExtractor, name="read_blackrock_sorting" -) +define_function_from_class(source_class=BlackrockRecordingExtractor, name="read_blackrock") +define_function_from_class(source_class=BlackrockSortingExtractor, name="read_blackrock_sorting") diff --git a/src/spikeinterface/extractors/neoextractors/ced.py b/src/spikeinterface/extractors/neoextractors/ced.py index bc7a7c41ad..93d9c94943 100644 --- a/src/spikeinterface/extractors/neoextractors/ced.py +++ b/src/spikeinterface/extractors/neoextractors/ced.py @@ -56,4 +56,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_ced = define_function_from_class(source_class=CedRecordingExtractor, name="read_ced") +define_function_from_class(source_class=CedRecordingExtractor, name="read_ced") diff --git a/src/spikeinterface/extractors/neoextractors/edf.py b/src/spikeinterface/extractors/neoextractors/edf.py index 35f41b7dd1..b856cdb5a1 100644 --- a/src/spikeinterface/extractors/neoextractors/edf.py +++ b/src/spikeinterface/extractors/neoextractors/edf.py @@ -55,4 +55,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_edf = define_function_from_class(source_class=EDFRecordingExtractor, name="read_edf") +define_function_from_class(source_class=EDFRecordingExtractor, name="read_edf") diff --git a/src/spikeinterface/extractors/neoextractors/intan.py b/src/spikeinterface/extractors/neoextractors/intan.py index 793026ea94..f118d4f0b0 100644 --- a/src/spikeinterface/extractors/neoextractors/intan.py +++ b/src/spikeinterface/extractors/neoextractors/intan.py @@ -104,7 +104,7 @@ def _add_channel_groups(self): self.set_property(key="group_names", values=group_names) -read_intan = define_function_from_class(source_class=IntanRecordingExtractor, name="read_intan") +define_function_from_class(source_class=IntanRecordingExtractor, name="read_intan") class IntanSplitFilesRecordingExtractor(ConcatenateSegmentRecording, AppendSegmentRecording): @@ -205,6 +205,4 @@ def __init__( ) -read_split_intan_files = define_function_from_class( - source_class=IntanSplitFilesRecordingExtractor, name="read_split_intan_files" -) +define_function_from_class(source_class=IntanSplitFilesRecordingExtractor, name="read_split_intan_files") diff --git a/src/spikeinterface/extractors/neoextractors/maxwell.py b/src/spikeinterface/extractors/neoextractors/maxwell.py index 38e65096c2..5241027e66 100644 --- a/src/spikeinterface/extractors/neoextractors/maxwell.py +++ b/src/spikeinterface/extractors/neoextractors/maxwell.py @@ -148,5 +148,5 @@ def get_events(self, channel_id, start_time, end_time): return event -read_maxwell = define_function_from_class(source_class=MaxwellRecordingExtractor, name="read_maxwell") -read_maxwell_event = define_function_from_class(source_class=MaxwellEventExtractor, name="read_maxwell_event") +define_function_from_class(source_class=MaxwellRecordingExtractor, name="read_maxwell") +define_function_from_class(source_class=MaxwellEventExtractor, name="read_maxwell_event") diff --git a/src/spikeinterface/extractors/neoextractors/mcsraw.py b/src/spikeinterface/extractors/neoextractors/mcsraw.py index 06bf9e3324..5f13122840 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_mcsraw") +define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_mcsraw") diff --git a/src/spikeinterface/extractors/neoextractors/neuralynx.py b/src/spikeinterface/extractors/neoextractors/neuralynx.py index 81f507535c..23510141a2 100644 --- a/src/spikeinterface/extractors/neoextractors/neuralynx.py +++ b/src/spikeinterface/extractors/neoextractors/neuralynx.py @@ -127,7 +127,5 @@ def map_to_neo_kwargs(cls, folder_path): return neo_kwargs -read_neuralynx = define_function_from_class(source_class=NeuralynxRecordingExtractor, name="read_neuralynx") -read_neuralynx_sorting = define_function_from_class( - source_class=NeuralynxSortingExtractor, name="read_neuralynx_sorting" -) +define_function_from_class(source_class=NeuralynxRecordingExtractor, name="read_neuralynx") +define_function_from_class(source_class=NeuralynxSortingExtractor, name="read_neuralynx_sorting") diff --git a/src/spikeinterface/extractors/neoextractors/neuroexplorer.py b/src/spikeinterface/extractors/neoextractors/neuroexplorer.py index 8a076264ec..f6b82f7ea5 100644 --- a/src/spikeinterface/extractors/neoextractors/neuroexplorer.py +++ b/src/spikeinterface/extractors/neoextractors/neuroexplorer.py @@ -71,4 +71,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_neuroexplorer = define_function_from_class(source_class=NeuroExplorerRecordingExtractor, name="read_neuroexplorer") +define_function_from_class(source_class=NeuroExplorerRecordingExtractor, name="read_neuroexplorer") diff --git a/src/spikeinterface/extractors/neoextractors/neuronexus.py b/src/spikeinterface/extractors/neoextractors/neuronexus.py index 92bb90a08b..4d94a9b30d 100644 --- a/src/spikeinterface/extractors/neoextractors/neuronexus.py +++ b/src/spikeinterface/extractors/neoextractors/neuronexus.py @@ -61,4 +61,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_neuronexus = define_function_from_class(source_class=NeuroNexusRecordingExtractor, name="read_neuronexus") +define_function_from_class(source_class=NeuroNexusRecordingExtractor, name="read_neuronexus") diff --git a/src/spikeinterface/extractors/neoextractors/neuroscope.py b/src/spikeinterface/extractors/neoextractors/neuroscope.py index 7064ea1303..3fe21f62b9 100644 --- a/src/spikeinterface/extractors/neoextractors/neuroscope.py +++ b/src/spikeinterface/extractors/neoextractors/neuroscope.py @@ -355,12 +355,8 @@ def _handle_xml_file_path(folder_path: str | Path, initial_xml_file_path: str | return xml_file_path -read_neuroscope_recording = define_function_from_class( - source_class=NeuroScopeRecordingExtractor, name="read_neuroscope_recording" -) -read_neuroscope_sorting = define_function_from_class( - source_class=NeuroScopeSortingExtractor, name="read_neuroscope_sorting" -) +define_function_from_class(source_class=NeuroScopeRecordingExtractor, name="read_neuroscope_recording") +define_function_from_class(source_class=NeuroScopeSortingExtractor, name="read_neuroscope_sorting") def read_neuroscope( diff --git a/src/spikeinterface/extractors/neoextractors/nix.py b/src/spikeinterface/extractors/neoextractors/nix.py index 9e959fc3ed..fbbe922caf 100644 --- a/src/spikeinterface/extractors/neoextractors/nix.py +++ b/src/spikeinterface/extractors/neoextractors/nix.py @@ -58,4 +58,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_nix = define_function_from_class(source_class=NixRecordingExtractor, name="read_nix") +define_function_from_class(source_class=NixRecordingExtractor, name="read_nix") diff --git a/src/spikeinterface/extractors/neoextractors/plexon.py b/src/spikeinterface/extractors/neoextractors/plexon.py index f8131a816a..fa72b6810d 100644 --- a/src/spikeinterface/extractors/neoextractors/plexon.py +++ b/src/spikeinterface/extractors/neoextractors/plexon.py @@ -90,5 +90,5 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_plexon = define_function_from_class(source_class=PlexonRecordingExtractor, name="read_plexon") -read_plexon_sorting = define_function_from_class(source_class=PlexonSortingExtractor, name="read_plexon_sorting") +define_function_from_class(source_class=PlexonRecordingExtractor, name="read_plexon") +define_function_from_class(source_class=PlexonSortingExtractor, name="read_plexon_sorting") diff --git a/src/spikeinterface/extractors/neoextractors/plexon2.py b/src/spikeinterface/extractors/neoextractors/plexon2.py index 6abba1ce58..bb0d049171 100644 --- a/src/spikeinterface/extractors/neoextractors/plexon2.py +++ b/src/spikeinterface/extractors/neoextractors/plexon2.py @@ -137,6 +137,6 @@ def map_to_neo_kwargs(cls, folder_path): return neo_kwargs -read_plexon2 = define_function_from_class(source_class=Plexon2RecordingExtractor, name="read_plexon2") -read_plexon2_sorting = define_function_from_class(source_class=Plexon2SortingExtractor, name="read_plexon2_sorting") -read_plexon2_event = define_function_from_class(source_class=Plexon2EventExtractor, name="read_plexon2_event") +define_function_from_class(source_class=Plexon2RecordingExtractor, name="read_plexon2") +define_function_from_class(source_class=Plexon2SortingExtractor, name="read_plexon2_sorting") +define_function_from_class(source_class=Plexon2EventExtractor, name="read_plexon2_event") diff --git a/src/spikeinterface/extractors/neoextractors/spike2.py b/src/spikeinterface/extractors/neoextractors/spike2.py index 6283176def..970fce1f3e 100644 --- a/src/spikeinterface/extractors/neoextractors/spike2.py +++ b/src/spikeinterface/extractors/neoextractors/spike2.py @@ -50,4 +50,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_spike2 = define_function_from_class(source_class=Spike2RecordingExtractor, name="read_spike2") +define_function_from_class(source_class=Spike2RecordingExtractor, name="read_spike2") diff --git a/src/spikeinterface/extractors/neoextractors/spikegadgets.py b/src/spikeinterface/extractors/neoextractors/spikegadgets.py index da4a66e1f5..3b7a4bc5a9 100644 --- a/src/spikeinterface/extractors/neoextractors/spikegadgets.py +++ b/src/spikeinterface/extractors/neoextractors/spikegadgets.py @@ -94,4 +94,4 @@ def map_to_neo_kwargs(cls, file_path): return neo_kwargs -read_spikegadgets = define_function_from_class(source_class=SpikeGadgetsRecordingExtractor, name="read_spikegadgets") +define_function_from_class(source_class=SpikeGadgetsRecordingExtractor, name="read_spikegadgets") diff --git a/src/spikeinterface/extractors/neoextractors/spikeglx.py b/src/spikeinterface/extractors/neoextractors/spikeglx.py index 7b45c8b3ab..86e731ebf8 100644 --- a/src/spikeinterface/extractors/neoextractors/spikeglx.py +++ b/src/spikeinterface/extractors/neoextractors/spikeglx.py @@ -121,7 +121,7 @@ def _handle_kwargs_backward_compatibility(cls, old_kwargs, full_dict): return new_kwargs -read_spikeglx = define_function_from_class(source_class=SpikeGLXRecordingExtractor, name="read_spikeglx") +define_function_from_class(source_class=SpikeGLXRecordingExtractor, name="read_spikeglx") class SpikeGLXEventExtractor(NeoBaseEventExtractor): diff --git a/src/spikeinterface/extractors/neoextractors/tdt.py b/src/spikeinterface/extractors/neoextractors/tdt.py index 593e623f19..ab5d9929f9 100644 --- a/src/spikeinterface/extractors/neoextractors/tdt.py +++ b/src/spikeinterface/extractors/neoextractors/tdt.py @@ -57,4 +57,4 @@ def map_to_neo_kwargs(cls, folder_path): return neo_kwargs -read_tdt = define_function_from_class(source_class=TdtRecordingExtractor, name="read_tdt") +define_function_from_class(source_class=TdtRecordingExtractor, name="read_tdt") diff --git a/src/spikeinterface/extractors/nwbextractors.py b/src/spikeinterface/extractors/nwbextractors.py index fd45aa07c2..ef911b97f5 100644 --- a/src/spikeinterface/extractors/nwbextractors.py +++ b/src/spikeinterface/extractors/nwbextractors.py @@ -1833,9 +1833,9 @@ def get_traces(self, start_frame, end_frame, channel_indices): # Create the reading function -read_nwb_recording = define_function_from_class(source_class=NwbRecordingExtractor, name="read_nwb_recording") -read_nwb_sorting = define_function_from_class(source_class=NwbSortingExtractor, name="read_nwb_sorting") -read_nwb_timeseries = define_function_from_class(source_class=NwbTimeSeriesExtractor, name="read_nwb_timeseries") +define_function_from_class(source_class=NwbRecordingExtractor, name="read_nwb_recording") +define_function_from_class(source_class=NwbSortingExtractor, name="read_nwb_sorting") +define_function_from_class(source_class=NwbTimeSeriesExtractor, name="read_nwb_timeseries") def read_nwb(file_path, load_recording=True, load_sorting=False, electrical_series_path=None): diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 92ae2a0437..7dd6d6ef56 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -345,8 +345,8 @@ def __init__(self, folder_path: Path | str, keep_good_only: bool = False, remove self._kwargs = {"folder_path": str(Path(folder_path).absolute()), "keep_good_only": keep_good_only} -read_phy = define_function_from_class(source_class=PhySortingExtractor, name="read_phy") -read_kilosort = define_function_from_class(source_class=KiloSortSortingExtractor, name="read_kilosort") +define_function_from_class(source_class=PhySortingExtractor, name="read_phy") +define_function_from_class(source_class=KiloSortSortingExtractor, name="read_kilosort") def read_kilosort_as_analyzer(folder_path, unwhiten=True, gain_to_uV=None, offset_to_uV=None) -> SortingAnalyzer: diff --git a/src/spikeinterface/extractors/shybridextractors.py b/src/spikeinterface/extractors/shybridextractors.py index eca7d46724..429c4942d8 100644 --- a/src/spikeinterface/extractors/shybridextractors.py +++ b/src/spikeinterface/extractors/shybridextractors.py @@ -238,10 +238,8 @@ def get_unit_spike_train( return train[idxs] -read_shybrid_recording = define_function_from_class( - source_class=SHYBRIDRecordingExtractor, name="read_shybrid_recording" -) -read_shybrid_sorting = define_function_from_class(source_class=SHYBRIDSortingExtractor, name="read_shybrid_sorting") +define_function_from_class(source_class=SHYBRIDRecordingExtractor, name="read_shybrid_recording") +define_function_from_class(source_class=SHYBRIDSortingExtractor, name="read_shybrid_sorting") class GeometryNotLoadedError(Exception): diff --git a/src/spikeinterface/extractors/sinapsrecordingextractors.py b/src/spikeinterface/extractors/sinapsrecordingextractors.py index f47a83bc47..1bedc84151 100644 --- a/src/spikeinterface/extractors/sinapsrecordingextractors.py +++ b/src/spikeinterface/extractors/sinapsrecordingextractors.py @@ -192,11 +192,9 @@ def get_traces(self, start_frame=None, end_frame=None, channel_indices=None): return traces -read_sinaps_research_platform = define_function_from_class( - source_class=SinapsResearchPlatformRecordingExtractor, name="read_sinaps_research_platform" -) +define_function_from_class(source_class=SinapsResearchPlatformRecordingExtractor, name="read_sinaps_research_platform") -read_sinaps_research_platform_h5 = define_function_from_class( +define_function_from_class( source_class=SinapsResearchPlatformH5RecordingExtractor, name="read_sinaps_research_platform_h5" ) diff --git a/src/spikeinterface/extractors/spykingcircusextractors.py b/src/spikeinterface/extractors/spykingcircusextractors.py index 0697bf949d..efd1fa6627 100644 --- a/src/spikeinterface/extractors/spykingcircusextractors.py +++ b/src/spikeinterface/extractors/spykingcircusextractors.py @@ -109,4 +109,4 @@ def _load_sample_rate(params_file): return sample_rate -read_spykingcircus = define_function_from_class(source_class=SpykingCircusSortingExtractor, name="read_spykingcircus") +define_function_from_class(source_class=SpykingCircusSortingExtractor, name="read_spykingcircus") diff --git a/src/spikeinterface/extractors/tridesclousextractors.py b/src/spikeinterface/extractors/tridesclousextractors.py index 146872ab9b..c4b9bf68a4 100644 --- a/src/spikeinterface/extractors/tridesclousextractors.py +++ b/src/spikeinterface/extractors/tridesclousextractors.py @@ -71,4 +71,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return spike_times.copy() -read_tridesclous = define_function_from_class(source_class=TridesclousSortingExtractor, name="read_tridesclous") +define_function_from_class(source_class=TridesclousSortingExtractor, name="read_tridesclous") diff --git a/src/spikeinterface/extractors/waveclussnippetstextractors.py b/src/spikeinterface/extractors/waveclussnippetstextractors.py index 72bdbd089a..74800d0aa1 100644 --- a/src/spikeinterface/extractors/waveclussnippetstextractors.py +++ b/src/spikeinterface/extractors/waveclussnippetstextractors.py @@ -149,6 +149,4 @@ def get_frames(self, indices=None): return self._spikestimes[indices] -read_waveclus_snippets = define_function_from_class( - source_class=WaveClusSnippetsExtractor, name="read_waveclus_snippets" -) +define_function_from_class(source_class=WaveClusSnippetsExtractor, name="read_waveclus_snippets") diff --git a/src/spikeinterface/extractors/waveclustextractors.py b/src/spikeinterface/extractors/waveclustextractors.py index d8b02ba1f5..7040aa7558 100644 --- a/src/spikeinterface/extractors/waveclustextractors.py +++ b/src/spikeinterface/extractors/waveclustextractors.py @@ -60,4 +60,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_waveclus = define_function_from_class(source_class=WaveClusSortingExtractor, name="read_waveclus") +define_function_from_class(source_class=WaveClusSortingExtractor, name="read_waveclus") diff --git a/src/spikeinterface/extractors/whitematterrecordingextractor.py b/src/spikeinterface/extractors/whitematterrecordingextractor.py index d1abc090ed..c2104e5668 100644 --- a/src/spikeinterface/extractors/whitematterrecordingextractor.py +++ b/src/spikeinterface/extractors/whitematterrecordingextractor.py @@ -91,4 +91,4 @@ def __init__( # Define function equivalent for convenience -read_whitematter = define_function_from_class(source_class=WhiteMatterRecordingExtractor, name="read_whitematter") +define_function_from_class(source_class=WhiteMatterRecordingExtractor, name="read_whitematter") diff --git a/src/spikeinterface/extractors/xclustextractors.py b/src/spikeinterface/extractors/xclustextractors.py index 44f95f5623..d6b47a796a 100644 --- a/src/spikeinterface/extractors/xclustextractors.py +++ b/src/spikeinterface/extractors/xclustextractors.py @@ -165,4 +165,4 @@ def get_unit_spike_train_in_seconds(self, unit_id, start_time=None, end_time=Non return spike_times[start_index:end_index] -read_xclust = define_function_from_class(source_class=XClustSortingExtractor, name="read_xclust") +define_function_from_class(source_class=XClustSortingExtractor, name="read_xclust") diff --git a/src/spikeinterface/extractors/yassextractors.py b/src/spikeinterface/extractors/yassextractors.py index 38efeae996..9e29369466 100644 --- a/src/spikeinterface/extractors/yassextractors.py +++ b/src/spikeinterface/extractors/yassextractors.py @@ -66,4 +66,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return times -read_yass = define_function_from_class(source_class=YassSortingExtractor, name="read_yass") +define_function_from_class(source_class=YassSortingExtractor, name="read_yass") diff --git a/src/spikeinterface/generation/noise_tools.py b/src/spikeinterface/generation/noise_tools.py index 87c9c52387..096cad95ef 100644 --- a/src/spikeinterface/generation/noise_tools.py +++ b/src/spikeinterface/generation/noise_tools.py @@ -223,9 +223,7 @@ def get_traces( return traces -noise_generator_recording = define_function_from_class( - source_class=NoiseGeneratorRecording, name="noise_generator_recording" -) +define_function_from_class(source_class=NoiseGeneratorRecording, name="noise_generator_recording") def generate_noise( diff --git a/src/spikeinterface/postprocessing/alignsorting.py b/src/spikeinterface/postprocessing/alignsorting.py index cf4189a3c7..5ee035beff 100644 --- a/src/spikeinterface/postprocessing/alignsorting.py +++ b/src/spikeinterface/postprocessing/alignsorting.py @@ -51,4 +51,4 @@ def get_unit_spike_train(self, unit_id, start_frame, end_frame): return original_spike_train - self._unit_peak_shifts[unit_id] -align_sorting = define_function_from_class(source_class=AlignSortingExtractor, name="align_sorting") +define_function_from_class(source_class=AlignSortingExtractor, name="align_sorting")