From b5dfc376ef781c5b883c355f31c44711ecfa9658 Mon Sep 17 00:00:00 2001 From: Fredrik Willumsen Haug Date: Sun, 19 Apr 2026 16:25:40 -0400 Subject: [PATCH 1/2] [sorter] activate amplitude-profile spatial features for C>=8 --- .../auto_amplitude_profile_benchmark.py | 171 ++++++++++++++++++ python/src/sorter.rs | 8 + src/sorter.rs | 154 +++++++++++++++- 3 files changed, 332 insertions(+), 1 deletion(-) create mode 100644 benchmarks/auto_amplitude_profile_benchmark.py diff --git a/benchmarks/auto_amplitude_profile_benchmark.py b/benchmarks/auto_amplitude_profile_benchmark.py new file mode 100644 index 0000000..d806f2e --- /dev/null +++ b/benchmarks/auto_amplitude_profile_benchmark.py @@ -0,0 +1,171 @@ +"""Benchmark: auto_amplitude_profile — amplitude-profile spatial features for large probes. + +Baseline (auto_amplitude_profile=False): Easy 82.7%, Medium 73.4%, Hard 64.2%, Avg 73.4% + +auto_amplitude_profile=True activates use_amplitude_profile for C>=8. +The amplitude profile encodes neighbor-channel energy bleed as a physics-based +spatial fingerprint in feature dimension K-2, replacing the weaker spike half-width. +On large probes, multiple units often share the same peak channel; their waveform +shapes are similar, but their spatial decay profiles differ. The ratio: + neighbor_sum / (n_neighbors * peak_amp) +ranges [0.3, 2.0]: tightly-localized units ~0.5, spread units ~1.2. +profile_scale = cluster_threshold * 2.0 = 14 maps this to feature range [4.2, 28]. +""" + +import json +import time +from datetime import datetime + +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.8270, "medium": 0.7340, "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) + defaults = dict( + threshold=5.0, refractory=15, spatial_radius=75.0, temporal_radius=5, + align_half_window=15, pre_samples=20, cluster_threshold=7.0, + cluster_max_count=1000, whitening_epsilon=1e-6, detection_mode="amplitude", + sneo_smooth_window=3, matched_filter_detect=True, + matched_filter_threshold=3.5, gmm_max_iter=10, + bandpass_low=0.0, bandpass_high=0.0, sample_rate=fs, + common_median_ref=False, merge_dprime_threshold=2.0, + use_amplitude_profile=False, amplitude_profile_neighbors=4, + ccg_merge=True, auto_cluster_threshold=True, + svd_init=False, auto_svd_init=True, gmm_refine=False, + min_cluster_snr=2.5, refinement_iterations=1, auto_refine=True, + refine_collapse_guard=True, refine_isi_guard=False, + refine_isi_tolerance=0.1, auto_threshold=True, + auto_refine_iterations=3, ncc_threshold=0.70, + template_subtract_passes=2, auto_amplitude_profile=True, + ) + defaults.update(kwargs) + t0 = time.perf_counter() + result = zbci.sort_multichannel(data, probe, **defaults) + 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:<78} acc=0.000 0 spk {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:<78} acc={acc:.3f} ({delta:+.3f}) wd={wd}/{n_gt} {n_spikes}spk/{n_cl}cl {elapsed:.1f}s") + per_unit_df = cmp.get_performance() + if detail: + for uid in per_unit_df.index: + print(f" unit {uid}: acc={float(per_unit_df.loc[uid]['accuracy']):.3f}") + return acc, {str(uid): float(per_unit_df.loc[uid]["accuracy"]) + for uid in per_unit_df.index} + + +def main(): + results = {} + + for difficulty in ["easy", "medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + detail = difficulty == "hard" + + # Baseline: half-width feature (old behavior, auto_amplitude_profile=False) + acc_off, _ = run(difficulty, "auto_amplitude_profile=False (half-width fallback, baseline)", + auto_amplitude_profile=False) + + # New default: auto_amplitude_profile=True + acc_on, pu_on = run(difficulty, "auto_amplitude_profile=True (amplitude profile for C>=8)", + auto_amplitude_profile=True, detail=detail) + + # Forced on for all probes: use_amplitude_profile=True, auto=False + acc_forced, _ = run(difficulty, "use_amplitude_profile=True (forced, all probe sizes)", + use_amplitude_profile=True, auto_amplitude_profile=False) + + # Profile with more neighbors: 8 instead of 4 + acc_8nbr, _ = run(difficulty, "auto_amplitude_profile=True, neighbors=8", + auto_amplitude_profile=True, amplitude_profile_neighbors=8) + + results[difficulty] = { + "off": acc_off, + "auto_on": acc_on, + "forced": acc_forced, + "neighbors8": acc_8nbr, + "per_unit_auto_on": {str(k): float(v) for k, v in pu_on.items()}, + } + + print(f"\n{'='*92}") + print(f"{'Config':<62} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*92}") + for label, key in [ + ("auto_amplitude_profile=False (half-width, baseline)", "off"), + ("auto_amplitude_profile=True (default, C>=8 gate)", "auto_on"), + ("use_amplitude_profile=True (forced all sizes)", "forced"), + ("auto_amplitude_profile=True, neighbors=8", "neighbors8"), + ]: + e = results["easy"][key] + m = results["medium"][key] + h = results["hard"][key] + avg = (e + m + h) / 3.0 + print(f" {label:<60} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*92}") + + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + out = {"date": ts, "baseline_avg": sum(BASELINE.values()) / 3.0, "results": results} + out_path = f"benchmarks/results/auto_amplitude_profile_{ts}.json" + with open(out_path, "w") as f: + json.dump(out, f, indent=2) + print(f"Saved: {out_path}") + + +if __name__ == "__main__": + main() diff --git a/python/src/sorter.rs b/python/src/sorter.rs index 6805105..803c361 100644 --- a/python/src/sorter.rs +++ b/python/src/sorter.rs @@ -97,6 +97,11 @@ fn sort_error_to_py(e: SortError) -> PyErr { /// active (C >= 8). Effective iterations = max(refinement_iterations, /// auto_refine_iterations). Early convergence exit prevents wasted work when /// the assignment stabilizes in fewer passes. Default: 3. +/// auto_amplitude_profile (bool): Auto-activate amplitude profile spatial features for +/// recordings with 8 or more channels. On large probes multiple units share the +/// same peak channel; the amplitude profile encodes energy bleed into neighboring +/// channels as a physics-based spatial fingerprint, replacing the weaker half-width +/// feature. Has no effect for C < 8. Default: True. /// /// Returns: /// dict: Sorting results with keys: @@ -176,6 +181,7 @@ fn sort_error_to_py(e: SortError) -> PyErr { auto_svd_init = true, auto_threshold = true, auto_refine_iterations = 3usize, + auto_amplitude_profile = true, ))] #[allow(clippy::too_many_arguments)] fn sort_multichannel<'py>( @@ -238,6 +244,7 @@ fn sort_multichannel<'py>( auto_svd_init: bool, auto_threshold: bool, auto_refine_iterations: usize, + auto_amplitude_profile: bool, ) -> PyResult { let shape = data.shape(); let n_samples = shape[0]; @@ -313,6 +320,7 @@ fn sort_multichannel<'py>( auto_svd_init, auto_threshold, auto_refine_iterations, + auto_amplitude_profile, }; // W=48 (captures full biphasic waveform), K=4 (3 PCA + 1 channel), diff --git a/src/sorter.rs b/src/sorter.rs index d034c24..cc8b78c 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -101,6 +101,7 @@ pub enum DetectionMode { /// assert!(!config.use_localization); /// assert!(!config.use_amplitude_profile); /// assert_eq!(config.amplitude_profile_neighbors, 4); +/// assert!(config.auto_amplitude_profile); /// assert!(config.auto_cmr); /// assert!(config.coincidence_detect); /// assert!((config.coincidence_primary_threshold - 3.5).abs() < 1e-12); @@ -265,6 +266,14 @@ pub struct SortConfig { /// The profile includes the peak channel + this many nearest neighbors. /// Default: 4 (5 channels total: peak + 4 neighbors). pub amplitude_profile_neighbors: usize, + /// Auto-activate the amplitude profile feature for recordings with C≥8 channels. + /// On large probes, multiple units share the same peak channel and half-width + /// cannot separate them. The amplitude profile encodes how much energy bleeds + /// into neighboring channels — a physics-based spatial fingerprint that varies + /// with unit position and depth even when peak channel is identical. + /// When true and C≥8 and K≥4, activates amplitude profile encoding regardless + /// of `use_amplitude_profile`. Has no effect for C<8. Default: true. + pub auto_amplitude_profile: bool, /// Auto-apply CMR when channel count >= 8. /// Common Median Reference removes shared noise across channels before /// whitening. For multi-channel recordings (C >= 8), correlated noise @@ -402,6 +411,7 @@ impl Default for SortConfig { use_localization: false, use_amplitude_profile: false, amplitude_profile_neighbors: 4, + auto_amplitude_profile: true, auto_cmr: true, coincidence_detect: true, coincidence_primary_threshold: 3.5, @@ -2939,7 +2949,7 @@ pub fn sort_multichannel< feature_buf[i][K - 1] = norm * spatial_scale; } } - } else if config.use_amplitude_profile && C > 1 && K >= 4 { + } else if (config.use_amplitude_profile || (config.auto_amplitude_profile && C >= 8)) && C > 1 && K >= 4 { // Two-feature spatial encoding: // - K-1: channel index (strong separation, same as fallback) // - K-2: amplitude profile ratio (fine spatial discrimination) @@ -4691,6 +4701,148 @@ mod tests { assert_eq!(config.template_min_count, 3); assert!(!config.use_amplitude_profile); assert_eq!(config.amplitude_profile_neighbors, 4); + assert!(config.auto_amplitude_profile); + } + + #[test] + fn test_auto_amplitude_profile_default() { + let config = SortConfig::default(); + assert!( + config.auto_amplitude_profile, + "auto_amplitude_profile must default to true" + ); + } + + #[test] + fn test_auto_amplitude_profile_activates_for_large_probe() { + // With auto_amplitude_profile=true and C=8, the amplitude profile branch + // must execute (use_amplitude_profile false, but auto gate fires). + // We verify by confirming the sort completes without panic and finds spikes. + let n = 4000; + let mut data = vec![[0.0f64; 8]; n]; + let mut rng = Rng::new(77); + for s in data.iter_mut() { + for v in s.iter_mut() { + *v = rng.gaussian(0.0, 1.0); + } + } + // Inject a strong spike on channel 3 with some energy on neighbors + let mut pos = 200; + while pos + 8 < n { + for dt in 0..8 { + if pos + dt < n { + data[pos + dt][3] += + -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + data[pos + dt][4] += + -6.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + } + } + pos += 160; + } + let probe = ProbeLayout::<8>::linear(25.0); + let config = SortConfig { + use_amplitude_profile: false, + auto_amplitude_profile: true, + threshold_multiplier: 4.5, + pre_samples: 2, + refractory_samples: 10, + ..SortConfig::default() + }; + let mut scratch = vec![0.0; n]; + let mut events = vec![ + MultiChannelEvent { + sample: 0, + channel: 0, + amplitude: 0.0, + }; + 400 + ]; + let mut wf = vec![[0.0; 8]; 400]; + let mut feat = vec![[0.0; 4]; 400]; + let mut lab = vec![0usize; 400]; + let result = sort_multichannel::<8, 64, 8, 4, 64, 16>( + &config, + &probe, + &mut data, + &mut scratch, + &mut events, + &mut wf, + &mut feat, + &mut lab, + ); + assert!( + result.is_ok(), + "auto_amplitude_profile sort must not error" + ); + let sr = result.unwrap(); + assert!( + sr.n_spikes > 0, + "auto_amplitude_profile must detect injected spikes" + ); + } + + #[test] + fn test_auto_amplitude_profile_no_effect_small_probe() { + // With C=4 (< 8), auto_amplitude_profile gate must not fire. + // Sort must complete and the feature encoding falls through to the + // half-width / channel-index fallback branch (no panic, spikes detected). + let n = 4000; + let mut data = vec![[0.0f64; 4]; n]; + let mut rng = Rng::new(55); + for s in data.iter_mut() { + for v in s.iter_mut() { + *v = rng.gaussian(0.0, 1.0); + } + } + let mut pos = 200; + while pos + 8 < n { + for dt in 0..8 { + if pos + dt < n { + data[pos + dt][1] += + -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + } + } + pos += 160; + } + let probe = ProbeLayout::<4>::linear(25.0); + let config = SortConfig { + auto_amplitude_profile: true, + threshold_multiplier: 4.5, + pre_samples: 2, + refractory_samples: 10, + ..SortConfig::default() + }; + let mut scratch = vec![0.0; n]; + let mut events = vec![ + MultiChannelEvent { + sample: 0, + channel: 0, + amplitude: 0.0, + }; + 400 + ]; + let mut wf = vec![[0.0; 8]; 400]; + let mut feat = vec![[0.0; 4]; 400]; + let mut lab = vec![0usize; 400]; + let result = sort_multichannel::<4, 16, 8, 4, 64, 16>( + &config, + &probe, + &mut data, + &mut scratch, + &mut events, + &mut wf, + &mut feat, + &mut lab, + ); + assert!( + result.is_ok(), + "small probe sort with auto_amplitude_profile must not error" + ); + let sr = result.unwrap(); + assert!( + sr.n_spikes > 0, + "small probe sort must still detect spikes" + ); } #[test] From a6c821accb3109a78afd62b7564ad606a1ae3451 Mon Sep 17 00:00:00 2001 From: Fredrik Willumsen Haug Date: Sun, 19 Apr 2026 16:44:51 -0400 Subject: [PATCH 2/2] [format] CI fix --- src/sorter.rs | 24 +++++++++--------------- 1 file changed, 9 insertions(+), 15 deletions(-) diff --git a/src/sorter.rs b/src/sorter.rs index cc8b78c..68eef97 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -2949,7 +2949,10 @@ pub fn sort_multichannel< feature_buf[i][K - 1] = norm * spatial_scale; } } - } else if (config.use_amplitude_profile || (config.auto_amplitude_profile && C >= 8)) && C > 1 && K >= 4 { + } else if (config.use_amplitude_profile || (config.auto_amplitude_profile && C >= 8)) + && C > 1 + && K >= 4 + { // Two-feature spatial encoding: // - K-1: channel index (strong separation, same as fallback) // - K-2: amplitude profile ratio (fine spatial discrimination) @@ -4731,10 +4734,8 @@ mod tests { while pos + 8 < n { for dt in 0..8 { if pos + dt < n { - data[pos + dt][3] += - -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); - data[pos + dt][4] += - -6.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + data[pos + dt][3] += -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + data[pos + dt][4] += -6.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); } } pos += 160; @@ -4770,10 +4771,7 @@ mod tests { &mut feat, &mut lab, ); - assert!( - result.is_ok(), - "auto_amplitude_profile sort must not error" - ); + assert!(result.is_ok(), "auto_amplitude_profile sort must not error"); let sr = result.unwrap(); assert!( sr.n_spikes > 0, @@ -4798,8 +4796,7 @@ mod tests { while pos + 8 < n { for dt in 0..8 { if pos + dt < n { - data[pos + dt][1] += - -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); + data[pos + dt][1] += -12.0 * f64::exp(-0.5 * ((dt as f64 - 2.0) / 1.5).powi(2)); } } pos += 160; @@ -4839,10 +4836,7 @@ mod tests { "small probe sort with auto_amplitude_profile must not error" ); let sr = result.unwrap(); - assert!( - sr.n_spikes > 0, - "small probe sort must still detect spikes" - ); + assert!(sr.n_spikes > 0, "small probe sort must still detect spikes"); } #[test]