From fd899c13ca0af6daf996e25555fca0448493bb93 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Royeth?= Date: Fri, 4 Sep 2026 12:38:33 -0400 Subject: [PATCH] Use channel offsets, not gains, when scaling amplitudes to uV --- .../postprocessing/amplitude_scalings.py | 2 +- .../postprocessing/spike_amplitudes.py | 2 +- .../tests/test_spike_amplitudes.py | 27 +++++++++++++++++++ 3 files changed, 29 insertions(+), 2 deletions(-) diff --git a/src/spikeinterface/postprocessing/amplitude_scalings.py b/src/spikeinterface/postprocessing/amplitude_scalings.py index 9a1ae02431..1bf33a1910 100644 --- a/src/spikeinterface/postprocessing/amplitude_scalings.py +++ b/src/spikeinterface/postprocessing/amplitude_scalings.py @@ -179,7 +179,7 @@ def __init__( if return_in_uV and recording.has_scaleable_traces(): self._dtype = np.float32 self._gains = recording.get_channel_gains() - self._offsets = recording.get_channel_gains() + self._offsets = recording.get_channel_offsets() else: self._dtype = recording.get_dtype() self._gains = None diff --git a/src/spikeinterface/postprocessing/spike_amplitudes.py b/src/spikeinterface/postprocessing/spike_amplitudes.py index a1b6fef8d2..3c2020a64e 100644 --- a/src/spikeinterface/postprocessing/spike_amplitudes.py +++ b/src/spikeinterface/postprocessing/spike_amplitudes.py @@ -69,7 +69,7 @@ def __init__( if return_in_uV and recording.has_scaleable_traces(): self._dtype = np.float32 self._gains = recording.get_channel_gains() - self._offsets = recording.get_channel_gains() + self._offsets = recording.get_channel_offsets() else: self._dtype = recording.get_dtype() self._gains = None diff --git a/src/spikeinterface/postprocessing/tests/test_spike_amplitudes.py b/src/spikeinterface/postprocessing/tests/test_spike_amplitudes.py index a68483a1b2..f32bf90ed5 100644 --- a/src/spikeinterface/postprocessing/tests/test_spike_amplitudes.py +++ b/src/spikeinterface/postprocessing/tests/test_spike_amplitudes.py @@ -1,8 +1,35 @@ +import numpy as np + +from spikeinterface.core import NumpyRecording, NumpySorting +from spikeinterface.core.node_pipeline import SpikeRetriever from spikeinterface.postprocessing import ComputeSpikeAmplitudes from spikeinterface.postprocessing.tests.common_extension_tests import AnalyzerExtensionCommonTestSuite +from spikeinterface.postprocessing.spike_amplitudes import SpikeAmplitudeNode class TestComputeSpikeAmplitudes(AnalyzerExtensionCommonTestSuite): def test_extension(self): self.run_extension_tests(ComputeSpikeAmplitudes, params=dict()) + + def test_scaled_amplitude_uses_channel_offset(self): + traces = np.zeros((200, 2), dtype=np.int16) + traces[100, 0] = 10 + recording = NumpyRecording([traces], sampling_frequency=1000) + recording.set_property("gain_to_uV", [2.0, 3.0]) + recording.set_property("offset_to_uV", [-100.0, -200.0]) + sorting = NumpySorting.from_times_and_labels( + [np.array([100], dtype=np.int64)], + [np.array([1], dtype=np.int64)], + sampling_frequency=1000, + unit_ids=[1], + ) + sorting.set_property("main_channel_id", [0]) + spike_retriever = SpikeRetriever(sorting, recording, channel_from_template=True) + node = SpikeAmplitudeNode(recording, parents=[spike_retriever], peak_shifts={1: 0}, return_in_uV=True) + peaks = np.zeros(1, dtype=[("sample_index", "int64"), ("unit_index", "int64"), ("channel_index", "int64")]) + peaks[0] = (100, 0, 0) + + amplitudes = node.compute(traces, peaks) + + np.testing.assert_array_equal(amplitudes, [-80.0])