diff --git a/benchmarks/neighbor_reassign_benchmark.py b/benchmarks/neighbor_reassign_benchmark.py new file mode 100644 index 0000000..876626a --- /dev/null +++ b/benchmarks/neighbor_reassign_benchmark.py @@ -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() diff --git a/src/sorter.rs b/src/sorter.rs index 44a064a..ce93f52 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -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]; @@ -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; @@ -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; } @@ -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;