Skip to content

Commit 01dd054

Browse files
committed
Create SortingAnalyzer directly in create_*
1 parent 60aac0d commit 01dd054

1 file changed

Lines changed: 34 additions & 6 deletions

File tree

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 34 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -414,8 +414,8 @@ def __init__(
414414
self,
415415
sorting: BaseSorting,
416416
recording: BaseRecording | None = None,
417+
format: Literal["memory", "binary_folder", "zarr"] = "memory",
417418
rec_attributes: dict | None = None,
418-
format: str | None = None,
419419
sparsity: ChannelSparsity | None = None,
420420
return_in_uV: bool = True,
421421
peak_sign: PeakSignType = "both",
@@ -697,7 +697,7 @@ def create_binary_folder(
697697
json.dump(check_json(settings), f, indent=4)
698698

699699
# Save the sorting output
700-
sorting.save(folder=sorting_folder)
700+
sorting = sorting.save(folder=sorting_folder)
701701

702702
# Dump sorting provenance
703703
if sorting.check_serializability("json"):
@@ -740,7 +740,21 @@ def create_binary_folder(
740740
if probegroup is not None:
741741
probeinterface.write_probeinterface(probegroup_file, probegroup)
742742

743-
return cls.load_from_binary_folder(folder, recording=recording, backend_options=backend_options)
743+
# Create SortingAnalyzer
744+
sorting_analyzer = SortingAnalyzer(
745+
sorting=sorting,
746+
recording=recording,
747+
rec_attributes={**rec_attributes_to_save, 'probegroup': probegroup},
748+
format="binary_folder",
749+
sparsity=sparsity,
750+
return_in_uV=return_in_uV,
751+
peak_sign=peak_sign,
752+
peak_mode=peak_mode,
753+
backend_options=backend_options,
754+
)
755+
sorting_analyzer.folder = folder
756+
757+
return sorting_analyzer
744758

745759
@classmethod
746760
def _handle_backward_compatibility_settings_pre_init(cls, settings: dict[str, Any]):
@@ -1063,7 +1077,21 @@ def create_zarr(
10631077
# Consolidate metadata (for faster reads)
10641078
zarr.consolidate_metadata(zarr_root.store)
10651079

1066-
return cls.load_from_zarr(folder, recording=recording, backend_options=backend_options)
1080+
# Create SortingAnalyzer
1081+
sorting_analyzer = SortingAnalyzer(
1082+
sorting=NumpySorting.from_sorting(sorting, with_metadata=True, copy_spike_vector=True),
1083+
recording=recording,
1084+
rec_attributes={**rec_attributes_to_save, 'probegroup': probegroup},
1085+
format="zarr",
1086+
sparsity=sparsity,
1087+
return_in_uV=return_in_uV,
1088+
peak_sign=peak_sign,
1089+
peak_mode=peak_mode,
1090+
backend_options=backend_options,
1091+
)
1092+
sorting_analyzer.folder = folder
1093+
1094+
return sorting_analyzer
10671095

10681096
@classmethod
10691097
def load_from_zarr(
@@ -1095,9 +1123,9 @@ def load_from_zarr(
10951123
"This may lead to unexpected behavior in loading extensions. "
10961124
"Consider re-generating the SortingAnalyzer object."
10971125
)
1098-
1126+
10991127
# Check all required inputs exist
1100-
if (zarr_root.attrs.get("settings") is None or zarr_root.get("sorting") is None
1128+
if (zarr_root.attrs.get("settings") is None or zarr_root.get("sorting") is None
11011129
or zarr_root.get("recording_info") is None
11021130
or zarr_root["recording_info"].attrs.get("recording_attributes") is None
11031131
):

0 commit comments

Comments
 (0)