Seizure Detection#

Seizure probability from intracranial EEG with a CNN + bidirectional LSTM applied to spectrograms (Sladky et al. 2022, Brain Communications, doi:10.1093/braincomms/fcac115).

Bundled models (load_trained_model()):

  • 'modelA': the model from the published work.

  • 'modelB': the same architecture trained on an extended data set.

The model was designed for 300 s inputs; its outputs near the ends of an input have little LSTM context and are less reliable, which is why the high-level function drops discard_edges_s at each window edge.

Two ways to use it:

  1. High level (recommended): predict_channel_seizure_probability() turns one long channel into a continuous probability trace on a 0.5 s grid. It handles the windowing, gaps, flat segments and the recording edges.

  2. Low level: preprocess_input() + infer_seizure_probability() for batches of windows you cut yourself.

Conventions (both levels)#

  • fs must be a whole, even number of Hz, at least 200 Hz. A float such as 500.0 (typical of .mat / MEF headers) is accepted. Odd or fractional rates raise ValueError: resample first (anti-aliasing is the caller’s job). These rules come from the spectrogram grid (whole samples per 1 s segment, an integer 0.5 s hop, 100 bins below Nyquist), not from the training data. At 200-256 Hz the anti-alias filter of the acquisition or resampling usually attenuates the upper bins (from ~0.4 fs = 80-100 Hz), which the model did not see in training: prefer higher rates (see preprocess_input()).

  • Spectrogram: 1 s segments (nperseg = fs), 0.5 s hop, bins 0-99 Hz at 1 Hz. Column j of a window describes the 1 s of signal centred (j + 1) * 0.5 s after the window’s first sample.

  • In the output of predict_channel_seizure_probability(), NaN means “not evaluated”, never “no seizure”: t = 0 and every 1 s segment containing NaN/inf or a flat signal are NaN; every other segment has a value. Combine channels or time with NaN-aware functions (np.nanmax) and never replace NaN by 0.

  • The low-level preprocess_input() zero-fills NaN samples and does not mark them; masking is done by the high-level function.

Example (high level, one channel)#

import numpy as np
from brainmaze_torch.seizure_detection import predict_channel_seizure_probability

fs = 500
x = np.random.randn(fs * 600)                 # 10 min, one channel
x[100 * fs:130 * fs] = np.nan                 # a 30 s gap
t, p = predict_channel_seizure_probability(x, fs, model='modelA')
# t = 0, 0.5, 1.0, ... s; p is NaN at t = 0 and over the gap

Example (low level, batch of 300 s windows)#

import numpy as np
from brainmaze_torch.seizure_detection import (
    load_trained_model, preprocess_input, infer_seizure_probability)

model = load_trained_model('modelA')
fs = 500
x = np.random.randn(3, fs * 300)              # 3 windows (rows) of 300 s
t, f, sxx = preprocess_input(x, fs, return_axes=True)   # sxx: (3, 100, 599)
y = infer_seizure_probability(sxx, model)     # (3, 599); y[:, j] belongs to t[j]
brainmaze_torch.seizure_detection.infer_seizure_probability(x, model, use_cuda=False, cuda_number=0)#

Run the seizure model on a batch of preprocessed spectrograms.

Parameters:
  • x (array_like, shape (batch_size, 100, n_times)) – Output of preprocess_input(). Must be finite and have n_times >= 2. Recommended window length is 300 s (599 columns).

  • model (torch.nn.Module) – A loaded seizure model, e.g. from load_trained_model('modelA'). The model is temporarily put in eval mode (and moved to the requested device); its original device and train/eval mode are restored on return, so the caller’s object is not mutated.

  • use_cuda (bool, optional) – Run inference on cuda:<cuda_number>. Default False (CPU).

  • cuda_number (int, optional) – CUDA device index used when use_cuda is True. Default 0.

Returns:

Seizure probability (softmax probability of class index 3 of the model’s 4 output classes) for every spectrogram column. Column j corresponds to the time stamp t[j] returned by preprocess_input() ((j + 1) * 0.5 s from window start). Values are in [0, 1].

Return type:

numpy.ndarray, shape (batch_size, n_times), float32

Raises:

ValueError – If x has the wrong shape or contains non-finite values.

Notes

Inference runs under torch.inference_mode(); no autograd graph is kept.

brainmaze_torch.seizure_detection.load_trained_model(model_name)#

Load one of the bundled, pre-trained seizure-detection models.

Parameters:

