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:
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.Low level:
preprocess_input()+infer_seizure_probability()for batches of windows you cut yourself.
Conventions (both levels)#
fsmust be a whole, even number of Hz, at least 200 Hz. A float such as500.0(typical of .mat / MEF headers) is accepted. Odd or fractional rates raiseValueError: 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 (seepreprocess_input()).Spectrogram: 1 s segments (
nperseg = fs), 0.5 s hop, bins 0-99 Hz at 1 Hz. Columnjof a window describes the 1 s of signal centred(j + 1) * 0.5s after the window’s first sample.In the output of
predict_channel_seizure_probability(), NaN means “not evaluated”, never “no seizure”:t = 0and 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 haven_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_cudais True. Default 0.
- Returns:
Seizure probability (softmax probability of class index 3 of the model’s 4 output classes) for every spectrogram column. Column
jcorresponds to the time stampt[j]returned bypreprocess_input()((j + 1) * 0.5s from window start). Values are in [0, 1].- Return type:
numpy.ndarray, shape (batch_size, n_times), float32
- Raises:
ValueError – If
xhas 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.Moduleon the CPU, in eval mode, with the weights loaded strictly (every parameter must match).- Return type:
SeizureDetectModel
- Raises:
KeyError – If
model_nameis 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_sseconds everystep_sseconds (the “regular” windows of the published method). Each window is converted to a spectrogram and run through the model;discard_edges_sseconds 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.0is 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.5so that the kept (non-discarded) parts of consecutive windows leave no hole (a window has2 * window_s - 1half-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_sseconds 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 timetis at index2 * 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_ss iffill_recording_edgesis 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_fractionis.
- Raises:
ValueError – If
xis not 1D,fsis invalid, the window parameters are inconsistent, or the recording is shorter thanwindow_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_fractionto 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
xis processed independently: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),
replace NaN / +-inf samples by 0,
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),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
fssamples (1 s).fs (int or float) – Sampling rate in Hz. Must be a whole, even number >= 200; a float such as
500.0is accepted. Odd or fractional rates raiseValueError– 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 - 1for whole-second inputs). Always finite.
- Raises:
ValueError – If
fsis invalid,xis 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
fsrules come from: they follow from the spectrogram grid, not from the training data (whose sampling rate is not documented in this repository).nperseg = fsneeds a whole number of samples per second; the 0.5 s hop (fs / 2samples) 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.