11import warnings
2+ from pathlib import Path
23from packaging import version
34
4- from spikeinterface .core import write_binary_recording
5+ import numpy as np
6+
7+ from spikeinterface .core import write_binary_recording , Motion , BaseRecording
58from spikeinterface .sorters .basesorter import BaseSorter , get_job_kwargs
69from .kilosortbase import KilosortBase
710from spikeinterface .sorters .basesorter import get_job_kwargs
@@ -77,7 +80,7 @@ def is_installed(cls):
7780
7881 @classmethod
7982 def get_sorter_version (cls ):
80- """kilosort.__version__ <4.0.10 is always '4'"""
83+ """kilosort.__version__ < 4.0.10 is always '4'"""
8184 return importlib_version ("kilosort" )
8285
8386 @classmethod
@@ -122,12 +125,11 @@ def initialize_folder(cls, recording, output_folder, verbose, remove_existing_fo
122125 @classmethod
123126 def check_sorter_version (cls ):
124127 kilosort_version = version .parse (cls .get_sorter_version ())
125- if kilosort_version < version .parse ("4.0.16" ):
126- raise Exception (
127- f"""SpikeInterface only supports kilosort versions 4.0.16 and above. You are running version { kilosort_version } . To install the latest version, run:
128- >>> pip install kilosort --upgrade
129- """
130- )
128+ if kilosort_version < version .parse ("4.1.1" ):
129+ raise Exception (f"""SpikeInterface only supports kilosort versions 4.1.1 and above (which support numpy>=2).
130+ You are running version { kilosort_version } . To install the latest version, run:
131+ >>> pip install kilosort --upgrade
132+ """ )
131133
132134 @classmethod
133135 def _setup_recording (cls , recording , sorter_output_folder , params , verbose ):
@@ -177,7 +179,6 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
177179
178180 import time
179181 import torch
180- import numpy as np
181182 import logging
182183
183184 if version .parse (cls .get_sorter_version ()) < version .parse ("4.0.16" ):
@@ -461,6 +462,11 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
461462 if (sorter_output_folder / "recording.dat" ).is_file ():
462463 (sorter_output_folder / "recording.dat" ).unlink ()
463464
465+ # close logger
466+ for handler in logger .handlers .copy ():
467+ logger .removeHandler (handler )
468+ handler .close ()
469+
464470 @classmethod
465471 def _get_result_from_folder (cls , sorter_output_folder ):
466472 return KilosortBase ._get_result_from_folder (sorter_output_folder )
@@ -469,7 +475,6 @@ def _get_result_from_folder(cls, sorter_output_folder):
469475 def _setup_json_probe_map (cls , recording , sorter_output_folder ):
470476 """Create a JSON probe map file for Kilosort4."""
471477 from kilosort .io import save_probe
472- import numpy as np
473478
474479 groups = recording .get_channel_groups ()
475480 positions = np .array (recording .get_channel_locations ())
@@ -492,3 +497,47 @@ def _setup_json_probe_map(cls, recording, sorter_output_folder):
492497 "n_chan" : n_chan ,
493498 }
494499 save_probe (probe , str (sorter_output_folder / "chanMap.json" ))
500+
501+
502+ def read_kilosort4_motion (sorter_output_folder : str | Path , recording : BaseRecording | None = None ) -> Motion :
503+ """Reads the motion information from a Kilosort4 output folder and returns a Motion object.
504+
505+ Parameters
506+ ----------
507+ sorter_output_folder: str or Path
508+ The path to the Kilosort4 output folder.
509+ recording: BaseRecording, optional
510+ The recording object. If provided, the temporal bins will be estimated based on the recording's
511+ start and end times. If not provided, the temporal bins will be estimated based on the number
512+ of batches in the ops file.
513+
514+ Returns
515+ -------
516+ Motion
517+ A Motion object containing the displacement, temporal bins, and spatial bins.
518+
519+ """
520+ sorter_output_folder = Path (sorter_output_folder )
521+ ops_file = sorter_output_folder / "ops.npy"
522+ if not ops_file .is_file ():
523+ raise FileNotFoundError ("'ops.npy' file not found!" )
524+ ops = np .load (ops_file , allow_pickle = True ).item ()
525+ yblk = ops .get ("yblk" )
526+ dshift = ops .get ("dshift" )
527+ if yblk is None or dshift is None :
528+ raise Exception ("'yblk' and 'dshift' fields not found in ops file!" )
529+ displacement = dshift
530+ spatial_bins_um = yblk
531+ # estimate temporal bins
532+ batch_size = ops ["batch_size" ]
533+ fs = ops ["fs" ]
534+ t_bin = batch_size / fs
535+ if recording is not None :
536+ t_start = recording .get_start_time ()
537+ t_end = recording .get_end_time ()
538+ temporal_bins_s = np .linspace (t_start + t_bin / 2 , t_end - t_bin / 2 , displacement .shape [0 ])
539+ else :
540+ temporal_bins_s = np .arange (displacement .shape [0 ]) * t_bin + t_bin / 2
541+
542+ motion = Motion (displacement = displacement , temporal_bins_s = temporal_bins_s , spatial_bins_um = spatial_bins_um )
543+ return motion
0 commit comments