model_name ({'modelA', 'modelB'}) – 'modelA': the model from the published work (Sladky et al. 2022). 'modelB': the same architecture trained on an extended data set.

Returns:

A torch.nn.Module on the CPU, in eval mode, with the weights loaded strictly (every parameter must match).

Return type:

SeizureDetectModel

Raises:

KeyError – If model_name is not one of the bundled models.

brainmaze_torch.seizure_detection.predict_channel_seizure_probability(x, fs, model='modelA', use_cuda=False, cuda_number=0, n_batch=128, window_s=300, step_s=20, discard_edges_s=10, min_valid_fraction=0.0, fill_recording_edges=True)#

Continuous seizure-probability trace for one (long) iEEG channel.

The signal is cut into overlapping windows of window_s seconds every step_s seconds (the “regular” windows of the published method). Each window is converted to a spectrogram and run through the model; discard_edges_s seconds are dropped at both window edges (the BiLSTM has little context there), and overlapping windows are combined with a NaN-ignoring maximum. If the recording end does not fall on the step grid, one extra window aligned to the recording end covers the tail; it only fills time that no regular window covers and never changes a value produced by the regular windows. With the default parameters the output is therefore identical to the published method (brainmaze-torch <= 0.1.1) wherever that method produced an estimate, except that invalid (gap / flat) segments are NaN instead of a number.

Parameters:
  • x (array_like, shape (n_samples,)) – Single-channel raw signal (1D). Lists and integer arrays are accepted. NaN / +-inf mark missing data (gaps).

  • fs (int or float) – Sampling rate in Hz. Must be a whole, even number >= 200 Hz; a float such as 500.0 is accepted. Resample other rates first.

  • model (str or torch.nn.Module, optional) – 'modelA' (published model), 'modelB' (extended training set) or an already-loaded model. Default 'modelA'. A passed model is not mutated (its device and train/eval mode are restored).

  • use_cuda (bool, optional) – Run inference on cuda:<cuda_number>. Default False.

  • cuda_number (int, optional) – CUDA device index. Default 0.

  • n_batch (int, optional) – Windows per inference batch (memory/speed trade-off only; does not change the result). Default 128.

  • window_s (float, optional) – Window length in seconds; multiple of 0.5, >= 2. Default 300 (the length the model was designed for). The recording must be at least this long.

  • step_s (float, optional) – Step between window starts in seconds; multiple of 0.5. Must satisfy step_s <= window_s - 2 * discard_edges_s - 0.5 so that the kept (non-discarded) parts of consecutive windows leave no hole (a window has 2 * window_s - 1 half-second columns). Default 20.

  • discard_edges_s (float, optional) – Seconds dropped at each window edge; multiple of 0.5, may be 0. Default 10.

  • min_valid_fraction (float, optional) – Window preference near gaps. A window is preferred if at least this fraction of its 1 s segments (spectrogram columns) are valid (finite samples, not flat). A valid segment gets the maximum over the preferred windows covering it; only if none of them covers it, the maximum over all windows covering it is used. This never turns a valid segment into NaN; it only lets gap-heavy (mostly zero-filled) windows be ignored where a better window exists. Default 0.0: every window is preferred, i.e. the published method. Values > 0 change probabilities next to long gaps compared with the published method (not validated). The choice is made per time point, so the result is not monotone in this value: a point switches from “preferred windows” to the fallback (all windows) as soon as the last preferred window covering it drops below the threshold, and different points switch at different values. If the value exceeds the valid fraction of every window, no window is preferred and the result equals that of 0.0. Choose one value for a study and do not interpolate between values.

  • fill_recording_edges (bool, optional) – The first and last discard_edges_s seconds of the recording lie in no window’s kept interior. If True (default), they are taken from the edge of the first/last window (shorter model context, hence less reliable). If False, they are NaN.

Returns:

  • t (numpy.ndarray, shape (n_out,)) – Time in seconds from the first sample: t[k] = k * 0.5 (so the value for time t is at index 2 * t), n_out = n_samples // (fs // 2). The last stamp is the centre of the last complete 1 s segment.

  • prob (numpy.ndarray, shape (n_out,), float64) – Seizure probability in [0, 1] for the 1 s segment centred on t[k], or NaN where no estimate exists. NaN is returned for:

    • t = 0 (a 1 s segment centred at 0 s would start before the data),

    • every 1 s segment containing at least one NaN/inf sample (gaps),

    • every 1 s segment over which the signal is constant (flat line, e.g. disconnected or saturated channel),

    • the first/last discard_edges_s s if fill_recording_edges is False.

    NaN therefore always means “not evaluated”, never “no seizure”; a value is never 0 merely because a time point was not covered. Every other (valid) segment always gets a value, whatever min_valid_fraction is.

