Skip to content

Commit b70ae9b

Browse files
committed
clean and debug
1 parent 2a4809b commit b70ae9b

9 files changed

Lines changed: 20 additions & 1843 deletions

File tree

src/spikeinterface/sorters/internal/tests/test_spykingcircus2.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,16 @@
44

55
from spikeinterface.sorters import Spykingcircus2Sorter
66

7+
from pathlib import Path
78

89
class SpykingCircus2SorterCommonTestSuite(SorterCommonTestSuite, unittest.TestCase):
910
SorterClass = Spykingcircus2Sorter
1011

1112

1213
if __name__ == "__main__":
14+
from spikeinterface import set_global_job_kwargs
15+
set_global_job_kwargs(n_jobs=1, progress_bar=False)
1316
test = SpykingCircus2SorterCommonTestSuite()
17+
test.cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "sorters"
1418
test.setUp()
1519
test.test_with_run()

src/spikeinterface/sorters/internal/tests/test_tridesclous2.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,15 @@
44

55
from spikeinterface.sorters import Tridesclous2Sorter
66

7+
from pathlib import Path
8+
79

810
class Tridesclous2SorterCommonTestSuite(SorterCommonTestSuite, unittest.TestCase):
911
SorterClass = Tridesclous2Sorter
1012

1113

1214
if __name__ == "__main__":
1315
test = Tridesclous2SorterCommonTestSuite()
16+
test.cache_folder = Path(__file__).resolve().parents[4] / "cache_folder" / "sorters"
1417
test.setUp()
1518
test.test_with_run()

src/spikeinterface/sorters/internal/tridesclous2.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -226,7 +226,8 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
226226
matching_method = params["matching"]["method"]
227227
matching_params = params["matching"]["method_kwargs"].copy()
228228
matching_params["templates"] = templates
229-
matching_params["noise_levels"] = noise_levels
229+
if params["matching"]["method"] in ("tdc-peeler", ):
230+
matching_params["noise_levels"] = noise_levels
230231
spikes = find_spikes_from_templates(
231232
recording_for_peeler, method=matching_method, method_kwargs=matching_params, **job_kwargs
232233
)

src/spikeinterface/sortingcomponents/matching/base.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -32,11 +32,12 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin):
3232
spikes["segment_index"] = segment_index
3333

3434
margin = self.get_trace_margin()
35-
if margin > 0:
35+
if margin > 0 and spikes.size > 0:
3636
keep = (spikes["sample_index"] >= margin) & (spikes["sample_index"] < (traces.shape[0] - margin))
3737
spikes = spikes[keep]
3838

39-
return spikes
39+
# node pipeline need to return a tuple
40+
return (spikes, )
4041

4142
def compute_matching(self, traces, start_frame, end_frame, segment_index):
4243
raise NotImplementedError

0 commit comments

Comments
 (0)