Skip to content

Commit b234713

Browse files
authored
Merge branch 'main' into curation-pot-merges
2 parents 238e694 + bd67da8 commit b234713

7 files changed

Lines changed: 96 additions & 24 deletions

File tree

.github/workflows/all-tests.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -139,7 +139,7 @@ jobs:
139139
140140
- name: Test streaming extractors
141141
shell: bash
142-
if: env.RUN_STREAMING_EXTRACTORS_TESTS
142+
if: env.RUN_STREAMING_EXTRACTORS_TESTS == 'true'
143143
run: |
144144
pip install -e .[streaming_extractors,test_extractors]
145145
./.github/run_tests.sh "streaming_extractors" --no-virtual-env
@@ -202,7 +202,7 @@ jobs:
202202
shell: bash
203203
if: env.RUN_WIDGETS_TESTS == 'true'
204204
run: |
205-
pip install -e .[full]
205+
pip install -e .[full,widgets]
206206
./.github/run_tests.sh widgets --no-virtual-env
207207
208208
- name: Test exporters

doc/api.rst

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,19 @@ Low-level
7373

7474
.. autoclass:: ChunkRecordingExecutor
7575

76+
77+
Back-compatibility with ``WaveformExtractor`` (version < 0.101.0)
78+
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
79+
80+
.. automodule:: spikeinterface.core
81+
:noindex:
82+
83+
.. autofunction:: extract_waveforms
84+
.. autofunction:: load_waveforms
85+
.. autofunction:: load_sorting_analyzer_or_waveforms
86+
87+
88+
7689
spikeinterface.extractors
7790
-------------------------
7891

src/spikeinterface/core/__init__.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -166,4 +166,8 @@
166166