Raises:

ValueError – If x is not 1D, fs is invalid, the window parameters are inconsistent, or the recording is shorter than window_s.

Notes

  • Gaps are not interpolated. Within an evaluated window the gap samples are zero-filled for the model (as in the original implementation), which can influence the probability of valid segments next to a gap via the BiLSTM context; only the gap segments themselves are masked to NaN. To fill short gaps instead (e.g. by interpolation), do so before calling this function; filled samples are then treated as valid data. See min_valid_fraction to ignore mostly-gap windows where possible.

  • Probabilities of overlapping windows are combined with a NaN-ignoring maximum (np.fmax), as in the published method.

Examples

>>> import numpy as np
>>> from brainmaze_torch.seizure_detection import predict_channel_seizure_probability
>>> fs = 500
>>> x = np.random.randn(fs * 600)          # 10 min, single channel
>>> t, p = predict_channel_seizure_probability(x, fs, model='modelA')
>>> t.shape == p.shape, float(t[1] - t[0])
(True, 0.5)
brainmaze_torch.seizure_detection.preprocess_input(x, fs, return_axes=False)#

Convert raw signal windows into normalised spectrograms for the model.

Each row of x is processed independently:

  1. z-score the raw signal (NaN-aware mean/std over the row; a row with zero or undefined std is only mean-centred, never divided by 0),

  2. replace NaN / +-inf samples by 0,

  3. spectrogram with nperseg = fs (1 s, 1 Hz bins) and a 0.5 s hop (noverlap = fs // 2), keeping the first 100 bins (0-99 Hz),

  4. z-score every frequency bin over time, using only columns with non-zero total power. Columns with zero power (e.g. a flat or zero-filled gap) are left at 0. If a bin has zero variance, or fewer than two columns have power, that bin is only mean-centred (no division by 0).

Parameters:
  • x (array_like, shape (n_samples,) or (batch_size, n_samples)) – Raw signal. A 1D input is treated as a single window. Rows are windows (or channels), columns are samples. Must contain at least fs samples (1 s).

  • fs (int or float) – Sampling rate in Hz. Must be a whole, even number >= 200; a float such as 500.0 is accepted. Odd or fractional rates raise ValueError – resample first (see Notes).

  • return_axes (bool, optional) – If True, also return the time and frequency axes. Default False.

Returns:

  • t (numpy.ndarray, shape (n_times,)) – Only if return_axes. Time stamp in seconds of every column, relative to the first sample of the window: (j + 1) * 0.5 – the centre of the 1 s segment [j * 0.5, j * 0.5 + 1).

  • f (numpy.ndarray, shape (100,)) – Only if return_axes. Frequencies in Hz (0, 1, …, 99).

  • sxx (numpy.ndarray, shape (batch_size, 100, n_times)) – Normalised spectrograms, n_times = floor((n_samples - fs) / (fs // 2)) + 1 (= 2 * duration_s - 1 for whole-second inputs). Always finite.

Raises:

ValueError – If fs is invalid, x is not 1D/2D, or shorter than 1 s.

Notes

NaN samples are zero-filled before the spectrogram, so the model still sees a value there. This low-level function does not mark such columns; use predict_channel_seizure_probability(), which reports any 1 s segment containing NaN/inf (or a flat signal) as NaN in its output.

Where the fs rules come from: they follow from the spectrogram grid, not from the training data (whose sampling rate is not documented in this repository). nperseg = fs needs a whole number of samples per second; the 0.5 s hop (fs / 2 samples) needs an even rate, otherwise the output would drift off the 0.5 s grid; and 100 one-Hz bins (0-99 Hz) need a Nyquist frequency of at least 100 Hz, i.e. fs >= 200.

Anti-aliasing caveat at 200-256 Hz: acquisition (or resampling) anti-alias filters usually start attenuating at about 0.4 * fs, i.e. 80-100 Hz at these rates. The upper spectrogram bins then hold less power than in recordings with a higher native rate, which the model never saw in that form: a silent domain shift. Prefer a native rate >= 250-500 Hz, and when you resample, keep the new rate high enough that the anti-alias filter’s transition band lies above 100 Hz.