From 26047fbc59d35f41bcdd0ea784a6c9c7cf323dd9 Mon Sep 17 00:00:00 2001 From: Fredrik Willumsen Haug Date: Thu, 23 Apr 2026 07:15:10 -0400 Subject: [PATCH] [sorter] iterate waveform reassignment to convergence --- benchmarks/iterative_reassign_benchmark.py | 99 ++++++++++++++++++ src/sorter.rs | 116 +++++++++++---------- 2 files changed, 160 insertions(+), 55 deletions(-) create mode 100644 benchmarks/iterative_reassign_benchmark.py diff --git a/benchmarks/iterative_reassign_benchmark.py b/benchmarks/iterative_reassign_benchmark.py new file mode 100644 index 0000000..4916fe7 --- /dev/null +++ b/benchmarks/iterative_reassign_benchmark.py @@ -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() diff --git a/src/sorter.rs b/src/sorter.rs index 1980290..44a064a 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -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 @@ -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::( - 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::( + 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 {