"""EEG Annotation Tool adapter for the experimental SenuaLab EEGPT linear probe.""" from __future__ import annotations from typing import Any import mne try: from braindecode.models import EEGPT except ImportError as exc: raise ImportError("This model requires braindecode[hub]==1.6.1") from exc from .preprocessing import INTERNATIONAL_10_20_CHANNELS, preprocess_batch as _preprocess EEGPT_CHANNELS = [ "Fp1", "Fp2", "F3", "F4", "C3", "C4", "P3", "P4", "O1", "O2", "F7", "F8", "T7", "T8", "P7", "P8", "Fz", "Cz", "Pz", ] MODEL_CHANNELS = INTERNATIONAL_10_20_CHANNELS MODEL_CHANNEL_ALIASES = {"T3": "T7", "T4": "T8", "T5": "P7", "T6": "P8"} MODEL_INPUT_SAMPLES = 1000 MODEL_SAMPLING_RATE_HZ = 250.0 MODEL_WINDOW_SECONDS = 4.0 MODEL_ENTRY_CLASS = "SenuaEEGPTLinearProbe" MODEL_NUM_CLASSES = 2 MODEL_CLASS_LABELS = ["Non-IED", "IED"] MODEL_NON_IED_CLASS_INDEX = 0 MODEL_IED_CLASS_INDICES = [1] MODEL_DESCRIPTION = "Experimental SenuaLab EEGPT frozen-encoder linear probe" MODEL_BATCH_PREPROCESSOR = "preprocess_batch" MODEL_REQUIRED_REFERENCE = "common average (applied by model adapter)" MODEL_REQUIRED_FILTERS = ["1-45 Hz zero-phase Butterworth (applied by model adapter)"] MODEL_REQUIRED_NORMALIZATION = "global four-second window z-score, clipped to [-8,8]" MODEL_INPUT_UNIT = "scale-invariant after window z-score" MODEL_SOURCE_SIGNAL_POLICY = "raw" MODEL_REQUIRES_FULL_WINDOW = True MODEL_REQUIRES_ALL_CHANNELS = True MODEL_DEFAULT_THRESHOLD = 0.40234375 MODEL_DEFAULT_STEP_MS = 500.0 MODEL_DEFAULT_PAD_POLICY = "skip" MODEL_DEFAULT_BATCH_SIZE = 16 MODEL_DEFAULT_BATCH_MEMORY_MB = 256.0 MODEL_VALIDATION_NOTE = "Experimental ablation; not recommended for deployment and not externally validated." def preprocess_batch(batch, source_sfreq=None, channel_names=None): return _preprocess( batch, source_sfreq=source_sfreq, target_sfreq=250, target_samples=1000, channel_names=channel_names, ) def _chs_info() -> list[dict[str, Any]]: info = mne.create_info(EEGPT_CHANNELS, sfreq=250.0, ch_types="eeg") info.set_montage("standard_1020") return info["chs"] class SenuaEEGPTLinearProbe(EEGPT): def __init__(self, num_channels: int = 19, num_classes: int = 2, input_length: int = 1000): if num_channels != 19 or input_length != 1000: raise ValueError("SenuaEEGPTLinearProbe requires 19 channels and 1,000 samples") super().__init__( n_outputs=num_classes, n_chans=19, chs_info=_chs_info(), n_times=1000, sfreq=250.0, chan_proj_type="none", return_encoder_output=False, )