Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions src/spikeinterface/core/core_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/core/segmentutils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/mcsraw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/astype.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Original file line number Diff line number Diff line change
Expand Up @@ -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",
)
6 changes: 2 additions & 4 deletions src/spikeinterface/preprocessing/clip.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/common_reference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/decimate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/depth_order.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/detect_artifacts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/detect_bad_channels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/directional_derivative.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
8 changes: 4 additions & 4 deletions src/spikeinterface/preprocessing/filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/filter_gaussian.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/highpass_spatial_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")


# -----------------------------------------------------------------------------------------------
Expand Down
6 changes: 2 additions & 4 deletions src/spikeinterface/preprocessing/interpolate_bad_channels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -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")
10 changes: 4 additions & 6 deletions src/spikeinterface/preprocessing/normalize_scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/phase_shift.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/rectify.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/remove_artifacts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/resample.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/silence_periods.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 1 addition & 3 deletions src/spikeinterface/preprocessing/unsigned_to_signed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/preprocessing/whiten.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
6 changes: 2 additions & 4 deletions src/spikeinterface/preprocessing/zero_channel_pad.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Loading