167167
# Important not for compatibility!!
168168
# This wil be uncommented after 0.100
169-
from .waveforms_extractor_backwards_compatibility import extract_waveforms, load_waveforms
169+
from .waveforms_extractor_backwards_compatibility import (
170+
extract_waveforms,
171+
load_waveforms,
172+
load_sorting_analyzer_or_waveforms,
173+
)

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -145,7 +145,7 @@ def create_sorting_analyzer(
145145
return sorting_analyzer
146146

147147

148-
def load_sorting_analyzer(folder, load_extensions=True, format="auto"):
148+
def load_sorting_analyzer(folder, load_extensions=True, format="auto", storage_options=None):
149149
"""
150150
Load a SortingAnalyzer object from disk.
151151
@@ -157,16 +157,17 @@ def load_sorting_analyzer(folder, load_extensions=True, format="auto"):
157157
Load all extensions or not.
158158
format : "auto" | "binary_folder" | "zarr"
159159
The format of the folder.
160+
storage_options : dict | None, default: None
161+
The storage options to specify credentials to remote zarr bucket.
162+
For open buckets, it doesn't need to be specified.
160163
161164
Returns
162165
-------
163166
sorting_analyzer : SortingAnalyzer
164167
The loaded SortingAnalyzer
165168
166169
"""
167-
if format == "zarr":
168-
folder = clean_zarr_folder_name(folder)
169-
return SortingAnalyzer.load(folder, load_extensions=load_extensions, format=format)
170+
return SortingAnalyzer.load(folder, load_extensions=load_extensions, format=format, storage_options=storage_options)
170171

171172

172173
class SortingAnalyzer:

src/spikeinterface/core/waveforms_extractor_backwards_compatibility.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -343,6 +343,31 @@ def get_template(
343343
return templates[0]
344344

345345

346+
def load_sorting_analyzer_or_waveforms(folder, sorting=None):
347+
"""
348+
Load a SortingAnalyzer from either a newly saved SortingAnalyzer folder or an old WaveformExtractor folder.
349+
350+
Parameters
351+
----------
352+
folder: str | Path
353+
The folder to the sorting analyzer or waveform extractor
354+
sorting: BaseSorting | None, default: None
355+
The sorting object to instantiate with the SortingAnalyzer (only used for old WaveformExtractor)
356+
357+
Returns
358+
-------
359+
sorting_analyzer: SortingAnalyzer
360+
The returned SortingAnalyzer.
361+
"""
362+
folder = Path(folder)
363+
if folder.suffix == ".zarr":
364+
return load_sorting_analyzer(folder)
365+
elif (folder / "spikeinterface_info.json").exists():
366+
return load_sorting_analyzer(folder)
367+
else:
368+
return load_waveforms(folder, sorting=sorting, output="SortingAnalyzer")
369+
370+
346371
def load_waveforms(
347372
folder,
348373
with_recording: bool = True,

src/spikeinterface/widgets/tests/test_widgets.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -539,6 +539,21 @@ def test_plot_sorting_summary(self):
539539
backend=backend,
540540
**self.backend_kwargs[backend],
541541
)
542+
# add unit_properties
543+
sw.plot_sorting_summary(
544+
self.sorting_analyzer_sparse,
545+
unit_table_properties=["firing_rate", "snr"],
546+
backend=backend,
547+
**self.backend_kwargs[backend],
548+
)
549+
# adding a missing property should raise a warning
550+
with self.assertWarns(UserWarning):
551+
sw.plot_sorting_summary(
552+
self.sorting_analyzer_sparse,
553+
unit_table_properties=["missing_property"],
554+
backend=backend,
555+
**self.backend_kwargs[backend],
556+
)
542557

543558
def test_plot_agreement_matrix(self):
544559
possible_backends = list(sw.AgreementMatrixWidget.get_possible_backends())

src/spikeinterface/widgets/utils_sortingview.py

Lines changed: 31 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
import numpy as np
44

5+
from ..core import SortingAnalyzer, BaseSorting
56
from ..core.core_tools import check_json
67
from warnings import warn
78

@@ -46,26 +47,42 @@ def handle_display_and_url(widget, view, **backend_kwargs):
4647
return url
4748

4849

49-
def generate_unit_table_view(analyzer, unit_properties=None, similarity_scores=None):
50+
def generate_unit_table_view(
51+
sorting_or_sorting_analyzer: SortingAnalyzer | BaseSorting,
52+
unit_properties: list[str] | None = None,
53+
similarity_scores: npndarray | None = None,
54+
):
5055
import sortingview.views as vv
5156

52-
sorting = analyzer.sorting
57+
if isinstance(sorting_or_sorting_analyzer, SortingAnalyzer):
58+
analyzer = sorting_or_sorting_analyzer
59+
sorting = analyzer.sorting
60+
else:
61+
sorting = sorting_or_sorting_analyzer
62+
analyzer = None
5363

5464
# Find available unit properties from all sources
5565
sorting_props = list(sorting.get_property_keys())
56-
if analyzer.get_extension("quality_metrics") is not None:
57-
qm_props = list(analyzer.get_extension("quality_metrics").get_data().columns)
58-
qm_data = analyzer.get_extension("quality_metrics").get_data()
66+
if analyzer is not None:
67+
if analyzer.get_extension("quality_metrics") is not None:
68+
qm_props = list(analyzer.get_extension("quality_metrics").get_data().columns)
69+
qm_data = analyzer.get_extension("quality_metrics").get_data()
70+
else:
71+
qm_props = []
72+
if analyzer.get_extension("template_metrics") is not None:
73+
tm_props = list(analyzer.get_extension("template_metrics").get_data().columns)
74+
tm_data = analyzer.get_extension("template_metrics").get_data()
75+
else:
76+
tm_props = []
77+
# Check for any overlaps and warn user if any
78+
all_props = sorting_props + qm_props + tm_props
5979
else:
80+
all_props = sorting_props
6081
qm_props = []
61-
if analyzer.get_extension("template_metrics") is not None:
62-
tm_props = list(analyzer.get_extension("template_metrics").get_data().columns)
63-
tm_data = analyzer.get_extension("template_metrics").get_data()
64-
else:
6582
tm_props = []
83+
qm_data = None
84+
tm_data = None
6685

67-
# Check for any overlaps and warn user if any
68-
all_props = sorting_props + qm_props + tm_props
6986
overlap_props = [prop for prop in all_props if all_props.count(prop) > 1]
7087
if len(overlap_props) > 0:
7188
warn(
@@ -93,7 +110,8 @@ def generate_unit_table_view(analyzer, unit_properties=None, similarity_scores=N
93110
elif prop_name in tm_props:
94111
property_values = tm_data[prop_name].values
95112
else:
96-
raise ValueError(f"Property '{prop_name}' not found in sorting, quality_metrics, or template_metrics")
113+
warn(f"Property '{prop_name}' not found in sorting, quality_metrics, or template_metrics")
114+
continue
97115

98116
# make dtype available
99117
val0 = np.array(property_values[0])
@@ -106,7 +124,7 @@ def generate_unit_table_view(analyzer, unit_properties=None, similarity_scores=N
106124
elif val0.dtype.kind == "b":
107125
dtype = "bool"
108126
else:
109-
print(f"Unsupported dtype {val0.dtype} for property {prop_name}. Skipping")
127+
warn(f"Unsupported dtype {val0.dtype} for property {prop_name}. Skipping")
110128
continue
111129
ut_columns.append(vv.UnitsTableColumn(key=prop_name, label=prop_name, dtype=dtype))
112130
valid_unit_properties.append(prop_name)
@@ -122,10 +140,6 @@ def generate_unit_table_view(analyzer, unit_properties=None, similarity_scores=N
122140
property_values = qm_data[prop_name].values
123141
elif prop_name in tm_props:
124142
property_values = tm_data[prop_name].values
125-
else:
126-
raise ValueError(
127-
f"Property '{prop_name}' not found in sorting, quality_metrics, or template_metrics"
128-
)
129143

130144
# Check for NaN values
131145
val0 = np.array(property_values[0])

0 commit comments

Comments
 (0)