Skip to content
Open
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
93 changes: 71 additions & 22 deletions src/spikeinterface/preprocessing/detect_bad_channels.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
from typing import Literal

from spikeinterface.core.core_tools import define_function_handling_dict_from_class
from spikeinterface.core.job_tools import TimeSeriesChunkExecutor, fix_job_kwargs
from spikeinterface.core.time_series_tools import get_random_sample_slices
from .filter import highpass_filter
from spikeinterface.core import get_random_data_chunks, order_channels_by_depth, BaseRecording
from spikeinterface.core.channelslice import ChannelSliceRecording
Expand Down Expand Up @@ -75,6 +77,12 @@
The random seed to extract chunks
channel_filters : set | None, default: None
For coherence+psd - only return `bad_channel_ids` whose labels are in the set `channel_filter`.
job_kwargs : dict | None, default: None
Keyword arguments for parallel processing. Only used for the "coherence+psd" method. Only the
execution-related keys (`pool_engine`, `n_jobs`, `progress_bar`, `mp_context`,
`max_threads_per_worker`) apply; the chunking size is fixed by `chunk_duration_s` and
`num_random_chunks` above, so `chunk_size`, `chunk_memory`, `total_memory` and `chunk_duration`
are not used here.
"""


Expand Down Expand Up @@ -153,6 +161,39 @@ def _get_all_detect_bad_channel_kwargs(detect_bad_channels_kwargs):
return all_detect_bad_channels_kwargs


def _detect_bad_channels_chunk_init(recording, method_kwargs):
return {"recording": recording, "method_kwargs": method_kwargs}


def _detect_bad_channels_chunk(segment_index, start_frame, end_frame, worker_context):
recording = worker_context["recording"]
method_kwargs = worker_context["method_kwargs"]

random_chunk = recording.get_traces(
start_frame=start_frame,
end_frame=end_frame,
segment_index=segment_index,
return_in_uV=True,
)

order_f = method_kwargs["order_f"]
order_r = method_kwargs["order_r"]
random_chunk_sorted = random_chunk[:, order_f] if order_f is not None else random_chunk
chunk_labels = detect_bad_channels_ibl(
raw=random_chunk_sorted,
fs=recording.sampling_frequency,
psd_hf_threshold=method_kwargs["psd_hf_threshold"],
dead_channel_thr=method_kwargs["dead_channel_threshold"],
noisy_channel_thr=method_kwargs["noisy_channel_threshold"],
outside_channel_thr=method_kwargs["outside_channel_threshold"],
n_neighbors=method_kwargs["n_neighbors"],
nyquist_threshold=method_kwargs["nyquist_threshold"],
welch_window_ms=method_kwargs["welch_window_ms"],
outside_channels_location=method_kwargs["outside_channels_location"],
)
return chunk_labels[order_r] if order_r is not None else chunk_labels


def detect_bad_channels(
recording: BaseRecording,
method: str = "coherence+psd",
Expand All @@ -173,6 +214,7 @@ def detect_bad_channels(
neighborhood_r2_radius_um: float = 30.0,
seed: int | None = None,
channel_filters: set | None = None,
job_kwargs: dict | None = None,
):
"""
Perform bad channel detection.
Expand Down Expand Up @@ -225,14 +267,12 @@ def detect_bad_channels(
if method in ("std", "mad"):
random_chunk_kwargs["return_in_uV"] = False
random_chunk_kwargs["concatenated"] = True
elif method == "coherence+psd":
random_chunk_kwargs["return_in_uV"] = True
random_chunk_kwargs["concatenated"] = False
elif method == "neighborhood_r2":
random_chunk_kwargs["return_in_uV"] = False
random_chunk_kwargs["concatenated"] = False

random_data = get_random_data_chunks(recording_hp, **random_chunk_kwargs)
if method != "coherence+psd":
random_data = get_random_data_chunks(recording_hp, **random_chunk_kwargs)

channel_labels = np.zeros(recording.get_num_channels(), dtype="U5")
channel_labels[:] = "good"
Expand All @@ -248,6 +288,9 @@ def detect_bad_channels(
channel_labels[mask] = "noise"

elif method == "coherence+psd":
job_kwargs = {} if job_kwargs is None else job_kwargs
job_kwargs = fix_job_kwargs(job_kwargs)

# some checks
assert recording.has_scaleable_traces(), (
"The 'coherence+psd' method uses thresholds assuming the traces are in uV, "
Expand All @@ -267,24 +310,30 @@ def detect_bad_channels(
order_f = None
order_r = None

# Create empty channel labels and fill with bad-channel detection estimate for each chunk
chunk_channel_labels = np.zeros((recording.get_num_channels(), len(random_data)), dtype=np.int8)

for i, random_chunk in enumerate(random_data):
random_chunk_sorted = random_chunk[:, order_f] if order_f is not None else random_chunk
chunk_labels = detect_bad_channels_ibl(
raw=random_chunk_sorted,
fs=recording.sampling_frequency,
psd_hf_threshold=psd_hf_threshold,
dead_channel_thr=dead_channel_threshold,
noisy_channel_thr=noisy_channel_threshold,
outside_channel_thr=outside_channel_threshold,
n_neighbors=n_neighbors,
nyquist_threshold=nyquist_threshold,
welch_window_ms=welch_window_ms,
outside_channels_location=outside_channels_location,
)
chunk_channel_labels[:, i] = chunk_labels[order_r] if order_r is not None else chunk_labels
method_kwargs = dict(
order_f=order_f,
order_r=order_r,
psd_hf_threshold=psd_hf_threshold,
dead_channel_threshold=dead_channel_threshold,
noisy_channel_threshold=noisy_channel_threshold,
outside_channel_threshold=outside_channel_threshold,
n_neighbors=n_neighbors,
nyquist_threshold=nyquist_threshold,
welch_window_ms=welch_window_ms,
outside_channels_location=outside_channels_location,
)
random_slices = get_random_sample_slices(recording_hp, **random_chunk_kwargs)
executor = TimeSeriesChunkExecutor(
recording_hp,
_detect_bad_channels_chunk,
_detect_bad_channels_chunk_init,
(recording_hp, method_kwargs),
handle_returns=True,
chunk_size=random_chunk_kwargs["chunk_size"],
job_name="detect_bad_channels",
**job_kwargs,
)
chunk_channel_labels = np.stack(executor.run(slices=random_slices), axis=1)

# Take the mode of the chunk estimates as final result. Convert to binary good / bad channel output.
mode_channel_labels, _ = mode(chunk_channel_labels, axis=1, keepdims=False)
Expand Down
65 changes: 65 additions & 0 deletions src/spikeinterface/preprocessing/tests/test_detect_bad_channels.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,71 @@ def test_detect_bad_channels_std_mad():
), "wrong channels locations."


@pytest.mark.parametrize("pool_engine", ["thread", "process"])
def test_detect_bad_channels_parallel(pool_engine):
recording = generate_recording(num_channels=16, durations=[1, 1], seed=0)
recording.set_channel_gains(1)
recording.set_channel_offsets(0)
method_kwargs = dict(
method="coherence+psd",
num_random_chunks=4,
chunk_duration_s=0.05,
seed=0,
)

expected_bad_channel_ids, expected_channel_labels = detect_bad_channels(recording, **method_kwargs)
job_kwargs = dict(n_jobs=2, pool_engine=pool_engine, max_threads_per_worker=1)
if pool_engine == "process":
job_kwargs["mp_context"] = "spawn"

bad_channel_ids, channel_labels = detect_bad_channels(
recording,
**method_kwargs,
job_kwargs=job_kwargs,
)

np.testing.assert_array_equal(bad_channel_ids, expected_bad_channel_ids)
np.testing.assert_array_equal(channel_labels, expected_channel_labels)


def test_detect_bad_channels_parallel_unfiltered_spawn():
"""
generate_recording() marks its output as already filtered, so the parallel test above never
exercises the highpass_filter() wrapper that detect_bad_channels builds internally for an
unfiltered recording. Use a plain NumpyRecording (is_filtered() defaults to False) with a
spawned process pool, so that wrapper has to survive cross-process serialization.
"""
num_channels = 16
sampling_frequency = 30000.0
rng = np.random.default_rng(0)
traces_list = [rng.standard_normal((int(sampling_frequency), num_channels)).astype("float32") for _ in range(2)]
recording = NumpyRecording(traces_list, sampling_frequency)
recording.set_channel_gains(1)
recording.set_channel_offsets(0)
probe = generate_linear_probe(num_elec=num_channels)
probe.set_device_channel_indices(np.arange(num_channels))
recording.set_probe(probe)
assert not recording.is_filtered()

method_kwargs = dict(
method="coherence+psd",
num_random_chunks=4,
chunk_duration_s=0.05,
seed=0,
)

expected_bad_channel_ids, expected_channel_labels = detect_bad_channels(recording, **method_kwargs)

bad_channel_ids, channel_labels = detect_bad_channels(
recording,
**method_kwargs,
job_kwargs=dict(n_jobs=2, pool_engine="process", mp_context="spawn", max_threads_per_worker=1),
)

np.testing.assert_array_equal(bad_channel_ids, expected_bad_channel_ids)
np.testing.assert_array_equal(channel_labels, expected_channel_labels)


@pytest.mark.parametrize("outside_channels_location", ["bottom", "top", "both"])
def test_detect_bad_channels_extremes(outside_channels_location):
num_channels = 64
Expand Down
Loading