File size: 2,040 Bytes
dc9d5f6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
"""Exact four-second preprocessing used by the released SenuaLab models."""

from __future__ import annotations

import math

import numpy as np
from scipy.signal import butter, resample, resample_poly, sosfiltfilt


NUM_CHANNELS = 19
INTERNATIONAL_10_20_CHANNELS = [
    "Fp1", "Fp2", "F3", "F4", "C3", "C4", "P3", "P4", "O1", "O2",
    "F7", "F8", "T3", "T4", "T5", "T6", "Fz", "Cz", "Pz",
]
DATASET_29_CHANNEL_ORDER = INTERNATIONAL_10_20_CHANNELS


def normalize_sample(sample):
    x = np.asarray(sample, dtype=np.float32)
    mean = x.mean()
    std = x.std()
    return np.clip((x - mean) / max(float(std), 1e-7), -8.0, 8.0).astype(np.float32)


def preprocess_batch(
    batch,
    *,
    source_sfreq=None,
    target_sfreq=250,
    target_samples=1000,
    channel_names=None,
):
    """Filter, average-reference, resample, and globally normalize `[B,C,T]`."""

    del channel_names  # IED Finder has already arranged the declared channel order.
    x = np.asarray(batch, dtype=np.float32)
    if x.ndim == 2:
        x = x[None]
    if x.ndim != 3 or x.shape[1] != NUM_CHANNELS:
        raise ValueError(f"Expected [batch,19,time], received {x.shape}")
    sfreq = float(source_sfreq or target_sfreq)
    if sfreq <= 90.0:
        raise ValueError(f"Sampling rate must be above the 45 Hz cutoff; received {sfreq}")
    sos = butter(4, [1.0, 45.0], btype="bandpass", fs=sfreq, output="sos")
    x = sosfiltfilt(sos, x, axis=-1).astype(np.float32, copy=False)
    x -= x.mean(axis=1, keepdims=True)
    if not math.isclose(sfreq, float(target_sfreq)):
        up = int(target_sfreq)
        down = int(round(sfreq))
        common = math.gcd(up, down)
        x = resample_poly(x, up // common, down // common, axis=-1).astype(np.float32)
    if x.shape[-1] != int(target_samples):
        x = resample(x, int(target_samples), axis=-1).astype(np.float32)
    mean = x.mean(axis=(1, 2), keepdims=True)
    std = x.std(axis=(1, 2), keepdims=True)
    return np.clip((x - mean) / np.maximum(std, 1e-7), -8.0, 8.0).astype(np.float32)