Skip to content

Commit d8229db

Browse files
committed
Fix tests.
1 parent 066caa0 commit d8229db

11 files changed

Lines changed: 44 additions & 33 deletions

src/spikeinterface/benchmark/benchmark_matching.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,7 @@ def plot_performances_vs_snr(self, case_keys=None, figsize=None, metrics=["accur
7777
if case_keys is None:
7878
case_keys = list(self.cases.keys())
7979

80+
import matplotlib.pyplot as plt
8081
fig, axs = plt.subplots(ncols=1, nrows=len(metrics), figsize=figsize, squeeze=False)
8182

8283
for count, k in enumerate(metrics):
Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,2 +0,0 @@
1-
from spikeinterface.benchmark.benchmark_base import Benchmark, BenchmarkStudy
2-
from spikeinterface.benchmark.benchmark_plot_tools import _simpleaxis

src/spikeinterface/benchmark/tests/test_benchmark_clustering.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,13 @@
33

44
import shutil
55

6-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import make_dataset
7-
from spikeinterface.sortingcomponents.benchmark.benchmark_clustering import ClusteringStudy
6+
from spikeinterface.benchmark.tests.common_benchmark_testing import make_dataset
7+
from spikeinterface.benchmark.benchmark_clustering import ClusteringStudy
88
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
99
from spikeinterface.core.template_tools import get_template_extremum_channel
1010

11+
from pathlib import Path
12+
1113

1214
@pytest.mark.skip()
1315
def test_benchmark_clustering(create_cache_folder):
@@ -78,4 +80,5 @@ def test_benchmark_clustering(create_cache_folder):
7880

7981

8082
if __name__ == "__main__":
81-
test_benchmark_clustering()
83+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
84+
test_benchmark_clustering(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_matching.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,19 @@
11
import pytest
22

33
import shutil
4+
from pathlib import Path
45

56

67
from spikeinterface.core import (
78
get_noise_levels,
89
compute_sparsity,
910
)
1011

11-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import (
12+
from spikeinterface.benchmark.tests.common_benchmark_testing import (
1213
make_dataset,
1314
compute_gt_templates,
1415
)
15-
from spikeinterface.sortingcomponents.benchmark.benchmark_matching import MatchingStudy
16+
from spikeinterface.benchmark.benchmark_matching import MatchingStudy
1617

1718

1819
@pytest.mark.skip()
@@ -72,4 +73,5 @@ def test_benchmark_matching(create_cache_folder):
7273

7374

7475
if __name__ == "__main__":
75-
test_benchmark_matching()
76+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
77+
test_benchmark_matching(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_motion_estimation.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,13 @@
22

33

44
import shutil
5+
from pathlib import Path
56

6-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import (
7+
from spikeinterface.benchmark.tests.common_benchmark_testing import (
78
make_drifting_dataset,
89
)
910

10-
from spikeinterface.sortingcomponents.benchmark.benchmark_motion_estimation import MotionEstimationStudy
11+
from spikeinterface.benchmark.benchmark_motion_estimation import MotionEstimationStudy
1112

1213

1314
@pytest.mark.skip()
@@ -75,4 +76,5 @@ def test_benchmark_motion_estimaton(create_cache_folder):
7576

7677

7778
if __name__ == "__main__":
78-
test_benchmark_motion_estimaton()
79+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
80+
test_benchmark_motion_estimaton(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_motion_interpolation.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,14 @@
44
import numpy as np
55

66
import shutil
7+
from pathlib import Path
78

8-
9-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import (
9+
from spikeinterface.benchmark.tests.common_benchmark_testing import (
1010
make_drifting_dataset,
1111
)
1212

13-
from spikeinterface.sortingcomponents.benchmark.benchmark_motion_interpolation import MotionInterpolationStudy
14-
from spikeinterface.sortingcomponents.benchmark.benchmark_motion_estimation import (
13+
from spikeinterface.benchmark.benchmark_motion_interpolation import MotionInterpolationStudy
14+
from spikeinterface.benchmark.benchmark_motion_estimation import (
1515
# get_unit_displacement,
1616
get_gt_motion_from_unit_displacement,
1717
)
@@ -139,4 +139,5 @@ def test_benchmark_motion_interpolation(create_cache_folder):
139139

140140

141141
if __name__ == "__main__":
142-
test_benchmark_motion_interpolation()
142+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
143+
test_benchmark_motion_interpolation(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_peak_detection.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,10 @@
11
import pytest
22

33
import shutil
4+
from pathlib import Path
45

5-
6-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import make_dataset
7-
from spikeinterface.sortingcomponents.benchmark.benchmark_peak_detection import PeakDetectionStudy
6+
from spikeinterface.benchmark.tests.common_benchmark_testing import make_dataset
7+
from spikeinterface.benchmark.benchmark_peak_detection import PeakDetectionStudy
88
from spikeinterface.core.sortinganalyzer import create_sorting_analyzer
99
from spikeinterface.core.template_tools import get_template_extremum_channel
1010

@@ -69,5 +69,5 @@ def test_benchmark_peak_detection(create_cache_folder):
6969

7070

7171
if __name__ == "__main__":
72-
# test_benchmark_peak_localization()
73-
test_benchmark_peak_detection()
72+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
73+
test_benchmark_peak_detection(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_peak_localization.py

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,12 @@
11
import pytest
22

33
import shutil
4+
from pathlib import Path
45

6+
from spikeinterface.benchmark.tests.common_benchmark_testing import make_dataset
57

6-
from spikeinterface.sortingcomponents.benchmark.tests.common_benchmark_testing import make_dataset
7-
8-
from spikeinterface.sortingcomponents.benchmark.benchmark_peak_localization import PeakLocalizationStudy
9-
from spikeinterface.sortingcomponents.benchmark.benchmark_peak_localization import UnitLocalizationStudy
8+
from spikeinterface.benchmark.benchmark_peak_localization import PeakLocalizationStudy
9+
from spikeinterface.benchmark.benchmark_peak_localization import UnitLocalizationStudy
1010

1111

1212
@pytest.mark.skip()
@@ -28,7 +28,8 @@ def test_benchmark_peak_localization(create_cache_folder):
2828
"init_kwargs": {"gt_positions": gt_sorting.get_property("gt_unit_locations")},
2929
"params": {
3030
"method": method,
31-
"method_kwargs": {"ms_before": 2},
31+
"ms_before": 2.0,
32+
"method_kwargs": {},
3233
},
3334
}
3435

@@ -60,7 +61,7 @@ def test_benchmark_unit_locations(create_cache_folder):
6061
cache_folder = create_cache_folder
6162
job_kwargs = dict(n_jobs=0.8, chunk_duration="100ms")
6263

63-
recording, gt_sorting = make_dataset()
64+
recording, gt_sorting, gt_analyzer = make_dataset()
6465

6566
# create study
6667
study_folder = cache_folder / "study_unit_locations"
@@ -71,7 +72,7 @@ def test_benchmark_unit_locations(create_cache_folder):
7172
"label": f"{method} on toy",
7273
"dataset": "toy",
7374
"init_kwargs": {"gt_positions": gt_sorting.get_property("gt_unit_locations")},
74-
"params": {"method": method, "method_kwargs": {"ms_before": 2}},
75+
"params": {"method": method, "ms_before": 2.0, "method_kwargs": {}},
7576
}
7677

7778
if study_folder.exists():
@@ -99,5 +100,6 @@ def test_benchmark_unit_locations(create_cache_folder):
99100

100101

101102
if __name__ == "__main__":
102-
# test_benchmark_peak_localization()
103-
test_benchmark_unit_locations()
103+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
104+
# test_benchmark_peak_localization(cache_folder)
105+
test_benchmark_unit_locations(cache_folder)
Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,12 @@
11
import pytest
22

3+
from pathlib import Path
34

45
@pytest.mark.skip()
56
def test_benchmark_peak_selection(create_cache_folder):
67
cache_folder = create_cache_folder
7-
pass
88

99

1010
if __name__ == "__main__":
11-
test_benchmark_peak_selection()
11+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks"
12+
test_benchmark_peak_selection(cache_folder)

src/spikeinterface/benchmark/tests/test_benchmark_sorter.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,5 +86,5 @@ def test_SorterStudy(setup_module):
8686

8787
if __name__ == "__main__":
8888
study_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "benchmarks" / "test_SorterStudy"
89-
# create_a_study(study_folder)
89+
create_a_study(study_folder)
9090
test_SorterStudy(study_folder)

0 commit comments

Comments
 (0)