1111import shutil
1212import warnings
1313import importlib
14+ from copy import copy
1415from packaging .version import parse
1516from time import perf_counter
1617
@@ -254,6 +255,7 @@ def create(
254255 sparsity = None ,
255256 return_scaled = True ,
256257 ):
258+ assert recording is not None , "To create a SortingAnalyzer you need to specify the recording"
257259 # some checks
258260 if sorting .sampling_frequency != recording .sampling_frequency :
259261 if math .isclose (sorting .sampling_frequency , recording .sampling_frequency , abs_tol = 1e-2 , rel_tol = 1e-5 ):
@@ -352,8 +354,6 @@ def create_memory(cls, sorting, recording, sparsity, return_scaled, rec_attribut
352354 def create_binary_folder (cls , folder , sorting , recording , sparsity , return_scaled , rec_attributes ):
353355 # used by create and save_as
354356
355- assert recording is not None , "To create a SortingAnalyzer you need to specify the recording"
356-
357357 folder = Path (folder )
358358 if folder .is_dir ():
359359 raise ValueError (f"Folder already exists { folder } " )
@@ -369,26 +369,34 @@ def create_binary_folder(cls, folder, sorting, recording, sparsity, return_scale
369369 json .dump (check_json (info ), f , indent = 4 )
370370
371371 # save a copy of the sorting
372- # NumpyFolderSorting.write_sorting(sorting, folder / "sorting")
373372 sorting .save (folder = folder / "sorting" )
374373
375- # save recording and sorting provenance
376- if recording .check_serializability ("json" ):
377- recording .dump (folder / "recording.json" , relative_to = folder )
378- elif recording .check_serializability ("pickle" ):
379- recording .dump (folder / "recording.pickle" , relative_to = folder )
374+ if recording is not None :
375+ # save recording and sorting provenance
376+ if recording .check_serializability ("json" ):
377+ recording .dump (folder / "recording.json" , relative_to = folder )
378+ elif recording .check_serializability ("pickle" ):
379+ recording .dump (folder / "recording.pickle" , relative_to = folder )
380+ else :
381+ warnings .warn ("The Recording is not serializable! The recording link will be lost for future load" )
382+ else :
383+ assert rec_attributes is not None , "recording or rec_attributes must be provided"
384+ warnings .warn ("Recording not provided, instntiating SortingAnalyzer in recordingless mode." )
380385
381386 if sorting .check_serializability ("json" ):
382387 sorting .dump (folder / "sorting_provenance.json" , relative_to = folder )
383388 elif sorting .check_serializability ("pickle" ):
384389 sorting .dump (folder / "sorting_provenance.pickle" , relative_to = folder )
390+ else :
391+ warnings .warn (
392+ "The sorting provenance is not serializable! The sorting provenance link will be lost for future load"
393+ )
385394
386395 # dump recording attributes
387396 probegroup = None
388397 rec_attributes_file = folder / "recording_info" / "recording_attributes.json"
389398 rec_attributes_file .parent .mkdir ()
390399 if rec_attributes is None :
391- assert recording is not None
392400 rec_attributes = get_rec_attributes (recording )
393401 rec_attributes_file .write_text (json .dumps (check_json (rec_attributes ), indent = 4 ), encoding = "utf8" )
394402 probegroup = recording .get_probegroup ()
@@ -519,20 +527,21 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
519527 zarr_root .attrs ["settings" ] = check_json (settings )
520528
521529 # the recording
522- rec_dict = recording .to_dict (relative_to = folder , recursive = True )
523-
524- if recording .check_serializability ("json" ):
525- # zarr_root.create_dataset("recording", data=rec_dict, object_codec=numcodecs.JSON())
526- zarr_rec = np .array ([check_json (rec_dict )], dtype = object )
527- zarr_root .create_dataset ("recording" , data = zarr_rec , object_codec = numcodecs .JSON ())
528- elif recording .check_serializability ("pickle" ):
529- # zarr_root.create_dataset("recording", data=rec_dict, object_codec=numcodecs.Pickle())
530- zarr_rec = np .array ([rec_dict ], dtype = object )
531- zarr_root .create_dataset ("recording" , data = zarr_rec , object_codec = numcodecs .Pickle ())
530+ if recording is not None :
531+ rec_dict = recording .to_dict (relative_to = folder , recursive = True )
532+ if recording .check_serializability ("json" ):
533+ # zarr_root.create_dataset("recording", data=rec_dict, object_codec=numcodecs.JSON())
534+ zarr_rec = np .array ([check_json (rec_dict )], dtype = object )
535+ zarr_root .create_dataset ("recording" , data = zarr_rec , object_codec = numcodecs .JSON ())
536+ elif recording .check_serializability ("pickle" ):
537+ # zarr_root.create_dataset("recording", data=rec_dict, object_codec=numcodecs.Pickle())
538+ zarr_rec = np .array ([rec_dict ], dtype = object )
539+ zarr_root .create_dataset ("recording" , data = zarr_rec , object_codec = numcodecs .Pickle ())
540+ else :
541+ warnings .warn ("The Recording is not serializable! The recording link will be lost for future load" )
532542 else :
533- warnings .warn (
534- "SortingAnalyzer with zarr : the Recording is not json serializable, the recording link will be lost for future load"
535- )
543+ assert rec_attributes is not None , "recording or rec_attributes must be provided"
544+ warnings .warn ("Recording not provided, instntiating SortingAnalyzer in recordingless mode." )
536545
537546 # sorting provenance
538547 sort_dict = sorting .to_dict (relative_to = folder , recursive = True )
@@ -542,14 +551,14 @@ def create_zarr(cls, folder, sorting, recording, sparsity, return_scaled, rec_at
542551 elif sorting .check_serializability ("pickle" ):
543552 zarr_sort = np .array ([sort_dict ], dtype = object )
544553 zarr_root .create_dataset ("sorting_provenance" , data = zarr_sort , object_codec = numcodecs .Pickle ())
545-
546- # else:
547- # warnings.warn("SortingAnalyzer with zarr : the sorting provenance is not json serializable, the sorting provenance link will be lost for futur load")
554+ else :
555+ warnings .warn (
556+ "The sorting provenance is not serializable! The sorting provenance link will be lost for future load"
557+ )
548558
549559 recording_info = zarr_root .create_group ("recording_info" )
550560
551561 if rec_attributes is None :
552- assert recording is not None
553562 rec_attributes = get_rec_attributes (recording )
554563 probegroup = recording .get_probegroup ()
555564 else :
@@ -605,11 +614,13 @@ def load_from_zarr(cls, folder, recording=None, storage_options=None):
605614
606615 # load recording if possible
607616 if recording is None :
608- rec_dict = zarr_root ["recording" ][0 ]
609- try :
610- recording = load_extractor (rec_dict , base_folder = folder )
611- except :
612- recording = None
617+ rec_field = zarr_root .get ("recording" )
618+ if rec_field is not None :
619+ rec_dict = rec_field [0 ]
620+ try :
621+ recording = load_extractor (rec_dict , base_folder = folder )
622+ except :
623+ recording = None
613624 else :
614625 # TODO maybe maybe not??? : do we need to check attributes match internal rec_attributes
615626 # Note this will make the loading too slow
@@ -2015,7 +2026,7 @@ def copy(self, new_sorting_analyzer, unit_ids=None):
20152026 new_extension .data = self .data
20162027 else :
20172028 new_extension .data = self ._select_extension_data (unit_ids )
2018- new_extension .run_info = self .run_info . copy ( )
2029+ new_extension .run_info = copy ( self .run_info )
20192030 new_extension .save ()
20202031 return new_extension
20212032
@@ -2033,7 +2044,7 @@ def merge(
20332044 new_extension .data = self ._merge_extension_data (
20342045 merge_unit_groups , new_unit_ids , new_sorting_analyzer , keep_mask , verbose = verbose , ** job_kwargs
20352046 )
2036- new_extension .run_info = self .run_info . copy ( )
2047+ new_extension .run_info = copy ( self .run_info )
20372048 new_extension .save ()
20382049 return new_extension
20392050
@@ -2251,15 +2262,16 @@ def _save_importing_provenance(self):
22512262 extension_group .attrs ["info" ] = info
22522263
22532264 def _save_run_info (self ):
2254- run_info = self .run_info .copy ()
2255-
2256- if self .format == "binary_folder" :
2257- extension_folder = self ._get_binary_extension_folder ()
2258- run_info_file = extension_folder / "run_info.json"
2259- run_info_file .write_text (json .dumps (run_info , indent = 4 ), encoding = "utf8" )
2260- elif self .format == "zarr" :
2261- extension_group = self ._get_zarr_extension_group (mode = "r+" )
2262- extension_group .attrs ["run_info" ] = run_info
2265+ if self .run_info is not None :
2266+ run_info = self .run_info .copy ()
2267+
2268+ if self .format == "binary_folder" :
2269+ extension_folder = self ._get_binary_extension_folder ()
2270+ run_info_file = extension_folder / "run_info.json"
2271+ run_info_file .write_text (json .dumps (run_info , indent = 4 ), encoding = "utf8" )
2272+ elif self .format == "zarr" :
2273+ extension_group = self ._get_zarr_extension_group (mode = "r+" )
2274+ extension_group .attrs ["run_info" ] = run_info
22632275
22642276 def get_pipeline_nodes (self ):
22652277 assert (
0 commit comments