2323
2424from .base import load_extractor
2525from .recording_tools import check_probe_do_not_overlap , get_rec_attributes , do_recording_attributes_match
26- from .core_tools import check_json , retrieve_importing_provenance , is_path_remote
26+ from .core_tools import check_json , retrieve_importing_provenance , is_path_remote , clean_zarr_folder_name
2727from .sorting_tools import generate_unit_ids_for_merge_group , _get_ids_after_merging
2828from .job_tools import split_job_kwargs
2929from .numpyextractors import NumpySorting
@@ -111,6 +111,8 @@ def create_sorting_analyzer(
111111 sparsity off (or give external sparsity) like this.
112112 """
113113 if format != "memory" :
114+ if format == "zarr" :
115+ folder = clean_zarr_folder_name (folder )
114116 if Path (folder ).is_dir ():
115117 if not overwrite :
116118 raise ValueError (f"Folder already exists { folder } ! Use overwrite=True to overwrite it." )
@@ -162,6 +164,8 @@ def load_sorting_analyzer(folder, load_extensions=True, format="auto"):
162164 The loaded SortingAnalyzer
163165
164166 """
167+ if format == "zarr" :
168+ folder = clean_zarr_folder_name (folder )
165169 return SortingAnalyzer .load (folder , load_extensions = load_extensions , format = format )
166170
167171
@@ -269,6 +273,8 @@ def create(
269273 sorting_analyzer = cls .load_from_binary_folder (folder , recording = recording )
270274 sorting_analyzer .folder = Path (folder )
271275 elif format == "zarr" :
276+ assert folder is not None , "For format='zarr' folder must be provided"
277+ folder = clean_zarr_folder_name (folder )
272278 cls .create_zarr (folder , sorting , recording , sparsity , return_scaled , rec_attributes = None )
273279 sorting_analyzer = cls .load_from_zarr (folder , recording = recording )
274280 sorting_analyzer .folder = Path (folder )
@@ -487,10 +493,7 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
487493 import zarr
488494 import numcodecs
489495
490- folder = Path (folder )
491- # force zarr sufix
492- if folder .suffix != ".zarr" :
493- folder = folder .parent / f"{ folder .stem } .zarr"
496+ folder = clean_zarr_folder_name (folder )
494497
495498 if folder .is_dir ():
496499 raise ValueError (f"Folder already exists { folder } " )
@@ -613,7 +616,7 @@ def load_from_zarr(cls, folder, recording=None, storage_options=None):
613616
614617 return sorting_analyzer
615618
616- def set_temporary_recording (self , recording : BaseRecording ):
619+ def set_temporary_recording (self , recording : BaseRecording , check_dtype : bool = True ):
617620 """
618621 Sets a temporary recording object. This function can be useful to temporarily set
619622 a "cached" recording object that is not saved in the SortingAnalyzer object to speed up
@@ -625,12 +628,17 @@ def set_temporary_recording(self, recording: BaseRecording):
625628 ----------
626629 recording : BaseRecording
627630 The recording object to set as temporary recording.
631+ check_dtype : bool, default: True
632+ If True, check that the dtype of the temporary recording is the same as the original recording.
628633 """
629634 # check that recording is compatible
630- assert do_recording_attributes_match (recording , self .rec_attributes ), "Recording attributes do not match."
631- assert np .array_equal (
632- recording .get_channel_locations (), self .get_channel_locations ()
633- ), "Recording channel locations do not match."
635+ attributes_match , exception_str = do_recording_attributes_match (
636+ recording , self .rec_attributes , check_dtype = check_dtype
637+ )
638+ if not attributes_match :
639+ raise ValueError (exception_str )
640+ if not np .array_equal (recording .get_channel_locations (), self .get_channel_locations ()):
641+ raise ValueError ("Recording channel locations do not match." )
634642 if self ._recording is not None :
635643 warnings .warn ("SortingAnalyzer recording is already set. The current recording is temporarily replaced." )
636644 self ._temporary_recording = recording
@@ -768,9 +776,7 @@ def _save_or_select_or_merge(
768776
769777 elif format == "zarr" :
770778 assert folder is not None , "For format='zarr' folder must be provided"
771- folder = Path (folder )
772- if folder .suffix != ".zarr" :
773- folder = folder .parent / f"{ folder .stem } .zarr"
779+ folder = clean_zarr_folder_name (folder )
774780 SortingAnalyzer .create_zarr (
775781 folder , sorting_provenance , recording , sparsity , self .return_scaled , self .rec_attributes
776782 )
@@ -829,6 +835,8 @@ def save_as(self, format="memory", folder=None) -> "SortingAnalyzer":
829835 format : "memory" | "binary_folder" | "zarr", default: "memory"
830836 The new backend format to use
831837 """
838+ if format == "zarr" :
839+ folder = clean_zarr_folder_name (folder )
832840 return self ._save_or_select_or_merge (format = format , folder = folder )
833841
834842 def select_units (self , unit_ids , format = "memory" , folder = None ) -> "SortingAnalyzer" :
@@ -854,6 +862,8 @@ def select_units(self, unit_ids, format="memory", folder=None) -> "SortingAnalyz
854862 The newly create sorting_analyzer with the selected units
855863 """
856864 # TODO check that unit_ids are in same order otherwise many extension do handle it properly!!!!
865+ if format == "zarr" :
866+ folder = clean_zarr_folder_name (folder )
857867 return self ._save_or_select_or_merge (format = format , folder = folder , unit_ids = unit_ids )
858868
859869 def remove_units (self , remove_unit_ids , format = "memory" , folder = None ) -> "SortingAnalyzer" :
@@ -880,6 +890,8 @@ def remove_units(self, remove_unit_ids, format="memory", folder=None) -> "Sortin
880890 """
881891 # TODO check that unit_ids are in same order otherwise many extension do handle it properly!!!!
882892 unit_ids = self .unit_ids [~ np .isin (self .unit_ids , remove_unit_ids )]
893+ if format == "zarr" :
894+ folder = clean_zarr_folder_name (folder )
883895 return self ._save_or_select_or_merge (format = format , folder = folder , unit_ids = unit_ids )
884896
885897 def merge_units (
@@ -938,6 +950,9 @@ def merge_units(
938950 The newly create `SortingAnalyzer` with the selected units
939951 """
940952
953+ if format == "zarr" :
954+ folder = clean_zarr_folder_name (folder )
955+
941956 assert merging_mode in ["soft" , "hard" ], "Merging mode should be either soft or hard"
942957
943958 if len (merge_unit_groups ) == 0 :
@@ -1016,6 +1031,9 @@ def has_temporary_recording(self) -> bool:
10161031 def is_sparse (self ) -> bool :
10171032 return self .sparsity is not None
10181033
1034+ def is_filtered (self ) -> bool :
1035+ return self .rec_attributes ["is_filtered" ]
1036+
10191037 def get_sorting_provenance (self ):
10201038 """
10211039 Get the original sorting if possible otherwise return None
0 commit comments