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
8 changes: 2 additions & 6 deletions src/spikeinterface/comparison/multicomparisons.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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")
8 changes: 3 additions & 5 deletions src/spikeinterface/comparison/paircomparisons.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/core/binaryfolder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/core/binaryrecordingextractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
16 changes: 5 additions & 11 deletions src/spikeinterface/core/core_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,20 +56,14 @@ 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
# 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, name: str) -> None:
"Wrapper to inject source_class into the caller's module namespace under name."


def define_function_from_class(source_class: Callable[P, T], name: str) -> Callable[P, T]:
"Wrapper to change the name of a class"

return source_class
sys._getframe(1).f_globals[name] = source_class


def read_python(path):
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/core/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ##
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/core/npyfoldersnippets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 2 additions & 2 deletions src/spikeinterface/core/npysnippetsextractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/core/npzsortingextractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
18 changes: 7 additions & 11 deletions 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")
define_function_from_class(source_class=AppendSegmentRecording, name="append_recordings")


class ConcatenateSegmentRecording(BaseRecording):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down
6 changes: 2 additions & 4 deletions src/spikeinterface/core/sortingfolder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/core/unitsaggregationsorting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 2 additions & 2 deletions src/spikeinterface/core/zarrextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/curation/curationsorting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/curation/mergeunitssorting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
4 changes: 1 addition & 3 deletions src/spikeinterface/curation/remove_duplicated_spikes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/curation/splitunitsorting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/alfsortingextractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/cbin_ibl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/combinatoextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/hdsortextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/herdingspikesextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 2 additions & 2 deletions src/spikeinterface/extractors/iblextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/klustaextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/mclustextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/mcsh5extractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
4 changes: 2 additions & 2 deletions src/spikeinterface/extractors/mdaextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
4 changes: 2 additions & 2 deletions src/spikeinterface/extractors/neoextractors/alphaomega.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/axon.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/axona.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/biocam.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
6 changes: 2 additions & 4 deletions src/spikeinterface/extractors/neoextractors/blackrock.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/ced.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
2 changes: 1 addition & 1 deletion src/spikeinterface/extractors/neoextractors/edf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
6 changes: 2 additions & 4 deletions src/spikeinterface/extractors/neoextractors/intan.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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")
4 changes: 2 additions & 2 deletions src/spikeinterface/extractors/neoextractors/maxwell.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
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")
define_function_from_class(source_class=MCSRawRecordingExtractor, name="read_mcsraw")
Loading
Loading