Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 99 additions & 0 deletions benchmarks/neighbor_reassign_benchmark.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
"""Benchmark: neighbor-channel augmented waveform reassignment.

Step 9h now also computes neighbor-channel templates and uses them
(weighted at 0.5x) in the L2 distance for reassignment. This adds
spatial discrimination for units that share a peak channel.
"""

import time
import numpy as np
import spikeinterface.core as si
from spikeinterface.comparison import compare_sorter_to_ground_truth
from spikeinterface.core import NumpySorting
import zpybci as zbci

TOLERANCE_MS = 0.4
_cache = {}
BASELINE = {"easy": 0.8280, "medium": 0.7880, "hard": 0.6420}


def gen(difficulty):
if difficulty in _cache:
return _cache[difficulty]
cfgs = {
"easy": {"num_channels": 4, "num_units": 3, "duration": 30.0, "noise_levels": 3.0},
"medium": {"num_channels": 16, "num_units": 8, "duration": 30.0, "noise_levels": 5.0},
"hard": {"num_channels": 32, "num_units": 15, "duration": 30.0, "noise_levels": 8.0},
}
c = cfgs[difficulty]
rec, gt = si.generate_ground_truth_recording(
durations=[c["duration"]], sampling_frequency=30000.0,
num_channels=c["num_channels"], num_units=c["num_units"],
seed=42, noise_kwargs={"noise_levels": c["noise_levels"], "strategy": "on_the_fly"},
)
traces = rec.get_traces(return_in_uV=True).astype(np.float64)
fs = rec.get_sampling_frequency()
gt_trains = {}
for uid in gt.get_unit_ids():
gt_trains[uid] = np.sort(gt.get_unit_spike_train(uid))
all_t = np.concatenate([gt_trains[u] for u in sorted(gt_trains)])
all_l = np.concatenate([np.full(len(gt_trains[u]), u, dtype=np.int64)
for u in sorted(gt_trains)])
idx = np.argsort(all_t)
gt_sorting = NumpySorting.from_samples_and_labels(
[all_t[idx]], [all_l[idx]], sampling_frequency=fs
)
_cache[difficulty] = (traces, fs, c["num_channels"], gt_trains, gt_sorting)
return _cache[difficulty]


def run(difficulty, label, detail=False, **kwargs):
traces, fs, n_ch, gt_trains, gt_sorting = gen(difficulty)
data = traces.copy()
probe = zbci.ProbeLayout.linear(n_ch, 25.0)
t0 = time.perf_counter()
result = zbci.sort_multichannel(data, probe, **kwargs)
elapsed = time.perf_counter() - t0
n_spikes = result["n_spikes"]
n_cl = result["n_clusters"]
n_gt = len(gt_trains)
if n_spikes == 0:
print(f" {label:<55} acc=0.000 {elapsed:.1f}s")
return 0.0
spike_times = np.asarray(result["spike_times"][:n_spikes], dtype=np.int64)
labels_out = np.asarray(result["labels"][:n_spikes], dtype=np.int64)
sorting = NumpySorting.from_samples_and_labels(
[spike_times], [labels_out], sampling_frequency=fs
)
cmp = compare_sorter_to_ground_truth(
gt_sorting, sorting, exhaustive_gt=True,
delta_time=TOLERANCE_MS, match_mode="hungarian"
)
avg = cmp.get_performance(method="pooled_with_average")
acc = float(avg["accuracy"])
wd = cmp.count_well_detected_units(well_detected_score=0.8)
delta = acc - BASELINE[difficulty]
print(f" {label:<55} acc={acc:.3f} ({delta:+.3f}) wd={wd}/{n_gt} {n_spikes}spk/{n_cl}cl {elapsed:.1f}s")
if detail:
per_unit = cmp.get_performance()
for uid in per_unit.index:
a = float(per_unit.loc[uid]["accuracy"])
if a < 0.01:
continue
r = float(per_unit.loc[uid]["recall"])
p = float(per_unit.loc[uid]["precision"])
print(f" unit {uid:>2}: acc={a:.3f} R={r:.3f} P={p:.3f}")
return acc


def main():
for difficulty in ["easy", "medium", "hard"]:
print(f"\n=== {difficulty.upper()} ===")
detail = difficulty in ("medium", "hard")
run(difficulty, "neighbor-augmented reassignment (new)", detail=detail)

