Skip to content

Commit b0c2bae

Browse files
authored
Merge pull request #3443 from alejoe91/fix-save-as-recordingless
Allow to save recordingless analyzer as
2 parents 0ae32a3 + c338658 commit b0c2bae

2 files changed

Lines changed: 58 additions & 46 deletions

File tree

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 54 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
import shutil
1212
import warnings
1313
import importlib
14+
from copy import copy
1415
from packaging.version import parse
1516
from 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 (

src/spikeinterface/preprocessing/tests/test_filter.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ def test_causal_filter_main_kwargs(self, recording_and_data):
4646

4747
filt_data = causal_filter(recording, direction="forward", **options, margin_ms=0).get_traces()
4848

49-
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-4)
49+
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-2)
5050

5151
# Then, change all kwargs to ensure they are propagated
5252
# and check the backwards version.
@@ -66,7 +66,7 @@ def test_causal_filter_main_kwargs(self, recording_and_data):
6666

6767
filt_data = causal_filter(recording, direction="backward", **options, margin_ms=0).get_traces()
6868

69-
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-4)
69+
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-2)
7070

7171
def test_causal_filter_custom_coeff(self, recording_and_data):
7272
"""
@@ -89,7 +89,7 @@ def test_causal_filter_custom_coeff(self, recording_and_data):
8989

9090
filt_data = causal_filter(recording, direction="forward", **options, margin_ms=0).get_traces()
9191

92-
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-4, equal_nan=True)
92+
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-2, equal_nan=True)
9393

9494
# Next, in "sos" mode
9595
options["filter_mode"] = "sos"
@@ -100,7 +100,7 @@ def test_causal_filter_custom_coeff(self, recording_and_data):
100100

101101
filt_data = causal_filter(recording, direction="forward", **options, margin_ms=0).get_traces()
102102

103-
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-4, equal_nan=True)
103+
assert np.allclose(test_data, filt_data, rtol=0, atol=1e-2, equal_nan=True)
104104

105105
def test_causal_kwarg_error_raised(self, recording_and_data):
106106
"""

0 commit comments

Comments
 (0)