diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 6578b0f390..8c57d3a9c5 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -67,6 +67,7 @@ def __init__( keep_good_only: bool = False, remove_empty_units: bool = False, load_all_cluster_properties: bool = True, + channel_ids: list | np.ndarray | None = None, ): try: import pandas as pd @@ -229,6 +230,12 @@ def __init__( values_ = cluster_info[prop_name].values self.set_property(key=prop_name, values=values_) + if channel_ids is not None: + main_channel_indices = _make_main_channel_indices_from_templates(phy_folder) + if main_channel_indices is not None: + main_channel_ids = channel_ids[main_channel_indices] + self.set_property(key="main_channel_id", values=main_channel_ids) + self.annotate(phy_folder=str(phy_folder.resolve())) self.add_sorting_segment(PhySortingSegment(spike_times_clean, spike_clusters_clean)) @@ -334,13 +341,20 @@ class KiloSortSortingExtractor(BasePhyKilosortSortingExtractor): The loaded Sorting object. """ - def __init__(self, folder_path: Path | str, keep_good_only: bool = False, remove_empty_units: bool = True): + def __init__( + self, + folder_path: Path | str, + keep_good_only: bool = False, + remove_empty_units: bool = True, + channel_ids: list | np.ndarray | None = None, + ): BasePhyKilosortSortingExtractor.__init__( self, folder_path, exclude_cluster_groups=None, keep_good_only=keep_good_only, remove_empty_units=remove_empty_units, + channel_ids=channel_ids, ) self._kwargs = {"folder_path": str(Path(folder_path).absolute()), "keep_good_only": keep_good_only} @@ -431,7 +445,7 @@ def read_kilosort_as_analyzer( recording = aggregate_channels(recordings) sparsity = _make_sparsity_from_templates(sorting, recording, phy_path) - main_channel_indices = _make_main_channel_indices_from_templates(sorting, recording, phy_path) + main_channel_indices = _make_main_channel_indices_from_templates(phy_path) sorting_analyzer = create_sorting_analyzer( sorting, recording, sparse=True, sparsity=sparsity, main_channel_indices=main_channel_indices @@ -503,14 +517,20 @@ def _make_sparsity_from_templates(sorting, recording, kilosort_output_path): return ChannelSparsity(mask, unit_ids=unit_ids, channel_ids=channel_ids) -def _make_main_channel_indices_from_templates(sorting, recording, kilosort_output_path): +def _make_main_channel_indices_from_templates(kilosort_output_path): """Constructs the `main_channel_indices` from kilosort output, by finding the channel containing the largest peak-to-peak value.""" - templates = np.load(kilosort_output_path / "templates.npy") - # main channel indices are the argmax of the ptp of the templates, which is the channel with - # the largest peak-to-peak amplitude - main_channel_indices = np.argmax(np.ptp(templates, axis=1), axis=1) + templates_filepath = kilosort_output_path / "templates.npy" + + if templates_filepath.is_file(): + templates = np.load(kilosort_output_path / "templates.npy") + # main channel indices are the argmax of the ptp of the templates, which is the channel with + # the largest peak-to-peak amplitude + main_channel_indices = np.argmax(np.ptp(templates, axis=1), axis=1) + else: + main_channel_indices = None + return main_channel_indices diff --git a/src/spikeinterface/sorters/external/kilosort4.py b/src/spikeinterface/sorters/external/kilosort4.py index bd856dc13b..ea5e596b79 100644 --- a/src/spikeinterface/sorters/external/kilosort4.py +++ b/src/spikeinterface/sorters/external/kilosort4.py @@ -7,7 +7,6 @@ from spikeinterface.core import write_binary_recording, Motion, BaseRecording from spikeinterface.sorters.basesorter import BaseSorter, get_job_kwargs from .kilosortbase import KilosortBase -from spikeinterface.sorters.basesorter import get_job_kwargs from importlib.metadata import version as importlib_version diff --git a/src/spikeinterface/sorters/external/kilosortbase.py b/src/spikeinterface/sorters/external/kilosortbase.py index 4cdb51e21e..99ebbab758 100644 --- a/src/spikeinterface/sorters/external/kilosortbase.py +++ b/src/spikeinterface/sorters/external/kilosortbase.py @@ -9,7 +9,7 @@ from spikeinterface.sorters.utils import ShellScript, get_matlab_shell_name, get_bash_path from spikeinterface.sorters.basesorter import get_job_kwargs from spikeinterface.extractors.extractor_classes import KiloSortSortingExtractor -from spikeinterface.core import write_binary_recording +from spikeinterface.core import write_binary_recording, load from spikeinterface.preprocessing.zero_channel_pad import TracePaddedRecording @@ -254,6 +254,20 @@ def _get_result_from_folder(cls, sorter_output_folder): params_file = sorter_output_folder / "spikeinterface_params.json" with params_file.open("r") as f: sorter_params = json.load(f)["sorter_params"] + + recording_file = sorter_output_folder.parent / "spikeinterface_recording.json" + if recording_file.is_file(): + try: + # TODO: load the channel ids without loading the recording + recording = load(recording_file) + channel_ids = recording.channel_ids + except: + channel_ids = None + else: + channel_ids = None + keep_good_only = sorter_params.get("keep_good_only", False) - sorting = KiloSortSortingExtractor(folder_path=sorter_output_folder, keep_good_only=keep_good_only) + sorting = KiloSortSortingExtractor( + folder_path=sorter_output_folder, keep_good_only=keep_good_only, channel_ids=channel_ids + ) return sorting