Skip to content

Commit e782cbb

Browse files
committed
oups : spikes in margin
1 parent ed7e636 commit e782cbb

6 files changed

Lines changed: 19 additions & 17 deletions

File tree

src/spikeinterface/sortingcomponents/matching/base.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,18 @@ def get_dtype(self):
2727
def get_trace_margin(self):
2828
raise NotImplementedError
2929

30-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
30+
def compute(self, traces, start_frame, end_frame, segment_index, max_margin):
31+
spikes = self.compute_matching(traces, start_frame, end_frame, segment_index)
32+
spikes["segment_index"] = segment_index
33+
34+
margin = self.get_trace_margin()
35+
if margin > 0:
36+
keep = (spikes["sample_index"] >= margin) & (spikes["sample_index"] < (traces.shape[0] - margin))
37+
spikes = spikes[keep]
38+
39+
return spikes
40+
41+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
3142
raise NotImplementedError
3243

3344
def get_extra_outputs(self):

src/spikeinterface/sortingcomponents/matching/circus.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -274,7 +274,7 @@ def get_extra_outputs(self):
274274
def get_trace_margin(self):
275275
return self.margin
276276

277-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
277+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
278278
import scipy.spatial
279279
import scipy
280280

@@ -478,8 +478,6 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *ar
478478
order = np.argsort(spikes["sample_index"])
479479
spikes = spikes[order]
480480

481-
spikes["segment_index"] = segment_index
482-
483481
return spikes
484482

485483

@@ -1024,7 +1022,7 @@ def get_trace_margin(self):
10241022
return self.margin
10251023

10261024

1027-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
1025+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
10281026

10291027
neighbor_window = self.num_samples - 1
10301028

@@ -1107,8 +1105,6 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *ar
11071105
order = np.argsort(spikes["sample_index"])
11081106
spikes = spikes[order]
11091107

1110-
spikes["segment_index"] = segment_index
1111-
11121108
return spikes
11131109

11141110

src/spikeinterface/sortingcomponents/matching/naive.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,7 @@ def get_trace_margin(self):
5151
return self.margin
5252

5353

54-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
54+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
5555

5656
if self.margin > 0:
5757
peak_traces = traces[self.margin:-self.margin, :]
@@ -78,8 +78,6 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *ar
7878
spikes["cluster_index"][i] = cluster_index
7979
spikes["amplitude"][i] = 0.0
8080

81-
spikes["segment_index"] = segment_index
82-
8381
return spikes
8482

8583

src/spikeinterface/sortingcomponents/matching/tdc.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,7 @@ def __init__(self, recording, return_output=True, parents=None,
160160
def get_trace_margin(self):
161161
return self.margin
162162

163-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
163+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
164164
traces = traces.copy()
165165

166166
all_spikes = []
@@ -186,8 +186,6 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *ar
186186
else:
187187
all_spikes = np.zeros(0, dtype=_base_matching_dtype)
188188

189-
all_spikes["segment_index"] = segment_index
190-
191189
return all_spikes
192190

193191
def _find_spikes_one_level(self, traces, level=0):

src/spikeinterface/sortingcomponents/matching/wobble.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -383,7 +383,7 @@ def __init__(self, recording, return_output=True, parents=None,
383383
def get_trace_margin(self):
384384
return self.margin
385385

386-
def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *args):
386+
def compute_matching(self, traces, start_frame, end_frame, segment_index):
387387

388388
# Unpack method_kwargs
389389
# nbefore, nafter = method_kwargs["nbefore"], method_kwargs["nafter"]
@@ -450,8 +450,7 @@ def compute(self, traces, start_frame, end_frame, segment_index, max_margin, *ar
450450
spikes["cluster_index"] = spike_train[:, 1]
451451
spikes["channel_index"] = channel_inds
452452
spikes["amplitude"] = amplitudes
453-
spikes["segment_index"] = segment_index
454-
453+
455454
return spikes
456455

457456
# TODO: Replace this method with equivalent from spikeinterface

src/spikeinterface/sortingcomponents/tests/test_template_matching.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -77,7 +77,7 @@ def test_find_spikes_from_templates(method, sorting_analyzer):
7777
# comp = si.compare_sorter_to_ground_truth(gt_sorting, sorting)
7878
# si.plot_agreement_matrix(comp, ax=ax)
7979
# ax.set_title(method)
80-
# plt.show()
80+
plt.show()
8181

8282

8383
if __name__ == "__main__":

0 commit comments

Comments
 (0)