diff --git a/benchmarks/auto_gmm_refine_benchmark.py b/benchmarks/auto_gmm_refine_benchmark.py new file mode 100644 index 0000000..b2ae954 --- /dev/null +++ b/benchmarks/auto_gmm_refine_benchmark.py @@ -0,0 +1,128 @@ +"""Benchmark: auto_gmm_refine — auto-gate GMM refinement for C>=8. + +Diagnostic found gmm_refine=True gives Medium +4.1% but Easy -3.4%. +Easy has C=4, Medium C=16, Hard C=32. An auto-gate for C>=8 would +capture the medium improvement without harming easy. + +Question: does GMM actually help medium with the new spatial_merge_dprime=2.0? +Previous test was with smd=1.5. Need to reconfirm. +""" + +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.7650, "hard": 0.6400} + + +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:<65} 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:<65} 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 (new defaults, smd=2.0)", detail=detail) + + r["gmm"] = run(difficulty, "gmm_refine=True", + gmm_refine=True, detail=detail) + + r["gmm_iter20"] = run(difficulty, "gmm_refine=True, gmm_max_iter=20", + gmm_refine=True, gmm_max_iter=20, detail=detail) + + r["gmm_iter5"] = run(difficulty, "gmm_refine=True, gmm_max_iter=5", + gmm_refine=True, gmm_max_iter=5, detail=detail) + + results[difficulty] = r + + print(f"\n{'='*85}") + print(f"{'Config':<50} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*85}") + for label, key in [ + ("baseline (smd=2.0 defaults)", "baseline"), + ("gmm_refine=True", "gmm"), + ("gmm_refine, max_iter=20", "gmm_iter20"), + ("gmm_refine, max_iter=5", "gmm_iter5"), + ]: + e = results["easy"][key] + m = results["medium"][key] + h = results["hard"][key] + avg = (e + m + h) / 3.0 + print(f" {label:<48} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*85}") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/precision_sweep_benchmark.py b/benchmarks/precision_sweep_benchmark.py new file mode 100644 index 0000000..09ecd4c --- /dev/null +++ b/benchmarks/precision_sweep_benchmark.py @@ -0,0 +1,151 @@ +"""Benchmark: targeting precision-limited hard units 1, 8, 12. + +Unit 1: R=0.927, P=0.544 — cluster absorbs ~40% false positives +Unit 8: R=0.938, P=0.655 — ~30% false positives +Unit 12: R=0.996, P=0.718 — ~28% false positives + +These clusters have correct spikes but too many wrong ones mixed in. +Try: isi_split_threshold, cluster_threshold, matched_filter_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 = {} +BASELINE = {"easy": 0.8270, "medium": 0.7650, "hard": 0.6400} + + +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:<68} 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:<68} 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"]) + if a < 0.01: + continue + 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 == "hard" + r = {} + + r["baseline"] = run(difficulty, "baseline", detail=detail) + + # ISI split: more aggressive splitting of contaminated clusters + for ist in [0.08, 0.05, 0.03]: + r[f"ist_{ist}"] = run(difficulty, f"isi_split_threshold={ist}", + isi_split_threshold=ist, detail=detail) + + # Lower cluster_threshold: create more initial clusters + for ct in [6.0, 5.0]: + r[f"ct_{ct}"] = run(difficulty, f"cluster_threshold={ct}", + cluster_threshold=ct, detail=detail) + + # Higher matched_filter_threshold: fewer false-positive MF detections + for mft in [4.0, 4.5, 5.0]: + r[f"mft_{mft}"] = run(difficulty, f"matched_filter_threshold={mft}", + matched_filter_threshold=mft, detail=detail) + + # Tighter merge_dprime (reduce over-merging that mixes units) + r["mdp_2.5"] = run(difficulty, "merge_dprime_threshold=2.5", + merge_dprime_threshold=2.5, detail=detail) + + # refine_isi_guard=True (prevent refinement from contaminating) + r["isi_guard"] = run(difficulty, "refine_isi_guard=True", + refine_isi_guard=True, detail=detail) + + results[difficulty] = r + + print(f"\n{'='*85}") + print(f"{'Config':<50} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*85}") + for label, key in [ + ("baseline", "baseline"), + ("isi_split=0.08", "ist_0.08"), + ("isi_split=0.05", "ist_0.05"), + ("isi_split=0.03", "ist_0.03"), + ("cluster_threshold=6.0", "ct_6.0"), + ("cluster_threshold=5.0", "ct_5.0"), + ("matched_filter_threshold=4.0", "mft_4.0"), + ("matched_filter_threshold=4.5", "mft_4.5"), + ("matched_filter_threshold=5.0", "mft_5.0"), + ("merge_dprime=2.5", "mdp_2.5"), + ("refine_isi_guard=True", "isi_guard"), + ]: + 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:<48} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*85}") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/split_ccg_sweep_benchmark.py b/benchmarks/split_ccg_sweep_benchmark.py new file mode 100644 index 0000000..3e70aa7 --- /dev/null +++ b/benchmarks/split_ccg_sweep_benchmark.py @@ -0,0 +1,154 @@ +"""Benchmark: split_bimodality and ccg_template_corr thresholds. + +Hard units 1, 8, 12 are precision-limited. Maybe the contamination +comes from bad merges (CCG or d-prime) or insufficient splitting. +Also check split_min_cluster_size — maybe we're not splitting small +contaminated clusters. +""" + +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.7650, "hard": 0.6400} + + +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:<62} 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:<62} 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(): + results = {} + for difficulty in ["easy", "medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + detail = difficulty == "hard" + r = {} + + r["baseline"] = run(difficulty, "baseline", detail=detail) + + # More aggressive splitting + for sbt in [1.5, 1.0]: + r[f"sbt_{sbt}"] = run(difficulty, f"split_bimodality_threshold={sbt}", + split_bimodality_threshold=sbt, detail=detail) + + # Smaller min cluster size for splits + for smcs in [5, 3]: + r[f"smcs_{smcs}"] = run(difficulty, f"split_min_cluster_size={smcs}", + split_min_cluster_size=smcs, detail=detail) + + # CCG template correlation threshold + for ccg in [0.7, 0.8, 0.9]: + r[f"ccg_{ccg}"] = run(difficulty, f"ccg_template_corr_threshold={ccg}", + ccg_template_corr_threshold=ccg, detail=detail) + + # Disable template subtraction entirely + r["no_ts"] = run(difficulty, "template_subtract=False", + template_subtract=False, detail=detail) + + # Disable matched filter + r["no_mf"] = run(difficulty, "matched_filter_detect=False", + matched_filter_detect=False, detail=detail) + + # Combination: aggressive split + tighter CCG + r["combo"] = run(difficulty, "split_bim=1.5 + ccg_corr=0.8", + split_bimodality_threshold=1.5, + ccg_template_corr_threshold=0.8, detail=detail) + + results[difficulty] = r + + print(f"\n{'='*85}") + print(f"{'Config':<50} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*85}") + for label, key in [ + ("baseline", "baseline"), + ("split_bimodality=1.5", "sbt_1.5"), + ("split_bimodality=1.0", "sbt_1.0"), + ("split_min_cluster=5", "smcs_5"), + ("split_min_cluster=3", "smcs_3"), + ("ccg_corr=0.7", "ccg_0.7"), + ("ccg_corr=0.8", "ccg_0.8"), + ("ccg_corr=0.9", "ccg_0.9"), + ("no template subtract", "no_ts"), + ("no matched filter", "no_mf"), + ("split_bim=1.5 + ccg_corr=0.8", "combo"), + ]: + 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:<48} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*85}") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/waveform_reassign_margin_benchmark.py b/benchmarks/waveform_reassign_margin_benchmark.py new file mode 100644 index 0000000..efdc25a --- /dev/null +++ b/benchmarks/waveform_reassign_margin_benchmark.py @@ -0,0 +1,121 @@ +"""Benchmark: waveform_reassign_margin sweep. + +Step 9h reassigns spikes to the nearest same-channel template, but only +if the distance improvement exceeds the margin. Current default: 0.3 (30%). +Test: lowering the margin lets more borderline spikes be reassigned, +potentially cleaning up precision-limited units. +""" + +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.7650, "hard": 0.6400} + + +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(): + results = {} + for difficulty in ["easy", "medium", "hard"]: + print(f"\n=== {difficulty.upper()} ===") + detail = difficulty == "hard" + r = {} + for margin in [0.3, 0.2, 0.1, 0.05, 0.0]: + label = f"waveform_reassign_margin={margin}" + if margin == 0.3: + label += " (current default)" + r[f"m_{margin}"] = run(difficulty, label, + waveform_reassign_margin=margin, detail=detail) + results[difficulty] = r + + print(f"\n{'='*80}") + print(f"{'Margin':<45} {'Easy':>7} {'Medium':>7} {'Hard':>7} {'Avg':>7}") + print(f"{'='*80}") + for margin in [0.3, 0.2, 0.1, 0.05, 0.0]: + key = f"m_{margin}" + e = results["easy"][key] + m = results["medium"][key] + h = results["hard"][key] + avg = (e + m + h) / 3.0 + label = f"margin={margin}" + if margin == 0.3: + label += " (current)" + print(f" {label:<43} {e:>7.3f} {m:>7.3f} {h:>7.3f} {avg:>7.3f}") + print(f"{'='*80}") + + +if __name__ == "__main__": + main() diff --git a/python/src/sorter.rs b/python/src/sorter.rs index 804d0e6..a057b8f 100644 --- a/python/src/sorter.rs +++ b/python/src/sorter.rs @@ -102,6 +102,10 @@ fn sort_error_to_py(e: SortError) -> PyErr { /// 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. +/// waveform_reassign_margin (float): Margin for template-based waveform reassignment. +/// A spike is only reassigned to a closer same-channel template if the new L2 +/// distance is less than (1.0 - margin) times the old distance. Higher = more +/// conservative. 0.0 = always take nearest. Default: 0.05. /// /// Returns: /// dict: Sorting results with keys: @@ -182,6 +186,7 @@ fn sort_error_to_py(e: SortError) -> PyErr { auto_threshold = true, auto_refine_iterations = 3usize, auto_amplitude_profile = false, + waveform_reassign_margin = 0.05, ))] #[allow(clippy::too_many_arguments)] fn sort_multichannel<'py>( @@ -245,6 +250,7 @@ fn sort_multichannel<'py>( auto_threshold: bool, auto_refine_iterations: usize, auto_amplitude_profile: bool, + waveform_reassign_margin: f64, ) -> PyResult { let shape = data.shape(); let n_samples = shape[0]; @@ -321,6 +327,7 @@ fn sort_multichannel<'py>( auto_threshold, auto_refine_iterations, auto_amplitude_profile, + waveform_reassign_margin, }; // W=48 (captures full biphasic waveform), K=4 (3 PCA + 1 channel), diff --git a/src/sorter.rs b/src/sorter.rs index 63724c8..1980290 100644 --- a/src/sorter.rs +++ b/src/sorter.rs @@ -116,6 +116,7 @@ pub enum DetectionMode { /// assert!(!config.refine_isi_guard); /// assert!((config.refine_isi_tolerance - 0.1).abs() < 1e-12); /// assert!(config.ccg_merge); +/// assert!((config.waveform_reassign_margin - 0.05).abs() < 1e-12); /// ``` pub struct SortConfig { /// Threshold multiplier for spike detection (sigma units on whitened data). @@ -364,6 +365,13 @@ pub struct SortConfig { /// skip the iteration if violations increase by more than 10%. /// Only active when `refine_isi_guard` is true. Default: 0.1. pub refine_isi_tolerance: Float, + /// Margin for template-based waveform reassignment (step 9h). + /// A spike is only reassigned to a closer same-channel template if the + /// new L2 distance is less than `(1.0 - margin)` times the old distance. + /// Higher values require bigger improvements to reassign (more conservative). + /// 0.0 = always take the nearest template (most aggressive). + /// 0.3 = require 30% distance improvement. Default: 0.05. + pub waveform_reassign_margin: Float, } impl Default for SortConfig { @@ -425,6 +433,7 @@ impl Default for SortConfig { refine_collapse_guard: true, refine_isi_guard: false, refine_isi_tolerance: 0.1, + waveform_reassign_margin: 0.05, } } } @@ -3900,9 +3909,10 @@ pub fn sort_multichannel< best_c = c; } } - // Only reassign if new template is significantly closer (30% margin) - // This prevents marginal reassignments from breaking good clustering. - if best_c != old_label && best_d < old_d * 0.7 { + // 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; }