Skip to content

Commit 9f2a992

Browse files
authored
Merge branch 'main' into rt-sort
2 parents 996f597 + 64d253c commit 9f2a992

43 files changed

Lines changed: 1248 additions & 333 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.github/scripts/test_kilosort4_ci.py

Lines changed: 19 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@
2020

2121
import pytest
2222
import copy
23-
from typing import Any
23+
from packaging.version import parse
2424
from inspect import signature
2525

2626
import numpy as np
@@ -84,8 +84,6 @@
8484
"duplicate_spike_ms": 0.3,
8585
}
8686

87-
PARAMS_TO_TEST = list(PARAMS_TO_TEST_DICT.keys())
88-
8987
PARAMETERS_NOT_AFFECTING_RESULTS = [
9088
"artifact_threshold",
9189
"ccg_threshold",
@@ -95,13 +93,24 @@
9593
"duplicate_spike_ms", # this is because ground-truth spikes don't have violations
9694
]
9795

98-
# THIS IS A PLACEHOLDER FOR FUTURE PARAMS TO TEST
99-
# if parse(version("kilosort")) >= parse("4.0.X"):
100-
# PARAMS_TO_TEST_DICT.update(
101-
# [
102-
# {"new_param": new_value},
103-
# ]
104-
# )
96+
97+
# Add/Remove version specific parameters
98+
if parse(kilosort.__version__) >= parse("4.0.22"):
99+
PARAMS_TO_TEST_DICT.update(
100+
{"position_limit": 50}
101+
)
102+
# Position limit only affects computing spike locations after sorting
103+
PARAMETERS_NOT_AFFECTING_RESULTS.append("position_limit")
104+
105+
if parse(kilosort.__version__) >= parse("4.0.24"):
106+
PARAMS_TO_TEST_DICT.update(
107+
{"max_peels": 200},
108+
)
109+
# max_peels is not affecting the results in this short dataset
110+
PARAMETERS_NOT_AFFECTING_RESULTS.append("max_peels")
111+
112+
113+
PARAMS_TO_TEST = list(PARAMS_TO_TEST_DICT.keys())
105114

106115

107116
class TestKilosort4Long:

