Skip to content

Commit 5a66c36

Browse files
committed
Fix unicode size bug, rename functions, and add get_sorting_property
1 parent 91c2a0d commit 5a66c36

3 files changed

Lines changed: 35 additions & 11 deletions

File tree

src/spikeinterface/core/base.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -145,7 +145,7 @@ def ids_to_indices(
145145
non_existent_ids = [id for id in ids if id not in self._main_ids]
146146
if non_existent_ids:
147147
error_msg = (
148-
f"IDs {non_existent_ids} are not channel ids of the extractor. \n"
148+
f"IDs {non_existent_ids} are not ids of the extractor. \n"
149149
f"Available ids are {self._main_ids} with dtype {self._main_ids.dtype}"
150150
)
151151
raise ValueError(error_msg)
@@ -293,6 +293,11 @@ def set_property(
293293
), f"Mismatch between existing property dtype {existing_property.kind} and provided values dtype {dtype_kind}."
294294

295295
indices = self.ids_to_indices(ids)
296+
if dtype_kind == "U":
297+
# re-adjust the size of the property
298+
existing_unicode_size = max(len(s) for s in self._properties[key])
299+
new_unicode_size = max(max(len(s) for s in values), existing_unicode_size)
300+
self._properties[key] = self._properties[key].astype(f"<U{new_unicode_size}")
296301
self._properties[key][indices] = values
297302
else:
298303
indices = self.ids_to_indices(ids)

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 20 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -743,7 +743,7 @@ def set_temporary_recording(self, recording: BaseRecording, check_dtype: bool =
743743
warnings.warn("SortingAnalyzer recording is already set. The current recording is temporarily replaced.")
744744
self._temporary_recording = recording
745745

746-
def set_unit_property(
746+
def set_sorting_property(
747747
self,
748748
key,
749749
values: list | np.ndarray | tuple,
@@ -793,6 +793,25 @@ def set_unit_property(
793793
# IMPORTANT: we need to re-consolidate the zarr store!
794794
zarr.consolidate_metadata(zarr_root.store)
795795

796+
def get_sorting_property(self, key: str, ids: Optional[Iterable] = None) -> np.ndarray:
797+
"""
798+
Get property vector for unit ids.
799+
800+
Parameters
801+
----------
802+
key : str
803+
The property name
804+
ids : list/np.array, default: None
805+
List of subset of ids to get the values.
806+
if None all the ids are returned
807+
808+
Returns
809+
-------
810+
values : np.array
811+
Array of values for the property
812+
"""
813+
return self.sorting.get_property(key, ids=ids)
814+
796815
def _save_or_select_or_merge(
797816
self,
798817
format="binary_folder",

src/spikeinterface/core/tests/test_sortinganalyzer.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -66,9 +66,9 @@ def test_SortingAnalyzer_memory(tmp_path, dataset):
6666
)
6767
assert not sorting_analyzer.return_scaled
6868

69-
# test set_unit_property
70-
sorting_analyzer.set_unit_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
71-
sorting_analyzer.set_unit_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
69+
# test set_sorting_property
70+
sorting_analyzer.set_sorting_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
71+
sorting_analyzer.set_sorting_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
7272
assert "quality" in sorting_analyzer.sorting.get_property_keys()
7373
assert "number" in sorting_analyzer.sorting.get_property_keys()
7474

@@ -109,9 +109,9 @@ def test_SortingAnalyzer_binary_folder(tmp_path, dataset):
109109
assert not sorting_analyzer.return_scaled
110110
_check_sorting_analyzers(sorting_analyzer, sorting, cache_folder=tmp_path)
111111

112-
# test set_unit_property
113-
sorting_analyzer.set_unit_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
114-
sorting_analyzer.set_unit_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
112+
# test set_sorting_property
113+
sorting_analyzer.set_sorting_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
114+
sorting_analyzer.set_sorting_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
115115
assert "quality" in sorting_analyzer.sorting.get_property_keys()
116116
assert "number" in sorting_analyzer.sorting.get_property_keys()
117117
sorting_analyzer_reloded = load_sorting_analyzer(folder, format="auto")
@@ -191,9 +191,9 @@ def test_SortingAnalyzer_zarr(tmp_path, dataset):
191191
== LZMA.codec_id
192192
)
193193

194-
# test set_unit_property
195-
sorting_analyzer.set_unit_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
196-
sorting_analyzer.set_unit_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
194+
# test set_sorting_property
195+
sorting_analyzer.set_sorting_property(key="quality", values=["good"] * len(sorting_analyzer.unit_ids))
196+
sorting_analyzer.set_sorting_property(key="number", values=np.arange(len(sorting_analyzer.unit_ids)))
197197
assert "quality" in sorting_analyzer.sorting.get_property_keys()
198198
assert "number" in sorting_analyzer.sorting.get_property_keys()
199199
sorting_analyzer_reloded = load_sorting_analyzer(sorting_analyzer.folder, format="auto")

0 commit comments

Comments
 (0)