diff --git a/oephys2nix/tonix.py b/oephys2nix/tonix.py index 60c41da..bc13b8b 100644 --- a/oephys2nix/tonix.py +++ b/oephys2nix/tonix.py @@ -7,6 +7,7 @@ import numpy as np import rlxnix as rlx from IPython import embed from neo.io import OpenEphysBinaryIO +from scipy.signal import butter, sosfiltfilt from oephys2nix.logging import setup_logging from oephys2nix.metadata import create_dict_from_section, create_metadata_from_dict @@ -68,15 +69,6 @@ class RawToNix: def append_fish_lines(self) -> None: """Append fish lines from open-ephys.""" efishs = ["ttl-line", "global-eod", "stimulus", "local-eod", "sinus"] - - efish_types = [ - "open-ephys.data.sampled", - "open-ephys.data.sampled", - "open-ephys.data.sampled", - "open-ephys.data.sampled", - "open-ephys.data.sampled", - ] - efish_group = self.block.create_group("efish", "open-ephys.sampled") efish_neo_data = self._load_neo_object(["Data_ADC", "acquisition_board_ADC"]) @@ -86,7 +78,7 @@ class RawToNix: efish_neo_data_array = efish_neo_data[:, i] data_array = self.block.create_data_array( f"{efishs[i]}", - f"{efish_types[i]}", + "open-ephys.data.sampled", data=efish_neo_data_array.magnitude.flatten(), label="voltage", unit="V", @@ -178,6 +170,72 @@ class RawToNix: nix_data_array.append_sampled_dimension(1, label="channel") gr.data_arrays.append(nix_data_array) + preprocessed_data = self.preprocess_raw(nix_data_array) + + preprocessed_data.append_sampled_dimension( + 1 / raw_neo_data.sampling_rate.magnitude, label="time", unit="s" + ) + preprocessed_data.append_sampled_dimension(1, label="channel") + gr.data_arrays.append(nix_data_array) + + def preprocess_raw( + self, + data: nixio.DataArray, + batch_seconds: float = 20.0, + overlap_seconds: float = 1, + band_hz: tuple[float] = (300.0, 6000.0), + filter_order: int = 3, + ) -> None: + + fs = int(1 / data.dimensions[0].sampling_interval) + n_samples = data.shape[0] + + batch_samples = int(round(batch_seconds * fs)) + overlap_samples = int(round(overlap_seconds * fs)) + + low_hz, high_zh = band_hz + if not 0 < low_hz < high_zh < fs / 2: + raise ValueError("Band must lie striclty between 0 and Nyquist.") + + sos = butter( + filter_order, + band_hz, + btype="bandpass", + fs=fs, + output="sos", + ) + nix_data_array = self.block.create_data_array( + name="preprocessed-data", + array_type="open-ephys.data.sampled", + dtype=nixio.DataType.Float, + unit="uV", + shape=data.shape, + ) + + for start in range(0, n_samples, batch_samples): + stop = min(start + batch_samples, n_samples) + + # Read extra samples on both sides of the output batch. + read_start = max(0, start - overlap_samples) + read_stop = min(n_samples, stop + overlap_samples) + + block = data[read_start:read_stop].astype(np.float32) + + ref = np.median(block, axis=1, keepdims=True) + x = block - ref + + # Zero-phase filtering along time, independently per channel. + filtered = sosfiltfilt(sos, x, axis=0) + + # Discard the overlap; retain only this batch's central samples. + keep_start = start - read_start + keep_stop = keep_start + (stop - start) + + nix_data_array[start:stop] = filtered[keep_start:keep_stop] + + log.debug(f"Processed {stop:,} / {n_samples:,} samples") + return nix_data_array + def close(self) -> None: """Close all nix files.""" self.nix_file.close()