11import pytest
22
33import 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
101102if __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 )
0 commit comments