doc/releases/0.102.1.rst

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
.. _release0.102.1:
2+
3+
SpikeInterface 0.102.1 release notes
4+
------------------------------------
5+
6+
19th February 2025
7+
8+
Minor release with bug fixes
9+
10+
core:
11+
12+
* Add parents in HTML representation and always print class name (#3700)
13+
* Add `SortingAnalyzer.set_sorting_property()/get_sorting_property()` functions (#3694)
14+
* Add `_parent` to `select_segment` classes (#3692)
15+
* Fix `chunk_size_limit` in `get_random_recording_slices` (#3691)
16+
* Fix bug in `super_zarr_open` (#3686)
17+
* Fix `si.load` for `WaveformExtarctor` (#3680)
18+
19+
sorters:
20+
21+
* Add `BaseSorter._dynamic_params()` to be to retrieve parameters dynamically (#3697)
22+
23+
continuous integration:
24+
25+
* Add `torch` to test installation (#3706)
26+
27+
testing:
28+
29+
* Test Python 3.13 (#3683)
30+
31+
Contributors:
32+
33+
* @alejoe91

doc/whatisnew.rst

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ Release notes
88
.. toctree::
99
:maxdepth: 1
1010

11+
releases/0.102.1.rst
1112
releases/0.102.0.rst
1213
releases/0.101.2.rst
1314
releases/0.101.1.rst
@@ -46,6 +47,21 @@ Release notes
4647
releases/0.9.1.rst
4748

4849

50+
Version 0.102.1
51+
===============
52+
53+
* Minor release with bug fixes
54+
55+
Version 0.102.0
56+
===============
57+
58+
* Added auto-label functions in curation module (#2918)
59+
* Refactored and improved auto-merge functions in curation module (#3435, #3601)
60+
* Added `spikeinterface.load()` function to load any SpikeInterface object (#3613, #3651)
61+
* Improved handling of time in base recording (#3509, #3623)
62+
* Multi-segment handling of motion interpolation (#3659)
63+
* Support for Numpy 2.0 and Zarr<3.0 (#3481,#3598)
64+
4965
Version 0.101.2
5066
===============
5167

pyproject.toml

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[project]
22
name = "spikeinterface"
3-
version = "0.102.1"
3+
version = "0.102.2"
44
authors = [
55
{ name="Alessio Buccino", email="alessiop.buccino@gmail.com" },
66
{ name="Samuel Garcia", email="sam.garcia.die@gmail.com" },
@@ -159,12 +159,16 @@ test = [
159159
"s3fs",
160160

161161
# tridesclous
162-
"numba",
162+
"numba<0.61.0;python_version<'3.13'",
163+
"numba>=0.61.0;python_version>='3.13'",
163164
"hdbscan>=0.8.33", # Previous version had a broken wheel
164165

165166
# for sortingview backend
166167
"sortingview",
167168

169+
# for motion and sortingcomponents
170+
"torch",
171+
168172
# curation
169173
"skops",
170174
"huggingface_hub",

src/spikeinterface/benchmark/benchmark_base.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -228,11 +228,11 @@ def run(self, case_keys=None, keep=True, verbose=False, **job_kwargs):
228228
benchmark.result["run_time"] = float(t1 - t0)
229229
benchmark.save_main(bench_folder)
230230

231-
def set_colors(self, colors=None, map_name="tab20"):
231+
def set_colors(self, colors=None, map_name="tab10"):
232232
if colors is None:
233233
case_keys = list(self.cases.keys())
234234
self.colors = get_some_colors(
235-
case_keys, map_name=map_name, color_engine="matplotlib", shuffle=False, margin=0
235+
case_keys, map_name=map_name, color_engine="matplotlib", shuffle=False, margin=0, resample=False
236236
)
237237
else:
238238
self.colors = colors

src/spikeinterface/benchmark/benchmark_matching.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -138,7 +138,7 @@ def plot_unit_counts(self, case_keys=None, figsize=None):
138138

139139
plot_study_unit_counts(self, case_keys, figsize=figsize)
140140

141-
def plot_unit_losses(self, before, after, metric=["precision"], figsize=None):
141+
def plot_unit_losses(self, before, after, metric=["accuracy"], figsize=None):
142142
import matplotlib.pyplot as plt
143143

144144
fig, axs = plt.subplots(ncols=1, nrows=len(metric), figsize=figsize, squeeze=False)

src/spikeinterface/benchmark/benchmark_peak_detection.py

Lines changed: 27 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from .benchmark_base import Benchmark, BenchmarkStudy
1616
from spikeinterface.core.basesorting import minimum_spike_dtype
1717
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
18+
from .benchmark_plot_tools import fit_sigmoid, sigmoid
1819

1920

2021
class PeakDetectionBenchmark(Benchmark):
@@ -126,7 +127,7 @@ def create_benchmark(self, key):
126127
def plot_agreements_by_channels(self, case_keys=None, figsize=(15, 15)):
127128
if case_keys is None:
128129
case_keys = list(self.cases.keys())
129-
import pylab as plt
130+
import matplotlib.pyplot as plt
130131

131132
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
132133

@@ -138,7 +139,7 @@ def plot_agreements_by_channels(self, case_keys=None, figsize=(15, 15)):
138139
def plot_agreements_by_units(self, case_keys=None, figsize=(15, 15)):
139140
if case_keys is None:
140141
case_keys = list(self.cases.keys())
141-
import pylab as plt
142+
import matplotlib.pyplot as plt
142143

143144
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
144145

@@ -151,32 +152,39 @@ def plot_performances_vs_snr(self, case_keys=None, figsize=(15, 15), detect_thre
151152
if case_keys is None:
152153
case_keys = list(self.cases.keys())
153154

155+
import matplotlib.pyplot as plt
156+
154157
fig, axs = plt.subplots(ncols=1, nrows=3, figsize=figsize)
155158

156159
for count, k in enumerate(("accuracy", "recall", "precision")):
157160

158161
ax = axs[count]
159162
for key in case_keys:
163+
color = self.get_colors()[key]
160164
label = self.cases[key]["label"]
161165

162166
analyzer = self.get_sorting_analyzer(key)
163167
metrics = analyzer.get_extension("quality_metrics").get_data()
164168
x = metrics["snr"].values
165169
y = self.get_result(key)["sliced_gt_comparison"].get_performance()[k].values
166-
ax.scatter(x, y, marker=".", label=label)
170+
ax.scatter(x, y, marker=".", label=label, color=color)
167171
ax.set_title(k)
168172
if detect_threshold is not None:
169173
ymin, ymax = ax.get_ylim()
170174
ax.plot([detect_threshold, detect_threshold], [ymin, ymax], "k--")
171175

176+
popt = fit_sigmoid(x, y, p0=None)
177+
xfit = np.linspace(0, max(metrics["snr"].values), 100)
178+
ax.plot(xfit, sigmoid(xfit, *popt), color=color)
179+
172180
if count == 2:
173181
ax.legend()
174182

175183
def plot_detected_amplitudes(self, case_keys=None, figsize=(15, 5), detect_threshold=None):
176184

177185
if case_keys is None:
178186
case_keys = list(self.cases.keys())
179-
import pylab as plt
187+
import matplotlib.pyplot as plt
180188

181189
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
182190

@@ -201,7 +209,7 @@ def plot_deltas_per_cells(self, case_keys=None, figsize=(15, 5)):
201209

202210
if case_keys is None:
203211
case_keys = list(self.cases.keys())
204-
import pylab as plt
212+
import matplotlib.pyplot as plt
205213

206214
fig, axs = plt.subplots(ncols=len(case_keys), nrows=1, figsize=figsize, squeeze=False)
207215
for count, key in enumerate(case_keys):
@@ -222,12 +230,12 @@ def plot_template_similarities(self, case_keys=None, metric="l2", figsize=(15, 5
222230

223231
if case_keys is None:
224232
case_keys = list(self.cases.keys())
225-
import pylab as plt
233+
import matplotlib.pyplot as plt
226234

227235
fig, ax = plt.subplots(ncols=1, nrows=1, figsize=figsize, squeeze=True)
228236
for key in case_keys:
229237

230-
import sklearn
238+
import sklearn.metrics
231239

232240
gt_templates = self.get_result(key)["gt_templates"]
233241
found_templates = self.get_result(key)["templates"]
@@ -244,14 +252,22 @@ def plot_template_similarities(self, case_keys=None, metric="l2", figsize=(15, 5
244252
else:
245253
distances[i] = sklearn.metrics.pairwise_distances(a[None, :], b[None, :], metric)[0, 0]
246254

255+
color = self.get_colors()[key]
256+
247257
label = self.cases[key]["label"]
248258
analyzer = self.get_sorting_analyzer(key)
249259
metrics = analyzer.get_extension("quality_metrics").get_data()
250260
x = metrics["snr"].values
251-
ax.scatter(x, distances, marker=".", label=label)
252-
if detect_threshold is not None:
253-
ymin, ymax = ax.get_ylim()
254-
ax.plot([detect_threshold, detect_threshold], [ymin, ymax], "k--")
261+
y = distances
262+
ax.scatter(x, y, marker=".", label=label, color=color)
263+
264+
popt = fit_sigmoid(x, y, p0=None)
265+
xfit = np.linspace(0, max(metrics["snr"].values), 100)
266+
ax.plot(xfit, sigmoid(xfit, *popt), color=color)
267+
268+
if detect_threshold is not None:
269+
ymin, ymax = ax.get_ylim()
270+
ax.plot([detect_threshold, detect_threshold], [ymin, ymax], "k--")
255271

256272
ax.legend()
257273
ax.set_xlabel("snr")

src/spikeinterface/benchmark/benchmark_plot_tools.py

Lines changed: 25 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import numpy as np
2+
import warnings
23

34

45
def _simpleaxis(ax):
@@ -193,7 +194,7 @@ def plot_agreement_matrix(study, ordered=True, case_keys=None):
193194
case_keys = list(study.cases.keys())
194195

195196
num_axes = len(case_keys)
196-
fig, axs = plt.subplots(ncols=num_axes)
197+
fig, axs = plt.subplots(ncols=num_axes, squeeze=False)
197198

198199
for count, key in enumerate(case_keys):
199200
ax = axs.flatten()[count]
@@ -214,7 +215,9 @@ def plot_agreement_matrix(study, ordered=True, case_keys=None):
214215
ax.set_xticks([])
215216

216217

217-
def plot_performances_vs_snr(study, case_keys=None, figsize=None, metrics=["accuracy", "recall", "precision"]):
218+
def plot_performances_vs_snr(
219+
study, case_keys=None, figsize=None, metrics=["accuracy", "recall", "precision"], snr_dataset_reference=None
220+
):
218221
import matplotlib.pyplot as plt
219222

220223
if case_keys is None:
@@ -228,7 +231,13 @@ def plot_performances_vs_snr(study, case_keys=None, figsize=None, metrics=["accu
228231
for key in case_keys:
229232
label = study.cases[key]["label"]
230233

231-
analyzer = study.get_sorting_analyzer(key)
234+
if snr_dataset_reference is None:
235+
# use the SNR of each dataset
236+
analyzer = study.get_sorting_analyzer(key)
237+
else:
238+
# use the same SNR from a reference dataset
239+
analyzer = study.get_sorting_analyzer(dataset_key=snr_dataset_reference)
240+
232241
metrics = analyzer.get_extension("quality_metrics").get_data()
233242
x = metrics["snr"].values
234243
y = study.get_result(key)["gt_comparison"].get_performance()[k].values
@@ -303,3 +312,16 @@ def plot_performances_comparison(
303312
ax.legend(handles=patches)
304313
fig.tight_layout()
305314
return fig
315+
316+
317+
def sigmoid(x, x0, k, b):
318+
with warnings.catch_warnings(action="ignore"):
319+
out = (1 / (1 + np.exp(-k * (x - x0)))) + b
320+
return out
321+
322+
323+
def fit_sigmoid(xdata, ydata, p0=None):
324+
from scipy.optimize import curve_fit
325+
326+
popt, pcov = curve_fit(sigmoid, xdata, ydata, p0)
327+
return popt

0 commit comments

Comments
 (0)