print("\nDone.")


if __name__ == "__main__":
main()
73 changes: 71 additions & 2 deletions src/sorter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3857,12 +3857,28 @@ pub fn sort_multichannel<
// 2. The channel constraint prevents cross-channel contamination
// 3. It naturally handles the case where PCA discards discriminative info
//
// For multi-channel recordings (C > 1), also computes a neighbor-channel
// template: the mean waveform on the nearest neighbor of each cluster's
// peak channel. This adds spatial discrimination: units on the same peak
// channel but at different positions have different neighbor-channel
// waveforms. The neighbor distance is weighted at 0.5x the peak distance.
//
// Runs iteratively: each pass recomputes templates from updated labels,
// then reassigns. Converges when no spikes move.
if n_clusters > 1 && n_extracted > n_clusters {
let margin_factor = 1.0 - config.waveform_reassign_margin;
let max_passes = 5;

// Precompute nearest neighbor for each channel
let mut nearest_nbr = [0usize; C];
if C > 1 {
for (ch, nbr) in nearest_nbr.iter_mut().enumerate().take(C) {
let mut nbuf = [0usize; 1];
let n = probe.nearest_channels(ch, 1, &mut nbuf);
*nbr = if n > 0 { nbuf[0] } else { ch };
}
}

for _pass in 0..max_passes {
// Build templates: mean waveform per cluster, tracking peak channel
let mut tmpl_wf = [[0.0 as Float; W]; N];
Expand All @@ -3879,19 +3895,66 @@ pub fn sort_multichannel<
&mut tmpl_ch,
);

// Build neighbor-channel templates for spatial discrimination
let mut tmpl_nbr = [[0.0 as Float; W]; N];
if C > 1 {
let mut nbr_count = [0u32; N];
for i in 0..n_extracted {
let label = labels[i];
if label >= n_clusters || label >= N {
continue;
}
let peak_ch = event_buf[i].channel;
let nbr_ch = nearest_nbr[peak_ch];
let t = event_buf[i].sample;
if t < config.pre_samples {
continue;
}
let start = t - config.pre_samples;
if start + W > data.len() {
continue;
}
for w in 0..W {
tmpl_nbr[label][w] += data[start + w][nbr_ch];
}
nbr_count[label] += 1;
}
for c in 0..n_clusters.min(N) {
if nbr_count[c] > 0 {
let inv = 1.0 / nbr_count[c] as Float;
for val in tmpl_nbr[c].iter_mut().take(W) {
*val *= inv;
}
}
}
}

// Reassign: for each spike, find the nearest template on the same channel.
let mut changed = 0usize;
for i in 0..n_extracted {
let spike_ch = event_buf[i].channel;
let old_label = labels[i];

// Compute distance to current template
// Extract neighbor waveform for this spike (on the fly)
let nbr_ch = nearest_nbr[spike_ch];
let t = event_buf[i].sample;
let has_nbr =
C > 1 && t >= config.pre_samples && t - config.pre_samples + W <= data.len();

// Compute distance to current template (peak + neighbor)
let mut old_d: Float = 0.0;
if old_label < n_clusters && old_label < N && tmpl_ch[old_label] == spike_ch {
for w in 0..W {
let diff = waveform_buf[i][w] - tmpl_wf[old_label][w];
old_d += diff * diff;
}
if has_nbr {
let start = t - config.pre_samples;
for w in 0..W {
let diff = data[start + w][nbr_ch] - tmpl_nbr[old_label][w];
old_d += 0.5 * diff * diff;
}
}
}

let mut best_c = old_label;
Expand All @@ -3900,7 +3963,6 @@ pub fn sort_multichannel<
if tmpl_count[c] < config.template_min_count as u32 {
continue;
}
// Only consider clusters on the same peak channel
if tmpl_ch[c] != spike_ch {
continue;
}
Expand All @@ -3909,6 +3971,13 @@ pub fn sort_multichannel<
let diff = waveform_buf[i][w] - tmpl_wf[c][w];
d += diff * diff;
}
if has_nbr {
let start = t - config.pre_samples;
for w in 0..W {
let diff = data[start + w][nbr_ch] - tmpl_nbr[c][w];
d += 0.5 * diff * diff;
}
}
if d < best_d {
best_d = d;
best_c = c;
Expand Down
Loading