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/iterative_reassign_benchmark.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
"""Benchmark: iterative waveform reassignment (multi-pass step 9h).

Previous: single-pass reassignment with margin=0.05.
New: up to 5 passes, recomputing templates after each pass, converging
when no spikes move.
"""

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.7820, "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:<50} 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:<50} 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, "iterative reassign (new, up to 5 passes)", detail=detail)

print("\nDone.")


if __name__ == "__main__":
main()
116 changes: 61 additions & 55 deletions src/sorter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3845,7 +3845,7 @@ pub fn sort_multichannel<
}
}

// 9h. Template-based waveform reassignment.
// 9h. Iterative template-based waveform reassignment.
//
// After all detection and clustering passes, reassign each spike to the
// cluster whose mean waveform it most closely matches, but ONLY among
Expand All @@ -3856,70 +3856,76 @@ pub fn sort_multichannel<
// 1. It uses all W waveform samples (not just K << W PCA features)
// 2. The channel constraint prevents cross-channel contamination
// 3. It naturally handles the case where PCA discards discriminative info
//
// Runs iteratively: each pass recomputes templates from updated labels,
// then reassigns. Converges when no spikes move.
if n_clusters > 1 && n_extracted > n_clusters {
// Build templates: mean waveform per cluster, tracking peak channel
let mut tmpl_wf = [[0.0 as Float; W]; N];
let mut tmpl_count = [0u32; N];
let mut tmpl_ch = [0usize; N];
compute_cluster_means::<W, N>(
waveform_buf,
labels,
event_buf,
n_extracted,
n_clusters,
&mut tmpl_wf,
&mut tmpl_count,
&mut tmpl_ch,
);
let margin_factor = 1.0 - config.waveform_reassign_margin;
let max_passes = 5;

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];
let mut tmpl_count = [0u32; N];
let mut tmpl_ch = [0usize; N];
compute_cluster_means::<W, N>(
waveform_buf,
labels,
event_buf,
n_extracted,
n_clusters,
&mut tmpl_wf,
&mut tmpl_count,
&mut tmpl_ch,
);

// Reassign: for each spike, find the nearest template on the same
// channel. Require a 30% distance improvement to prevent marginal
// reassignments that break good clustering.
let mut changed = 0usize;
for i in 0..n_extracted {
let spike_ch = event_buf[i].channel;
let old_label = labels[i];
// 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
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;
// Compute distance to current template
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;
}
}
}

let mut best_c = old_label;
let mut best_d: Float = Float::MAX;
for c in 0..n_clusters.min(N) {
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;
}
let mut d = 0.0;
for w in 0..W {
let diff = waveform_buf[i][w] - tmpl_wf[c][w];
d += diff * diff;
let mut best_c = old_label;
let mut best_d: Float = Float::MAX;
for c in 0..n_clusters.min(N) {
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;
}
let mut d = 0.0;
for w in 0..W {
let diff = waveform_buf[i][w] - tmpl_wf[c][w];
d += diff * diff;
}
if d < best_d {
best_d = d;
best_c = c;
}
}
if d < best_d {
best_d = d;
best_c = c;
if best_c != old_label && best_d < old_d * margin_factor {
labels[i] = best_c;
changed += 1;
}
}
// Only reassign if new template is sufficiently closer.
// The margin prevents marginal reassignments from breaking good clustering.
let margin_factor = 1.0 - config.waveform_reassign_margin;
if best_c != old_label && best_d < old_d * margin_factor {
labels[i] = best_c;
changed += 1;

// If no spikes moved, we've converged
if changed == 0 {
break;
}
}

// If reassignment changed labels, remove empty clusters
if changed > 0 {
// Remove empty clusters
let mut counts = [0u32; N];
for label in labels.iter().take(n_extracted) {
if *label < N {
Expand Down
Loading