From a8f22906d4f701d9695107aa678c362e874ffc1b Mon Sep 17 00:00:00 2001 From: Fredrik Willumsen Haug Date: Tue, 21 Apr 2026 11:25:59 -0400 Subject: [PATCH 1/2] [sorter] raise spatial_merge_dprime default from 1.5 to 2.0 --- .../amplitude_profile_scale_benchmark.py | 128 +++++++++++++++ benchmarks/diagnostic_benchmark.py | 143 ++++++++++++++++ benchmarks/spatial_merge_gmm_benchmark.py | 155 ++++++++++++++++++ benchmarks/verify_spatial_merge_benchmark.py | 100 +++++++++++ 4 files changed, 526 insertions(+) create mode 100644 benchmarks/amplitude_profile_scale_benchmark.py create mode 100644 benchmarks/diagnostic_benchmark.py create mode 100644 benchmarks/spatial_merge_gmm_benchmark.py create mode 100644 benchmarks/verify_spatial_merge_benchmark.py diff --git a/benchmarks/amplitude_profile_scale_benchmark.py b/benchmarks/amplitude_profile_scale_benchmark.py new file mode 100644 index 0000000..391f908 --- /dev/null +++ b/benchmarks/amplitude_profile_scale_benchmark.py @@ -0,0 +1,128 @@ +"""Benchmark: amplitude profile scale tuning. + +The amplitude_profile feature uses profile_scale = cluster_threshold * 2.0 = 14. +This scale may be too large, dominating the PCA component in K-2 and causing +cluster merging errors. Test lower scales by varying cluster_threshold with +use_amplitude_profile=True directly. + +Since profile_scale = cluster_threshold * 2.0, we can't decouple them directly. +Instead, test use_amplitude_profile at different cluster_threshold values and +compare to half-width (default) at the same threshold. +""" + +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 = {} + +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", + matched_filter_detect=True, matched_filter_threshold=3.5, + 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=False, + ) + 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:<70} 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) + print(f" {label:<70} acc={acc:.3f} wd={wd}/{n_gt} {n_spikes}spk/{n_cl}cl {elapsed:.1f}s") + if detail: + per_unit_df = cmp.get_performance() + for uid in per_unit_df.index: + print(f" unit {uid}: acc={float(per_unit_df.loc[uid]['accuracy']):.3f}") + return acc + + +def main(): + for difficulty in ["medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + detail = difficulty == "hard" + + # Baseline: half-width (no amplitude profile) + run(difficulty, "baseline (half-width, ct=7.0)", + use_amplitude_profile=False, detail=detail) + + # Amplitude profile with standard cluster_threshold=7.0 + run(difficulty, "amp_profile (ct=7.0, scale=14.0)", + use_amplitude_profile=True, detail=detail) + + # Try lower cluster thresholds to reduce profile_scale + for ct in [5.0, 4.0, 3.5]: + ps = ct * 2.0 + run(difficulty, f"amp_profile (ct={ct}, scale={ps})", + use_amplitude_profile=True, cluster_threshold=ct, detail=detail) + + # Try more neighbors + run(difficulty, "amp_profile (ct=7.0, nbr=2)", + use_amplitude_profile=True, amplitude_profile_neighbors=2, detail=detail) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/diagnostic_benchmark.py b/benchmarks/diagnostic_benchmark.py new file mode 100644 index 0000000..f1d9759 --- /dev/null +++ b/benchmarks/diagnostic_benchmark.py @@ -0,0 +1,143 @@ +"""Diagnostic: understand where the remaining accuracy loss comes from. + +For each difficulty, examine: +1. Baseline per-unit accuracy breakdown +2. Per-unit recall vs precision to identify whether the loss is misses or false assigns +3. Effect of merge_dprime_threshold (are units being over-merged?) +4. Effect of split thresholds (are units being under-split?) +""" + +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 = {} + + +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_and_analyze(difficulty, label, **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", + matched_filter_detect=True, matched_filter_threshold=3.5, + 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, + auto_threshold=True, auto_refine_iterations=3, ncc_threshold=0.70, + template_subtract_passes=2, auto_amplitude_profile=False, + ) + 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}: 0 spikes") + 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) + print(f"\n {label}: acc={acc:.3f} wd={wd}/{n_gt} {n_spikes}spk/{n_cl}cl {elapsed:.1f}s") + + per_unit = cmp.get_performance() + for uid in per_unit.index: + r = float(per_unit.loc[uid]["recall"]) + p = float(per_unit.loc[uid]["precision"]) + a = float(per_unit.loc[uid]["accuracy"]) + gt_count = len(gt_trains.get(uid, [])) + bottleneck = "RECALL" if r < p else "PRECISION" if p < r else "EQUAL" + if a < 0.01: + bottleneck = "UNDETECTED" + print(f" unit {uid:>2}: acc={a:.3f} recall={r:.3f} prec={p:.3f} gt={gt_count:>4} [{bottleneck}]") + + return acc + + +def main(): + for difficulty in ["medium", "hard"]: + print(f"\n{'='*80}") + print(f" {difficulty.upper()}") + print(f"{'='*80}") + + # Baseline + run_and_analyze(difficulty, "BASELINE (defaults)") + + # Looser merge (more merging) + run_and_analyze(difficulty, "merge_dprime=1.5 (more merging)", + merge_dprime_threshold=1.5) + + # Tighter merge (less merging) + run_and_analyze(difficulty, "merge_dprime=3.0 (less merging)", + merge_dprime_threshold=3.0) + + # Spatial merge tighter + run_and_analyze(difficulty, "spatial_merge_dprime=2.5 (less spatial merge)", + spatial_merge_dprime=2.5) + + # No CCG merge + run_and_analyze(difficulty, "ccg_merge=False", + ccg_merge=False) + + # Lower min_cluster_snr + run_and_analyze(difficulty, "min_cluster_snr=1.5", + min_cluster_snr=1.5) + + # GMM refine + run_and_analyze(difficulty, "gmm_refine=True", + gmm_refine=True) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/spatial_merge_gmm_benchmark.py b/benchmarks/spatial_merge_gmm_benchmark.py new file mode 100644 index 0000000..9961948 --- /dev/null +++ b/benchmarks/spatial_merge_gmm_benchmark.py @@ -0,0 +1,155 @@ +"""Benchmark: spatial_merge_dprime + gmm_refine tuning. + +Diagnostic found: +- Medium: spatial_merge_dprime=2.5 → +3.2% (over-merging at 1.5) +- Medium: gmm_refine=True → +4.1% +- Hard: units 1,8,12 are precision-limited (noise spikes misassigned) + +Test combinations and whether they can be auto-gated. +""" + +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.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", + matched_filter_detect=True, matched_filter_threshold=3.5, + 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, + auto_threshold=True, auto_refine_iterations=3, ncc_threshold=0.70, + template_subtract_passes=2, auto_amplitude_profile=False, + ) + 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:<72} 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:<72} 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"]) + 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(): + results = {} + for difficulty in ["easy", "medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + detail = difficulty in ("medium", "hard") + r = {} + + r["baseline"] = run(difficulty, "baseline", detail=detail) + + # spatial_merge_dprime sweep + for smd in [2.0, 2.5, 3.0]: + r[f"smd_{smd}"] = run(difficulty, f"spatial_merge_dprime={smd}", + spatial_merge_dprime=smd, detail=detail) + + # gmm_refine alone + r["gmm"] = run(difficulty, "gmm_refine=True", + gmm_refine=True, detail=detail) + + # gmm + spatial_merge_dprime=2.5 + r["gmm_smd25"] = run(difficulty, "gmm_refine + spatial_merge_dprime=2.5", + gmm_refine=True, spatial_merge_dprime=2.5, detail=detail) + + # gmm + spatial_merge_dprime=2.0 + r["gmm_smd20"] = run(difficulty, "gmm_refine + spatial_merge_dprime=2.0", + gmm_refine=True, spatial_merge_dprime=2.0, detail=detail) + + results[difficulty] = r + + print(f"\n{'='*90}") + print(f"{'Config':<55} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*90}") + for label, key in [ + ("baseline (current defaults)", "baseline"), + ("spatial_merge_dprime=2.0", "smd_2.0"), + ("spatial_merge_dprime=2.5", "smd_2.5"), + ("spatial_merge_dprime=3.0", "smd_3.0"), + ("gmm_refine=True", "gmm"), + ("gmm + spatial_merge_dprime=2.5", "gmm_smd25"), + ("gmm + spatial_merge_dprime=2.0", "gmm_smd20"), + ]: + e = results["easy"].get(key, 0) + m = results["medium"].get(key, 0) + h = results["hard"].get(key, 0) + avg = (e + m + h) / 3.0 + print(f" {label:<53} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*90}") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/verify_spatial_merge_benchmark.py b/benchmarks/verify_spatial_merge_benchmark.py new file mode 100644 index 0000000..c64144b --- /dev/null +++ b/benchmarks/verify_spatial_merge_benchmark.py @@ -0,0 +1,100 @@ +"""Verify: spatial_merge_dprime=2.0 (new default) vs 1.5 (old default). + +Expected: Medium +3.1%, Easy unchanged, Hard flat. +""" + +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 = {} + + +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, **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") + 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) + print(f" {label:<50} acc={acc:.3f} wd={wd}/{n_gt} {n_spikes}spk/{n_cl}cl {elapsed:.1f}s") + return acc + + +def main(): + results = {} + for difficulty in ["easy", "medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + # New default (spatial_merge_dprime=2.0 is now the default) + acc_new = run(difficulty, "NEW DEFAULT (spatial_merge_dprime=2.0)") + # Old default + acc_old = run(difficulty, "OLD DEFAULT (spatial_merge_dprime=1.5)", + spatial_merge_dprime=1.5) + results[difficulty] = {"new": acc_new, "old": acc_old} + + print(f"\n{'='*70}") + e_old, e_new = results["easy"]["old"], results["easy"]["new"] + m_old, m_new = results["medium"]["old"], results["medium"]["new"] + h_old, h_new = results["hard"]["old"], results["hard"]["new"] + avg_old = (e_old + m_old + h_old) / 3 + avg_new = (e_new + m_new + h_new) / 3 + print(f" OLD (smd=1.5): Easy {e_old:.3f} Medium {m_old:.3f} Hard {h_old:.3f} Avg {avg_old:.3f}") + print(f" NEW (smd=2.0): Easy {e_new:.3f} Medium {m_new:.3f} Hard {h_new:.3f} Avg {avg_new:.3f}") + print(f" Delta: Easy {e_new-e_old:+.3f} Medium {m_new-m_old:+.3f} Hard {h_new-h_old:+.3f} Avg {avg_new-avg_old:+.3f}") + print(f"{'='*70}") + + +if __name__ == "__main__": + main() From 17f99ab94fbc4e395fc4a3b77a0077032ed42179 Mon Sep 17 00:00:00 2001 From: Fredrik Willumsen Haug Date: Tue, 21 Apr 2026 11:26:14 -0400 Subject: [PATCH 2/2] [sorter] add auto_amplitude_profile field (default: false) --- python/src/sorter.rs | 6 +++--- src/sorter.rs | 19 ++++++++----------- 2 files changed, 11 insertions(+), 14 deletions(-) diff --git a/python/src/sorter.rs b/python/src/sorter.rs index 803c361..804d0e6 100644 --- a/python/src/sorter.rs +++ b/python/src/sorter.rs @@ -42,7 +42,7 @@ fn sort_error_to_py(e: SortError) -> PyErr { /// merge_isi_threshold (float): ISI violation threshold for cluster merging. Default: 0.05. /// split_min_cluster_size (int): Minimum spikes per cluster to attempt splitting. Default: 10. /// split_bimodality_threshold (float): Gap/std threshold for cluster splitting. Default: 2.0. -/// spatial_merge_dprime (float): D-prime threshold for cross-channel spatial merge. Default: 1.5. +/// spatial_merge_dprime (float): D-prime threshold for cross-channel spatial merge. Default: 2.0. /// template_subtract (bool): Enable template subtraction to recover masked spikes. Default: True. /// template_min_count (int): Minimum spikes per cluster to build a subtraction template. Default: 3. /// min_cluster_snr (float): Minimum SNR for cluster auto-curation. Default: 2.5. @@ -138,7 +138,7 @@ fn sort_error_to_py(e: SortError) -> PyErr { merge_isi_threshold = 0.05, split_min_cluster_size = 10, split_bimodality_threshold = 2.0, - spatial_merge_dprime = 1.5, + spatial_merge_dprime = 2.0, template_subtract = true, template_min_count = 3, min_cluster_snr = 2.5, @@ -181,7 +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, + auto_amplitude_profile = false, ))] #[allow(clippy::too_many_arguments)] fn sort_multichannel<'py>( diff --git a/src/sorter.rs b/src/sorter.rs index 68eef97..63724c8 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -82,7 +82,7 @@ pub enum DetectionMode { /// assert!((config.merge_isi_threshold - 0.05).abs() < 1e-12); /// assert_eq!(config.split_min_cluster_size, 10); /// assert!((config.split_bimodality_threshold - 2.0).abs() < 1e-12); -/// assert!((config.spatial_merge_dprime - 1.5).abs() < 1e-12); +/// assert!((config.spatial_merge_dprime - 2.0).abs() < 1e-12); /// assert!(config.template_subtract); /// assert_eq!(config.template_min_count, 3); /// assert!((config.min_cluster_snr - 2.5).abs() < 1e-12); @@ -101,7 +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_amplitude_profile); /// assert!(config.auto_cmr); /// assert!(config.coincidence_detect); /// assert!((config.coincidence_primary_threshold - 3.5).abs() < 1e-12); @@ -382,7 +382,7 @@ impl Default for SortConfig { merge_isi_threshold: 0.05, split_min_cluster_size: 10, split_bimodality_threshold: 2.0, - spatial_merge_dprime: 1.5, + spatial_merge_dprime: 2.0, template_subtract: true, template_min_count: 3, min_cluster_snr: 2.5, @@ -411,7 +411,7 @@ impl Default for SortConfig { use_localization: false, use_amplitude_profile: false, amplitude_profile_neighbors: 4, - auto_amplitude_profile: true, + auto_amplitude_profile: false, auto_cmr: true, coincidence_detect: true, coincidence_primary_threshold: 3.5, @@ -2949,10 +2949,7 @@ 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 && 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) @@ -4704,15 +4701,15 @@ 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); + 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" + !config.auto_amplitude_profile, + "auto_amplitude_profile must default to false" ); }