From 670e88ad514ebaf7330827387bf5c53317638d3e Mon Sep 17 00:00:00 2001 From: chrishalcrow Date: Fri, 31 Jul 2026 11:27:03 +0100 Subject: [PATCH] add compute main channel ids to kilosort run sorter --- .../extractors/phykilosortextractors.py | 4 ++-- src/spikeinterface/sorters/external/kilosort4.py | 15 ++++++++++++++- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 92ae2a0437..335b80e1d0 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -415,7 +415,7 @@ def read_kilosort_as_analyzer(folder_path, unwhiten=True, gain_to_uV=None, offse ) 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 @@ -490,7 +490,7 @@ 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.""" diff --git a/src/spikeinterface/sorters/external/kilosort4.py b/src/spikeinterface/sorters/external/kilosort4.py index bd856dc13b..7796f556f7 100644 --- a/src/spikeinterface/sorters/external/kilosort4.py +++ b/src/spikeinterface/sorters/external/kilosort4.py @@ -1,4 +1,5 @@ import warnings +import csv from pathlib import Path from packaging import version @@ -7,7 +8,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 @@ -457,6 +457,19 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose): save_preprocessed_copy=save_preprocessed_copy, ) + if (results_dir / "templates.npy").is_file(): + # Note: these are the whitened templates + templates = np.load(results_dir / "templates.npy") + # main channel indices are the argmax of the ptp of the templates + main_channel_indices = np.argmax(np.ptp(templates, axis=1), axis=1) + main_channel_ids = recording.channel_ids[main_channel_indices] + # save main_channel_ids + with open(results_dir / "cluster_main_channel_id.tsv", "w", newline="", encoding="utf-8") as f: + writer = csv.writer(f, delimiter="\t") + writer.writerow(["cluster_id", "main_channel_id"]) + for unit_index, item in enumerate(main_channel_ids): + writer.writerow([unit_index, item]) + if params["delete_recording_dat"]: # only delete dat file if it was created by the wrapper if (sorter_output_folder / "recording.dat").is_file():