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
171 changes: 171 additions & 0 deletions benchmarks/auto_amplitude_profile_benchmark.py
Original file line number Diff line number Diff line change
@@ -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()
8 changes: 8 additions & 0 deletions python/src/sorter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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>(
Expand Down Expand Up @@ -238,6 +244,7 @@ fn sort_multichannel<'py>(
auto_svd_init: bool,
auto_threshold: bool,
auto_refine_iterations: usize,
auto_amplitude_profile: bool,
) -> PyResult<PyObject> {
let shape = data.shape();
let n_samples = shape[0];
Expand Down Expand Up @@ -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),
Expand Down
148 changes: 147 additions & 1 deletion src/sorter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -2939,7 +2949,10 @@ 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)
Expand Down Expand Up @@ -4691,6 +4704,139 @@ 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]
Expand Down
Loading