Skip to content
Draft
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
34 changes: 27 additions & 7 deletions src/spikeinterface/extractors/phykilosortextractors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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


Expand Down
1 change: 0 additions & 1 deletion src/spikeinterface/sorters/external/kilosort4.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
18 changes: 16 additions & 2 deletions src/spikeinterface/sorters/external/kilosortbase.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we should also try the pickle if json do not exists no ?

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
Comment on lines +258 to +267

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Let's wait for #4671, where I added the main_ids to the to_dict.

So this should become:

with open(recording_file) as f:
    recording_dict = json.load(f)
    channel_ids = recording_dict.get("main_ids", None)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Then it will works only for recent run('kilosort') older folder will not have this main_channel_ids


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
Loading