From 1db746561883383fe342117558a2e3a520b12513 Mon Sep 17 00:00:00 2001 From: Pascal <65159092+pas-calc@users.noreply.github.com> Date: Sun, 8 Feb 2026 18:05:38 +0100 Subject: [PATCH 1/3] Add load_templates_clusters parameter (kilosort output // phy/manual refinement) Added a new parameter 'load_templates_clusters' to control loading of templates or clusters from Kilosort output. Updated logic to handle loading based on the parameter value. --- .../extractors/phykilosortextractors.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 68b16074fb..4d8c9e06eb 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -37,6 +37,9 @@ class BasePhyKilosortSortingExtractor(BaseSorting): If True, empty units are removed from the sorting extractor. load_all_cluster_properties : bool, default: True If True, all cluster properties are loaded from the tsv/csv files. + load_templates_clusters : str, templates|clusters|auto , default: "auto" + Defines whether to load templates (kilosort output) or clusters (after manual refinement) + If "auto", try to load clusters, fallback to templates if not existing Notes ----- @@ -68,6 +71,7 @@ def __init__( keep_good_only: bool = False, remove_empty_units: bool = False, load_all_cluster_properties: bool = True, + load_templates_clusters="auto", ): try: import pandas as pd @@ -77,10 +81,17 @@ def __init__( phy_folder = Path(folder_path) spike_times = np.load(phy_folder / "spike_times.npy").astype(int) - if (phy_folder / "spike_clusters.npy").is_file(): + if load_templates_clusters=="auto": + if (phy_folder / "spike_clusters.npy").is_file(): + spike_clusters = np.load(phy_folder / "spike_clusters.npy") + else: + spike_clusters = np.load(phy_folder / "spike_templates.npy") + elif load_templates_clusters=="templates": + spike_clusters = np.load(phy_folder / "spike_templates.npy") + elif load_templates_clusters=="clusters": spike_clusters = np.load(phy_folder / "spike_clusters.npy") else: - spike_clusters = np.load(phy_folder / "spike_templates.npy") + raise ValueError("Invalid value provided for load_templates_clusters: '{}'.".format(load_templates_clusters)) # spike_times and spike_clusters can be 2d sometimes --> convert to 1d. spike_times = np.atleast_1d(spike_times.squeeze()) From d2a32f405491fa8245a3b1f76bb889560b92cf2e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 8 Feb 2026 17:07:43 +0000 Subject: [PATCH 2/3] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/spikeinterface/extractors/phykilosortextractors.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 4d8c9e06eb..0401813e7c 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -81,17 +81,19 @@ def __init__( phy_folder = Path(folder_path) spike_times = np.load(phy_folder / "spike_times.npy").astype(int) - if load_templates_clusters=="auto": + if load_templates_clusters == "auto": if (phy_folder / "spike_clusters.npy").is_file(): spike_clusters = np.load(phy_folder / "spike_clusters.npy") else: spike_clusters = np.load(phy_folder / "spike_templates.npy") - elif load_templates_clusters=="templates": + elif load_templates_clusters == "templates": spike_clusters = np.load(phy_folder / "spike_templates.npy") - elif load_templates_clusters=="clusters": + elif load_templates_clusters == "clusters": spike_clusters = np.load(phy_folder / "spike_clusters.npy") else: - raise ValueError("Invalid value provided for load_templates_clusters: '{}'.".format(load_templates_clusters)) + raise ValueError( + "Invalid value provided for load_templates_clusters: '{}'.".format(load_templates_clusters) + ) # spike_times and spike_clusters can be 2d sometimes --> convert to 1d. spike_times = np.atleast_1d(spike_times.squeeze()) From a091d41e509310d617103067fa9116a2ef483e77 Mon Sep 17 00:00:00 2001 From: chrishalcrow Date: Tue, 15 Sep 2026 17:12:40 +0100 Subject: [PATCH 3/3] don't load metrics if using templates --- .../extractors/phykilosortextractors.py | 109 ++++++++++-------- 1 file changed, 59 insertions(+), 50 deletions(-) diff --git a/src/spikeinterface/extractors/phykilosortextractors.py b/src/spikeinterface/extractors/phykilosortextractors.py index 0401813e7c..c58b3e3bdf 100644 --- a/src/spikeinterface/extractors/phykilosortextractors.py +++ b/src/spikeinterface/extractors/phykilosortextractors.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Optional +from typing import Optional, Literal from pathlib import Path import warnings @@ -37,9 +37,9 @@ class BasePhyKilosortSortingExtractor(BaseSorting): If True, empty units are removed from the sorting extractor. load_all_cluster_properties : bool, default: True If True, all cluster properties are loaded from the tsv/csv files. - load_templates_clusters : str, templates|clusters|auto , default: "auto" - Defines whether to load templates (kilosort output) or clusters (after manual refinement) - If "auto", try to load clusters, fallback to templates if not existing + load_templates_or_cluster : str, Literal["templates, "clusters"] , default: "templates" + Defines whether to load templates (kilosort output) or clusters (after manual refinement + using phy). Notes ----- @@ -71,7 +71,7 @@ def __init__( keep_good_only: bool = False, remove_empty_units: bool = False, load_all_cluster_properties: bool = True, - load_templates_clusters="auto", + load_templates_or_cluster: Literal["templates", "clusters"] = "clusters", ): try: import pandas as pd @@ -81,19 +81,11 @@ def __init__( phy_folder = Path(folder_path) spike_times = np.load(phy_folder / "spike_times.npy").astype(int) - if load_templates_clusters == "auto": - if (phy_folder / "spike_clusters.npy").is_file(): - spike_clusters = np.load(phy_folder / "spike_clusters.npy") - else: - spike_clusters = np.load(phy_folder / "spike_templates.npy") - elif load_templates_clusters == "templates": - spike_clusters = np.load(phy_folder / "spike_templates.npy") - elif load_templates_clusters == "clusters": - spike_clusters = np.load(phy_folder / "spike_clusters.npy") + spike_times_filename = phy_folder / f"spike_{load_templates_or_cluster}.npy" + if spike_times_filename.is_file(): + spike_clusters = np.load(spike_times_filename) else: - raise ValueError( - "Invalid value provided for load_templates_clusters: '{}'.".format(load_templates_clusters) - ) + raise FileNotFoundError(f"Cannot find spike times file at {spike_times_filename}.") # spike_times and spike_clusters can be 2d sometimes --> convert to 1d. spike_times = np.atleast_1d(spike_times.squeeze()) @@ -104,43 +96,45 @@ def __init__( params = read_python(str(phy_folder / "params.py")) sampling_frequency = params["sample_rate"] - # try to load cluster info - cluster_info_files = [ - p for p in phy_folder.iterdir() if p.suffix in [".csv", ".tsv"] and "cluster_info" in p.name - ] + if load_templates_or_cluster == "clusters": - if len(cluster_info_files) == 1: - # load properties from cluster_info file - cluster_info_file = cluster_info_files[0] - if cluster_info_file.suffix == ".tsv": - delimiter = "\t" - else: - delimiter = "," - cluster_info = pd.read_csv(cluster_info_file, delimiter=delimiter) - else: - # load properties from other tsv/csv files - all_property_files = [p for p in phy_folder.iterdir() if p.suffix in [".csv", ".tsv"]] + # try to load cluster info + cluster_info_files = [ + p for p in phy_folder.iterdir() if p.suffix in [".csv", ".tsv"] and "cluster_info" in p.name + ] - cluster_info = None - for file in all_property_files: - if file.suffix == ".tsv": + if len(cluster_info_files) == 1: + # load properties from cluster_info file + cluster_info_file = cluster_info_files[0] + if cluster_info_file.suffix == ".tsv": delimiter = "\t" else: delimiter = "," - new_property = pd.read_csv(file, delimiter=delimiter) - - # Only merge files that contain a cluster_id column - # This prevents KeyError when extraneous files don't have cluster_id - # Typical aggregated files include cluster_group.tsv, cluster_info.tsv, cluster_KSLabel.tsv - # See Phy docs: https://phy.readthedocs.io/en/latest/sorting_user_guide/ - # See: https://github.com/SpikeInterface/spikeinterface/issues/4124 - if "cluster_id" not in new_property.columns: - continue - - if cluster_info is None: - cluster_info = new_property - else: - cluster_info = pd.merge(cluster_info, new_property, on="cluster_id", suffixes=[None, "_repeat"]) + cluster_info = pd.read_csv(cluster_info_file, delimiter=delimiter) + else: + # load properties from other tsv/csv files + all_property_files = [p for p in phy_folder.iterdir() if p.suffix in [".csv", ".tsv"]] + + cluster_info = None + for file in all_property_files: + if file.suffix == ".tsv": + delimiter = "\t" + else: + delimiter = "," + new_property = pd.read_csv(file, delimiter=delimiter) + + # Only merge files that contain a cluster_id column + # This prevents KeyError when extraneous files don't have cluster_id + # Typical aggregated files include cluster_group.tsv, cluster_info.tsv, cluster_KSLabel.tsv + # See Phy docs: https://phy.readthedocs.io/en/latest/sorting_user_guide/ + # See: https://github.com/SpikeInterface/spikeinterface/issues/4124 + if "cluster_id" not in new_property.columns: + continue + + if cluster_info is None: + cluster_info = new_property + else: + cluster_info = pd.merge(cluster_info, new_property, on="cluster_id", suffixes=[None, "_repeat"]) # in case no tsv/csv files are found populate cluster info with minimal info if cluster_info is None: @@ -268,6 +262,9 @@ class PhySortingExtractor(BasePhyKilosortSortingExtractor): Cluster groups to exclude (e.g. "noise" or ["noise", "mua"]). load_all_cluster_properties : bool, default: True If True, all cluster properties are loaded from the tsv/csv files. + load_templates_or_cluster : str, Literal["templates, "clusters"] , default: "clusters" + Defines whether to load templates (kilosort output) or clusters (after manual refinement + using phy). Returns ------- @@ -280,6 +277,7 @@ def __init__( folder_path: Path | str, exclude_cluster_groups: Optional[list[str] | str] = None, load_all_cluster_properties: bool = True, + load_templates_or_cluster: Literal["templates", "clusters"] = "clusters", ): BasePhyKilosortSortingExtractor.__init__( self, @@ -287,6 +285,7 @@ def __init__( exclude_cluster_groups, keep_good_only=False, load_all_cluster_properties=load_all_cluster_properties, + load_templates_or_cluster=load_templates_or_cluster, ) self._kwargs = { @@ -307,6 +306,9 @@ class KiloSortSortingExtractor(BasePhyKilosortSortingExtractor): If True, only Kilosort-labeled 'good' units are returned. remove_empty_units : bool, default: True If True, empty units are removed from the sorting extractor. + load_templates_or_cluster : str, Literal["templates, "clusters"] , default: "clusters" + Defines whether to load templates (kilosort output) or clusters (after manual refinement + using phy). Returns ------- @@ -314,13 +316,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, + load_templates_or_cluster: Literal["templates", "clusters"] = "clusters", + ): BasePhyKilosortSortingExtractor.__init__( self, folder_path, exclude_cluster_groups=None, keep_good_only=keep_good_only, remove_empty_units=remove_empty_units, + load_templates_or_cluster=load_templates_or_cluster, ) self._kwargs = {"folder_path": str(Path(folder_path).absolute()), "keep_good_only": keep_good_only}