braindecode 1.8.0.dev1125__tar.gz → 1.8.0.dev168891348__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {braindecode-1.8.0.dev1125/braindecode.egg-info → braindecode-1.8.0.dev168891348}/PKG-INFO +1 -1
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/base.py +15 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub.py +90 -25
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub_io.py +67 -5
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/atcnet.py +2 -2
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/attn_sleep.py +87 -28
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/base.py +43 -4
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/ctnet.py +9 -6
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegsimpleconv.py +8 -6
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/ifnet.py +8 -5
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sleep_stager_blanco_2020.py +5 -9
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sleep_stager_chambon_2018.py +1 -7
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sparcnet.py +2 -2
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/summary.csv +1 -1
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/tidnet.py +1 -7
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/windowers.py +126 -30
- braindecode-1.8.0.dev168891348/braindecode/version.py +1 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348/braindecode.egg-info}/PKG-INFO +1 -1
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/whats_new.rst +59 -1
- braindecode-1.8.0.dev1125/braindecode/version.py +0 -1
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/LICENSE.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/MANIFEST.in +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/NOTICE.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/README.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/augmentation/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/augmentation/base.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/augmentation/functional.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/augmentation/transforms.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/classifier.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bbci.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bcicomp.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/datasets.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/format.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub_format.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub_validation.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/iterable.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/chb_mit.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/collate.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/mne.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/moabb.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/nmt.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/registry.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/siena.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/sleep_physio_challe_18.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/sleep_physionet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/tuh.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/utils.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/xy.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datautil/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datautil/channel_utils.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datautil/hub_formats.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datautil/serialization.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datautil/util.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/eegneuralnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/functional/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/functional/functions.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/functional/initialization.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/attentionbasenet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/bendr.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/biot.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/brainmodule.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/cbramod.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/codebrain.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/config.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/contrawr.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/dance.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/deep4.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/deepsleepnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/dgcnn.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegconformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegdino.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eeginception_erp.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eeginception_mi.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegitnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegminer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegnex.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegpt.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegsym.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/eegtcnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/emg2qwerty.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/fbcnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/fblightconvnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/fbmsnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/hybrid.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/interpolated.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/labram.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/luna.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/medformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/meta_neuromotor.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/msvtnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/mvpformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/patchedtransformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/reve.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sccnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/shallow_fbcsp.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/signal_jepa.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sinc_shallow.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/sstdpn.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/steegformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/syncnet.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/tcformer.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/tcn.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/tsinception.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/usleep.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/util.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/zuna.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/activation.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/attention.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/blocks.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/convolution.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/dance_modules.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/filter.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/interpolation.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/layers.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/linear.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/parametrization.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/stats.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/util.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/modules/wrapper.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/eegprep_preprocess.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/mne_preprocess.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/preprocess.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/preprocessing/util.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/regressor.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/samplers/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/samplers/base.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/samplers/ssl.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/training/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/training/callbacks.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/training/losses.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/training/scoring.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/util.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/__init__.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/attribution.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/confusion_matrices.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/frequency.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/metrics.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/sanity.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/visualization/topology.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode.egg-info/SOURCES.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode.egg-info/dependency_links.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode.egg-info/requires.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode.egg-info/top_level.txt +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/Makefile +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/_templates/autosummary/class.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/_templates/autosummary/class_in_subdir.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/_templates/autosummary/function.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/_templates/autosummary/function_in_subdir.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/api.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/cite.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/conf.py +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/help.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/index.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/install/install.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/install/install_pip.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/install/install_source.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/attention.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/channel.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/convolution.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/filterbank.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/gnn.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/interpretable.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/lbm.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/recurrent.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/categorization/spd.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/models.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/models_categorization.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/models_table.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/models/models_visualization.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/docs/sg_execution_times.rst +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/pyproject.toml +0 -0
- {braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/setup.cfg +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: braindecode
|
|
3
|
-
Version: 1.8.0.
|
|
3
|
+
Version: 1.8.0.dev168891348
|
|
4
4
|
Summary: Deep learning software to decode EEG, ECG or MEG signals
|
|
5
5
|
Author-email: Robin Tibor Schirrmeister <robintibor@gmail.com>, Bruno Aristimunha Pinto <b.aristimunha@gmail.com>, Alexandre Gramfort <agramfort@meta.com>
|
|
6
6
|
Maintainer-email: Alexandre Gramfort <agramfort@meta.com>, Bruno Aristimunha Pinto <b.aristimunha@gmail.com>, Robin Tibor Schirrmeister <robintibor@gmail.com>
|
|
@@ -59,6 +59,7 @@ def _html_row(label, value):
|
|
|
59
59
|
|
|
60
60
|
_METADATA_INTERNAL_COLS = {
|
|
61
61
|
"i_window_in_trial",
|
|
62
|
+
"i_trial_in_dataset",
|
|
62
63
|
"i_start_in_trial",
|
|
63
64
|
"i_stop_in_trial",
|
|
64
65
|
"target",
|
|
@@ -1374,6 +1375,20 @@ class BaseConcatDataset(ConcatDataset, HubDatasetMixin, Generic[T]):
|
|
|
1374
1375
|
"datasets are WindowsDataset."
|
|
1375
1376
|
)
|
|
1376
1377
|
|
|
1378
|
+
for ds in self.datasets:
|
|
1379
|
+
if hasattr(ds, "_windows") and ds._windows is not None:
|
|
1380
|
+
df = ds._windows.metadata
|
|
1381
|
+
else:
|
|
1382
|
+
df = ds.metadata
|
|
1383
|
+
if (
|
|
1384
|
+
"i_trial_in_dataset" in df.columns
|
|
1385
|
+
and "i_trial_in_dataset" in ds.description
|
|
1386
|
+
):
|
|
1387
|
+
raise ValueError(
|
|
1388
|
+
"Dataset descriptions cannot contain the reserved window "
|
|
1389
|
+
"metadata key 'i_trial_in_dataset'."
|
|
1390
|
+
)
|
|
1391
|
+
|
|
1377
1392
|
all_dfs = list()
|
|
1378
1393
|
for ds in self.datasets:
|
|
1379
1394
|
if hasattr(ds, "_windows") and ds._windows is not None:
|
{braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub.py
RENAMED
|
@@ -28,6 +28,7 @@ The format follows a BIDS-inspired sourcedata structure:
|
|
|
28
28
|
# License: BSD (3-clause)
|
|
29
29
|
|
|
30
30
|
import contextlib
|
|
31
|
+
import copy
|
|
31
32
|
import json
|
|
32
33
|
import logging
|
|
33
34
|
import tempfile
|
|
@@ -56,6 +57,7 @@ from .hub_io import (
|
|
|
56
57
|
_load_eegwindows_from_zarr,
|
|
57
58
|
_load_raw_from_zarr,
|
|
58
59
|
_load_windows_from_zarr,
|
|
60
|
+
_prepare_info_for_json,
|
|
59
61
|
_save_eegwindows_to_zarr,
|
|
60
62
|
_save_raw_to_zarr,
|
|
61
63
|
_save_windows_to_zarr,
|
|
@@ -72,6 +74,33 @@ log = logging.getLogger(__name__)
|
|
|
72
74
|
_LOCK_FILE = "format_info.json"
|
|
73
75
|
|
|
74
76
|
|
|
77
|
+
def _normalize_kwargs_for_json(value, field_name):
|
|
78
|
+
"""Return one preprocessing-kwargs value as native strict JSON."""
|
|
79
|
+
|
|
80
|
+
def _convert(obj):
|
|
81
|
+
if isinstance(obj, np.ndarray):
|
|
82
|
+
return _convert(obj.tolist())
|
|
83
|
+
if isinstance(obj, dict):
|
|
84
|
+
return {key: _convert(item) for key, item in obj.items()}
|
|
85
|
+
if isinstance(obj, (list, tuple)):
|
|
86
|
+
return [_convert(item) for item in obj]
|
|
87
|
+
if isinstance(obj, np.generic):
|
|
88
|
+
return _convert(obj.item())
|
|
89
|
+
return obj
|
|
90
|
+
|
|
91
|
+
try:
|
|
92
|
+
converted = _convert(value)
|
|
93
|
+
if isinstance(converted, str):
|
|
94
|
+
raise ValueError(
|
|
95
|
+
"a root string is ambiguous with the legacy encoded format"
|
|
96
|
+
)
|
|
97
|
+
return json.loads(json.dumps(converted, allow_nan=False))
|
|
98
|
+
except (OverflowError, RecursionError, TypeError, ValueError) as error:
|
|
99
|
+
raise ValueError(
|
|
100
|
+
f"{field_name} must contain only finite JSON-serializable values"
|
|
101
|
+
) from error
|
|
102
|
+
|
|
103
|
+
|
|
75
104
|
class HubDatasetMixin:
|
|
76
105
|
"""
|
|
77
106
|
Mixin class for Hugging Face Hub integration with EEG datasets.
|
|
@@ -685,14 +714,57 @@ class HubDatasetMixin:
|
|
|
685
714
|
f"{output_path} already exists. Set overwrite=True to replace it."
|
|
686
715
|
)
|
|
687
716
|
|
|
688
|
-
# Create zarr store (zarr v3 API)
|
|
689
|
-
root = zarr.open(str(output_path), mode="w")
|
|
690
|
-
|
|
691
717
|
# Validate uniformity across all datasets using shared validation
|
|
692
718
|
dataset_type, _, _ = hub_validation.validate_dataset_uniformity(self.datasets)
|
|
693
719
|
|
|
694
|
-
#
|
|
695
|
-
|
|
720
|
+
# Normalize every JSON value before opening the output store. The cached
|
|
721
|
+
# values below are the exact values handed to the write helpers.
|
|
722
|
+
prepared_infos = []
|
|
723
|
+
for ds in self.datasets:
|
|
724
|
+
if dataset_type == "WindowsDataset":
|
|
725
|
+
info = ds.windows.info
|
|
726
|
+
elif dataset_type in ("EEGWindowsDataset", "RawDataset"):
|
|
727
|
+
info = ds.raw.info
|
|
728
|
+
prepared_infos.append(_prepare_info_for_json(info.to_json_dict()))
|
|
729
|
+
|
|
730
|
+
prepared_kwargs = {}
|
|
731
|
+
for kwarg_name in [
|
|
732
|
+
"raw_preproc_kwargs",
|
|
733
|
+
"window_kwargs",
|
|
734
|
+
"window_preproc_kwargs",
|
|
735
|
+
]:
|
|
736
|
+
expected_present = hasattr(self.datasets[0], kwarg_name)
|
|
737
|
+
expected_value = None
|
|
738
|
+
expected_token = None
|
|
739
|
+
for i_ds, ds in enumerate(self.datasets):
|
|
740
|
+
present = hasattr(ds, kwarg_name)
|
|
741
|
+
if present != expected_present:
|
|
742
|
+
raise ValueError(
|
|
743
|
+
f"{kwarg_name} on dataset {i_ds} has inconsistent presence; "
|
|
744
|
+
"the Zarr format stores one global value"
|
|
745
|
+
)
|
|
746
|
+
if not present:
|
|
747
|
+
continue
|
|
748
|
+
value = _normalize_kwargs_for_json(
|
|
749
|
+
getattr(ds, kwarg_name), f"{kwarg_name} on dataset {i_ds}"
|
|
750
|
+
)
|
|
751
|
+
value_token = json.dumps(
|
|
752
|
+
value, sort_keys=True, separators=(",", ":"), allow_nan=False
|
|
753
|
+
)
|
|
754
|
+
if i_ds == 0:
|
|
755
|
+
expected_value = value
|
|
756
|
+
expected_token = value_token
|
|
757
|
+
elif value_token != expected_token:
|
|
758
|
+
raise ValueError(
|
|
759
|
+
f"{kwarg_name} on dataset {i_ds} differs from dataset 0; "
|
|
760
|
+
"the Zarr format stores one global value"
|
|
761
|
+
)
|
|
762
|
+
if expected_present:
|
|
763
|
+
prepared_kwargs[kwarg_name] = expected_value
|
|
764
|
+
|
|
765
|
+
# Create compressor and zarr store (zarr v3 API) only after preflight.
|
|
766
|
+
compressor = _create_compressor(compression, compression_level)
|
|
767
|
+
root = zarr.open(str(output_path), mode="w")
|
|
696
768
|
|
|
697
769
|
# Store global metadata
|
|
698
770
|
root.attrs["n_datasets"] = len(self.datasets)
|
|
@@ -706,21 +778,10 @@ class HubDatasetMixin:
|
|
|
706
778
|
root.attrs["zarr_version"] = zarr.__version__
|
|
707
779
|
root.attrs["scipy_version"] = scipy.__version__
|
|
708
780
|
|
|
709
|
-
# Save preprocessing kwargs
|
|
710
|
-
#
|
|
711
|
-
for kwarg_name in
|
|
712
|
-
|
|
713
|
-
"window_kwargs",
|
|
714
|
-
"window_preproc_kwargs",
|
|
715
|
-
]:
|
|
716
|
-
# Check first dataset for these attributes
|
|
717
|
-
if hasattr(first_ds, kwarg_name):
|
|
718
|
-
kwargs = getattr(first_ds, kwarg_name)
|
|
719
|
-
if kwargs:
|
|
720
|
-
root.attrs[kwarg_name] = json.dumps(kwargs)
|
|
721
|
-
|
|
722
|
-
# Create compressor
|
|
723
|
-
compressor = _create_compressor(compression, compression_level)
|
|
781
|
+
# Save preprocessing kwargs from the preflight cache. These are
|
|
782
|
+
# typically set by windowing functions on individual datasets.
|
|
783
|
+
for kwarg_name, kwargs in prepared_kwargs.items():
|
|
784
|
+
root.attrs[kwarg_name] = kwargs
|
|
724
785
|
|
|
725
786
|
# Save each recording
|
|
726
787
|
for i_ds, ds in enumerate(self.datasets):
|
|
@@ -731,7 +792,7 @@ class HubDatasetMixin:
|
|
|
731
792
|
data = ds.windows.get_data()
|
|
732
793
|
metadata = ds.windows.metadata
|
|
733
794
|
description = ds.description
|
|
734
|
-
info_dict =
|
|
795
|
+
info_dict = prepared_infos[i_ds]
|
|
735
796
|
target_name = ds.target_name if hasattr(ds, "target_name") else None
|
|
736
797
|
|
|
737
798
|
# Save using inlined function
|
|
@@ -750,7 +811,7 @@ class HubDatasetMixin:
|
|
|
750
811
|
raw = ds.raw
|
|
751
812
|
metadata = ds.metadata
|
|
752
813
|
description = ds.description
|
|
753
|
-
info_dict =
|
|
814
|
+
info_dict = prepared_infos[i_ds]
|
|
754
815
|
targets_from = ds.targets_from
|
|
755
816
|
last_target_only = ds.last_target_only
|
|
756
817
|
|
|
@@ -771,7 +832,7 @@ class HubDatasetMixin:
|
|
|
771
832
|
# Get continuous raw data from RawDataset
|
|
772
833
|
raw = ds.raw
|
|
773
834
|
description = ds.description
|
|
774
|
-
info_dict =
|
|
835
|
+
info_dict = prepared_infos[i_ds]
|
|
775
836
|
target_name = ds.target_name if hasattr(ds, "target_name") else None
|
|
776
837
|
|
|
777
838
|
# Save using inlined function
|
|
@@ -956,10 +1017,14 @@ class HubDatasetMixin:
|
|
|
956
1017
|
"window_preproc_kwargs",
|
|
957
1018
|
]:
|
|
958
1019
|
if kwarg_name in root.attrs:
|
|
959
|
-
kwargs =
|
|
1020
|
+
kwargs = root.attrs[kwarg_name]
|
|
1021
|
+
if isinstance(kwargs, str):
|
|
1022
|
+
# Stores written by older braindecode versions kept
|
|
1023
|
+
# these attributes as double-encoded JSON strings.
|
|
1024
|
+
kwargs = json.loads(kwargs)
|
|
960
1025
|
# Set on each individual dataset (where they were originally stored)
|
|
961
1026
|
for ds in datasets:
|
|
962
|
-
setattr(ds, kwarg_name, kwargs)
|
|
1027
|
+
setattr(ds, kwarg_name, copy.deepcopy(kwargs))
|
|
963
1028
|
|
|
964
1029
|
return concat_ds
|
|
965
1030
|
|
{braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/datasets/bids/hub_io.py
RENAMED
|
@@ -8,6 +8,7 @@ These functions keep the Zarr serialization details isolated from hub.py.
|
|
|
8
8
|
from __future__ import annotations
|
|
9
9
|
|
|
10
10
|
import json
|
|
11
|
+
from numbers import Real
|
|
11
12
|
from pathlib import Path
|
|
12
13
|
|
|
13
14
|
import numpy as np
|
|
@@ -17,17 +18,78 @@ from mne.utils import _soft_import
|
|
|
17
18
|
zarr = _soft_import("zarr", purpose="hugging face integration", strict=False)
|
|
18
19
|
|
|
19
20
|
|
|
21
|
+
def _is_non_bool_real(value):
|
|
22
|
+
return not isinstance(value, (bool, np.bool_)) and isinstance(
|
|
23
|
+
value, (Real, np.integer, np.floating)
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _prepare_info_for_json(obj, path="info", *, _in_numeric_sequence=False):
|
|
28
|
+
"""Normalize an MNE Info value to strict JSON without losing sequence NaNs."""
|
|
29
|
+
if isinstance(obj, np.ndarray):
|
|
30
|
+
obj = obj.tolist()
|
|
31
|
+
|
|
32
|
+
if isinstance(obj, dict):
|
|
33
|
+
return {
|
|
34
|
+
key: _prepare_info_for_json(value, f"{path}.{key}")
|
|
35
|
+
for key, value in obj.items()
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
if isinstance(obj, (list, tuple)):
|
|
39
|
+
values = list(obj)
|
|
40
|
+
has_none = any(value is None for value in values)
|
|
41
|
+
has_number = any(_is_non_bool_real(value) for value in values)
|
|
42
|
+
all_none = bool(values) and all(value is None for value in values)
|
|
43
|
+
if all_none or (has_none and has_number):
|
|
44
|
+
raise ValueError(
|
|
45
|
+
f"{path} is ambiguous: numeric sequences cannot contain JSON null"
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
is_numeric_sequence = bool(values) and all(
|
|
49
|
+
_is_non_bool_real(value) for value in values
|
|
50
|
+
)
|
|
51
|
+
return [
|
|
52
|
+
_prepare_info_for_json(
|
|
53
|
+
value,
|
|
54
|
+
f"{path}[{index}]",
|
|
55
|
+
_in_numeric_sequence=is_numeric_sequence,
|
|
56
|
+
)
|
|
57
|
+
for index, value in enumerate(values)
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
if isinstance(obj, (bool, np.bool_)):
|
|
61
|
+
return bool(obj)
|
|
62
|
+
if isinstance(obj, (int, np.integer)):
|
|
63
|
+
return int(obj)
|
|
64
|
+
if _is_non_bool_real(obj):
|
|
65
|
+
value = float(obj)
|
|
66
|
+
if np.isnan(value):
|
|
67
|
+
if _in_numeric_sequence:
|
|
68
|
+
return None
|
|
69
|
+
raise ValueError(f"{path} contains unsupported NaN")
|
|
70
|
+
if np.isposinf(value):
|
|
71
|
+
raise ValueError(f"{path} contains positive infinity")
|
|
72
|
+
if np.isneginf(value):
|
|
73
|
+
raise ValueError(f"{path} contains negative infinity")
|
|
74
|
+
return value
|
|
75
|
+
if obj is None or isinstance(obj, str):
|
|
76
|
+
return obj
|
|
77
|
+
raise ValueError(f"{path} contains a non-JSON-serializable value")
|
|
78
|
+
|
|
79
|
+
|
|
20
80
|
def _restore_nan_from_json(obj):
|
|
21
|
-
"""Restore NaN values from None in
|
|
81
|
+
"""Restore NaN values from None in JSON-loaded attributes.
|
|
22
82
|
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
|
|
83
|
+
JSON null is reserved for NaN only in non-boolean numeric sequences.
|
|
84
|
+
All-null lists are therefore the representation of validated all-NaN
|
|
85
|
+
sequences, while ordinary null values elsewhere remain unchanged.
|
|
26
86
|
"""
|
|
27
87
|
if isinstance(obj, dict):
|
|
28
88
|
return {k: _restore_nan_from_json(v) for k, v in obj.items()}
|
|
29
89
|
if isinstance(obj, list):
|
|
30
|
-
if len(obj) > 0 and all(
|
|
90
|
+
if len(obj) > 0 and all(
|
|
91
|
+
value is None or _is_non_bool_real(value) for value in obj
|
|
92
|
+
):
|
|
31
93
|
return [np.nan if x is None else x for x in obj]
|
|
32
94
|
return [_restore_nan_from_json(v) for v in obj]
|
|
33
95
|
return obj
|
|
@@ -196,7 +196,7 @@ class ATCNet(EEGModuleMixin, nn.Module):
|
|
|
196
196
|
num_heads : int
|
|
197
197
|
Number of attention heads, denoted H in table 1 of the paper [1]_.
|
|
198
198
|
Defaults to 2 as in [1]_.
|
|
199
|
-
|
|
199
|
+
att_drop_prob : float
|
|
200
200
|
Dropout probability used in the attention block, denoted pa in table 1
|
|
201
201
|
of the paper [1]_. Defaults to 0.5 as in [1]_.
|
|
202
202
|
tcn_depth : int
|
|
@@ -206,7 +206,7 @@ class ATCNet(EEGModuleMixin, nn.Module):
|
|
|
206
206
|
tcn_kernel_size : int
|
|
207
207
|
Temporal kernel size used in TCN block, denoted Kt in table 1 of the
|
|
208
208
|
paper [1]_. Defaults to 4 as in [1]_.
|
|
209
|
-
|
|
209
|
+
tcn_drop_prob : float
|
|
210
210
|
Dropout probability used in the TCN block, denoted pt in table 1
|
|
211
211
|
of the paper [1]_. Defaults to 0.3 as in [1]_.
|
|
212
212
|
tcn_activation : torch.nn.Module
|
{braindecode-1.8.0.dev1125 → braindecode-1.8.0.dev168891348}/braindecode/models/attn_sleep.py
RENAMED
|
@@ -1,10 +1,36 @@
|
|
|
1
1
|
# Authors: Divyesh Narayanan <divyesh.narayanan@gmail.com>
|
|
2
|
+
# Sarthak Tayal <sarthaktayal2@gmail.com>
|
|
2
3
|
#
|
|
3
4
|
# License: BSD (3-clause)
|
|
5
|
+
#
|
|
6
|
+
# This implementation derives from https://github.com/emadeldeen24/AttnSleep:
|
|
7
|
+
#
|
|
8
|
+
# MIT License
|
|
9
|
+
#
|
|
10
|
+
# Copyright (c) 2020 Emadeldeen Eldele
|
|
11
|
+
#
|
|
12
|
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
13
|
+
# of this software and associated documentation files (the "Software"), to deal
|
|
14
|
+
# in the Software without restriction, including without limitation the rights
|
|
15
|
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
16
|
+
# copies of the Software, and to permit persons to whom the Software is
|
|
17
|
+
# furnished to do so, subject to the following conditions:
|
|
18
|
+
#
|
|
19
|
+
# The above copyright notice and this permission notice shall be included in all
|
|
20
|
+
# copies or substantial portions of the Software.
|
|
21
|
+
#
|
|
22
|
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
23
|
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
24
|
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
25
|
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
26
|
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
27
|
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
28
|
+
# SOFTWARE.
|
|
4
29
|
|
|
5
30
|
import math
|
|
6
31
|
import warnings
|
|
7
32
|
from copy import deepcopy
|
|
33
|
+
from numbers import Integral
|
|
8
34
|
|
|
9
35
|
import torch
|
|
10
36
|
import torch.nn.functional as F
|
|
@@ -33,7 +59,8 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
33
59
|
|
|
34
60
|
Warning - This model was designed for signals of 30 seconds at 100Hz or 125Hz (in which case
|
|
35
61
|
the reference architecture from [1]_ which was validated on SHHS dataset [2]_ will be used)
|
|
36
|
-
to use any other input is likely to make the model perform in unintended ways.
|
|
62
|
+
to use any other input is likely to make the model perform in unintended ways. Any other
|
|
63
|
+
window length also needs a ``d_model`` of its own, see the parameter below.
|
|
37
64
|
|
|
38
65
|
Parameters
|
|
39
66
|
----------
|
|
@@ -44,7 +71,10 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
44
71
|
Also the input dimension of the first FC layer in the feed forward
|
|
45
72
|
and the output of the second FC layer in the same.
|
|
46
73
|
Increase for higher sampling rate/signal length.
|
|
47
|
-
It should be divisible by n_attn_heads
|
|
74
|
+
It should be divisible by n_attn_heads. It must also equal the number
|
|
75
|
+
of time steps returned by the feature extractor: 80 for 30 seconds at
|
|
76
|
+
100 Hz and 100 for 30 seconds at 125 Hz. A construction error reports
|
|
77
|
+
the value needed for other window lengths.
|
|
48
78
|
d_ff : int
|
|
49
79
|
Output dimension of the first FC layer in the feed forward and the
|
|
50
80
|
input dimension of the second FC layer in the same.
|
|
@@ -58,15 +88,12 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
58
88
|
If True, return the features, i.e. the output of the feature extractor
|
|
59
89
|
(before the final linear layer). If False, pass the features through
|
|
60
90
|
the final linear layer.
|
|
61
|
-
n_classes : int
|
|
62
|
-
Alias for `n_outputs`.
|
|
63
|
-
input_size_s : float
|
|
64
|
-
Alias for `input_window_seconds`.
|
|
65
91
|
activation : nn.Module, default=nn.ReLU
|
|
66
|
-
Activation function class to apply
|
|
92
|
+
Activation function class to apply in the AFR block and the TCE
|
|
93
|
+
feed-forward block. Should be a PyTorch activation
|
|
67
94
|
module class like ``nn.ReLU`` or ``nn.ELU``. Default is ``nn.ReLU``.
|
|
68
|
-
activation_mrcnn : nn.Module, default=nn.
|
|
69
|
-
Activation function class to apply in the
|
|
95
|
+
activation_mrcnn : nn.Module, default=nn.GELU
|
|
96
|
+
Activation function class to apply in the multi-resolution CNN layer.
|
|
70
97
|
Should be a PyTorch activation module class like ``nn.ReLU`` or
|
|
71
98
|
``nn.GELU``. Default is ``nn.GELU``.
|
|
72
99
|
|
|
@@ -99,6 +126,14 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
99
126
|
n_chans=None,
|
|
100
127
|
n_times=None,
|
|
101
128
|
):
|
|
129
|
+
if (
|
|
130
|
+
sum(value is not None for value in (n_times, sfreq, input_window_seconds))
|
|
131
|
+
< 2
|
|
132
|
+
):
|
|
133
|
+
raise ValueError(
|
|
134
|
+
"AttnSleep requires at least two of n_times, sfreq, and "
|
|
135
|
+
"input_window_seconds."
|
|
136
|
+
)
|
|
102
137
|
super().__init__(
|
|
103
138
|
n_outputs=n_outputs,
|
|
104
139
|
n_chans=n_chans,
|
|
@@ -140,6 +175,14 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
140
175
|
activation=activation_mrcnn,
|
|
141
176
|
activation_se=activation,
|
|
142
177
|
)
|
|
178
|
+
feature_length = self._feature_length(mrcnn, self.n_times)
|
|
179
|
+
if feature_length != d_model:
|
|
180
|
+
raise ValueError(
|
|
181
|
+
f"d_model is {d_model} but the feature extractor returns "
|
|
182
|
+
f"{feature_length} time steps for an input of {self.n_times} "
|
|
183
|
+
f"samples at {self.sfreq} Hz. Set d_model={feature_length}, with "
|
|
184
|
+
"an n_attn_heads that divides it."
|
|
185
|
+
)
|
|
143
186
|
attn = _MultiHeadedAttention(n_attn_heads, d_model, after_reduced_cnn_size)
|
|
144
187
|
ff = _PositionwiseFeedForward(d_model, d_ff, drop_prob, activation=activation)
|
|
145
188
|
tce = _TCE(
|
|
@@ -150,23 +193,26 @@ class AttnSleep(EEGModuleMixin, nn.Module):
|
|
|
150
193
|
)
|
|
151
194
|
|
|
152
195
|
self.feature_extractor = nn.Sequential(mrcnn, tce)
|
|
153
|
-
self.len_last_layer =
|
|
196
|
+
self.len_last_layer = feature_length * after_reduced_cnn_size
|
|
154
197
|
self.return_feats = return_feats
|
|
155
198
|
|
|
156
199
|
# TODO: Add new way to handle return features
|
|
157
200
|
"""if return_feats:
|
|
158
201
|
raise ValueError("return_feat == True is not accepted anymore")"""
|
|
159
202
|
if not return_feats:
|
|
160
|
-
self.final_layer = nn.Linear(
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
165
|
-
|
|
166
|
-
|
|
167
|
-
|
|
168
|
-
|
|
169
|
-
|
|
203
|
+
self.final_layer = nn.Linear(self.len_last_layer, self.n_outputs)
|
|
204
|
+
|
|
205
|
+
@staticmethod
|
|
206
|
+
def _feature_length(mrcnn, n_times):
|
|
207
|
+
training_states = [(module, module.training) for module in mrcnn.modules()]
|
|
208
|
+
mrcnn.eval()
|
|
209
|
+
try:
|
|
210
|
+
with torch.no_grad():
|
|
211
|
+
out = mrcnn(torch.zeros(1, 1, n_times))
|
|
212
|
+
finally:
|
|
213
|
+
for module, was_training in training_states:
|
|
214
|
+
module.training = was_training
|
|
215
|
+
return out.shape[-1]
|
|
170
216
|
|
|
171
217
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
172
218
|
"""
|
|
@@ -195,7 +241,7 @@ class _SELayer(nn.Module):
|
|
|
195
241
|
self.avg_pool = nn.AdaptiveAvgPool1d(1)
|
|
196
242
|
self.fc = nn.Sequential(
|
|
197
243
|
nn.Linear(channel, channel // reduction, bias=False),
|
|
198
|
-
activation(
|
|
244
|
+
activation(),
|
|
199
245
|
nn.Linear(channel // reduction, channel, bias=False),
|
|
200
246
|
nn.Sigmoid(),
|
|
201
247
|
)
|
|
@@ -236,10 +282,10 @@ class _SEBasicBlock(nn.Module):
|
|
|
236
282
|
super(_SEBasicBlock, self).__init__()
|
|
237
283
|
self.conv1 = nn.Conv1d(inplanes, planes, stride)
|
|
238
284
|
self.bn1 = nn.BatchNorm1d(planes)
|
|
239
|
-
self.relu = activation(
|
|
285
|
+
self.relu = activation()
|
|
240
286
|
self.conv2 = nn.Conv1d(planes, planes, 1)
|
|
241
287
|
self.bn2 = nn.BatchNorm1d(planes)
|
|
242
|
-
self.se = _SELayer(planes, reduction)
|
|
288
|
+
self.se = _SELayer(planes, reduction, activation=activation)
|
|
243
289
|
self.downsample = downsample
|
|
244
290
|
self.stride = stride
|
|
245
291
|
self.features = nn.Sequential(
|
|
@@ -320,11 +366,11 @@ class _MRCNN(nn.Module):
|
|
|
320
366
|
self.dropout = nn.Dropout(drate)
|
|
321
367
|
self.inplanes = 128
|
|
322
368
|
self.AFR = self._make_layer(
|
|
323
|
-
_SEBasicBlock, after_reduced_cnn_size, 1,
|
|
369
|
+
_SEBasicBlock, after_reduced_cnn_size, 1, activation=activation_se
|
|
324
370
|
)
|
|
325
371
|
|
|
326
372
|
def _make_layer(
|
|
327
|
-
self, block, planes, blocks, stride=1,
|
|
373
|
+
self, block, planes, blocks, stride=1, activation: type[nn.Module] = nn.ReLU
|
|
328
374
|
): # makes residual SE block
|
|
329
375
|
downsample = None
|
|
330
376
|
if stride != 1 or self.inplanes != planes * block.expansion:
|
|
@@ -340,10 +386,12 @@ class _MRCNN(nn.Module):
|
|
|
340
386
|
)
|
|
341
387
|
|
|
342
388
|
layers = []
|
|
343
|
-
layers.append(
|
|
389
|
+
layers.append(
|
|
390
|
+
block(self.inplanes, planes, stride, downsample, activation=activation)
|
|
391
|
+
)
|
|
344
392
|
self.inplanes = planes * block.expansion
|
|
345
393
|
for i in range(1, blocks):
|
|
346
|
-
layers.append(block(self.inplanes, planes,
|
|
394
|
+
layers.append(block(self.inplanes, planes, activation=activation))
|
|
347
395
|
|
|
348
396
|
return nn.Sequential(*layers)
|
|
349
397
|
|
|
@@ -375,9 +423,20 @@ class _MultiHeadedAttention(nn.Module):
|
|
|
375
423
|
def __init__(self, h, d_model, after_reduced_cnn_size, dropout=0.1):
|
|
376
424
|
"""Take in model size and number of heads."""
|
|
377
425
|
super().__init__()
|
|
378
|
-
|
|
426
|
+
if (
|
|
427
|
+
isinstance(h, bool)
|
|
428
|
+
or not isinstance(h, Integral)
|
|
429
|
+
or h <= 0
|
|
430
|
+
or d_model % h != 0
|
|
431
|
+
):
|
|
432
|
+
raise ValueError(
|
|
433
|
+
"n_attn_heads must be a positive integer that divides d_model, "
|
|
434
|
+
f"got n_attn_heads={h!r} and d_model={d_model}."
|
|
435
|
+
)
|
|
436
|
+
h = int(h)
|
|
379
437
|
self.d_per_head = d_model // h
|
|
380
438
|
self.h = h
|
|
439
|
+
self.attn = torch.empty(0)
|
|
381
440
|
|
|
382
441
|
base_conv = CausalConv1d(
|
|
383
442
|
in_channels=after_reduced_cnn_size,
|
|
@@ -80,7 +80,7 @@ class _BraindecodeDocstringMeta(NumpyDocstringInheritanceInitMeta):
|
|
|
80
80
|
unwrapped function and correctly inherits ``cls.__doc__``.
|
|
81
81
|
"""
|
|
82
82
|
|
|
83
|
-
def __init__(cls, class_name, class_bases, class_dict):
|
|
83
|
+
def __init__(cls, class_name, class_bases, class_dict, **kwargs):
|
|
84
84
|
super().__init__(class_name, class_bases, class_dict)
|
|
85
85
|
# Only wrap subclass __init__s, not EEGModuleMixin itself.
|
|
86
86
|
# Wrapping the mixin would cause super().__init__() calls to
|
|
@@ -198,7 +198,36 @@ class EEGModuleMixin(_BaseHubMixin, metaclass=_BraindecodeDocstringMeta):
|
|
|
198
198
|
See :ref:`load-pretrained-models` for a complete tutorial.
|
|
199
199
|
"""
|
|
200
200
|
|
|
201
|
+
#: Attributes that :func:`torch.jit.script` must not introspect. The
|
|
202
|
+
#: signal-related properties raise :class:`ValueError` when their value was
|
|
203
|
+
#: neither given nor inferable, and ``mapping`` carries a postponed
|
|
204
|
+
#: annotation TorchScript cannot resolve. Either one aborts scripting
|
|
205
|
+
#: before ``forward`` is ever compiled.
|
|
206
|
+
__jit_ignored_attributes__ = [
|
|
207
|
+
*sorted(_EEG_PARAMS),
|
|
208
|
+
"_chs_info",
|
|
209
|
+
"input_shape",
|
|
210
|
+
"mapping",
|
|
211
|
+
]
|
|
212
|
+
#: Rich MNE channel dictionaries are intentionally unavailable in scripted
|
|
213
|
+
#: forwards; they remain unchanged on eager models and in saved configs.
|
|
214
|
+
__jit_unused_properties__ = ["chs_info"]
|
|
215
|
+
|
|
201
216
|
def __init_subclass__(cls, **kwargs):
|
|
217
|
+
license = kwargs.pop("license", "bsd-3-clause")
|
|
218
|
+
|
|
219
|
+
# TorchScript only honours ``__jit_ignored_attributes__`` for
|
|
220
|
+
# properties defined directly on the concrete class: it collects them
|
|
221
|
+
# with ``vars(type(module))``, which skips inherited ones. Rebinding
|
|
222
|
+
# the very same property objects on each subclass makes them visible
|
|
223
|
+
# there without changing any runtime behaviour.
|
|
224
|
+
for name in cls.__jit_ignored_attributes__:
|
|
225
|
+
if name in cls.__dict__:
|
|
226
|
+
continue
|
|
227
|
+
prop = getattr(cls, name, None)
|
|
228
|
+
if isinstance(prop, property):
|
|
229
|
+
setattr(cls, name, prop)
|
|
230
|
+
|
|
202
231
|
# Append model-specific Hub integration notes to the docstring.
|
|
203
232
|
# This runs before the metaclass __init__, so the Hub notes will
|
|
204
233
|
# be included in the docstring that the metaclass processes.
|
|
@@ -226,8 +255,6 @@ class EEGModuleMixin(_BaseHubMixin, metaclass=_BraindecodeDocstringMeta):
|
|
|
226
255
|
)
|
|
227
256
|
repo_url = kwargs.pop("repo_url", "https://braindecode.org")
|
|
228
257
|
library_name = kwargs.pop("library_name", "braindecode")
|
|
229
|
-
license = kwargs.pop("license", "bsd-3-clause")
|
|
230
|
-
|
|
231
258
|
# Register a coder so that type[nn.Module] parameters
|
|
232
259
|
# (e.g. activation=nn.ELU) are serialized as importable
|
|
233
260
|
# strings in config.json and decoded back on load.
|
|
@@ -286,6 +313,12 @@ class EEGModuleMixin(_BaseHubMixin, metaclass=_BraindecodeDocstringMeta):
|
|
|
286
313
|
self._chs_info = chs_info # type: ignore[assignment]
|
|
287
314
|
self._n_outputs = n_outputs # type: ignore[assignment]
|
|
288
315
|
self._n_chans = n_chans # type: ignore[assignment]
|
|
316
|
+
# TorchScript cannot represent the rich MNE dictionaries in _chs_info.
|
|
317
|
+
# Keep the original eager state above and expose only the derived scalar
|
|
318
|
+
# to scripted signal-property getters.
|
|
319
|
+
self._n_chans_for_jit = (
|
|
320
|
+
len(chs_info) if n_chans is None and chs_info is not None else n_chans
|
|
321
|
+
)
|
|
289
322
|
self._n_times = n_times # type: ignore[assignment]
|
|
290
323
|
self._sfreq = sfreq # type: ignore[assignment]
|
|
291
324
|
|
|
@@ -310,6 +343,12 @@ class EEGModuleMixin(_BaseHubMixin, metaclass=_BraindecodeDocstringMeta):
|
|
|
310
343
|
|
|
311
344
|
@property
|
|
312
345
|
def n_chans(self) -> int:
|
|
346
|
+
if torch.jit.is_scripting():
|
|
347
|
+
if self._n_chans_for_jit is None:
|
|
348
|
+
raise ValueError(
|
|
349
|
+
"n_chans could not be inferred. Either specify n_chans or chs_info."
|
|
350
|
+
)
|
|
351
|
+
return self._n_chans_for_jit
|
|
313
352
|
if self._n_chans is None and self._chs_info is not None:
|
|
314
353
|
return len(self._chs_info)
|
|
315
354
|
elif self._n_chans is None:
|
|
@@ -319,7 +358,7 @@ class EEGModuleMixin(_BaseHubMixin, metaclass=_BraindecodeDocstringMeta):
|
|
|
319
358
|
return self._n_chans
|
|
320
359
|
|
|
321
360
|
@property
|
|
322
|
-
def chs_info(self) -> list[
|
|
361
|
+
def chs_info(self) -> list[dict]:
|
|
323
362
|
if self._chs_info is None:
|
|
324
363
|
raise ValueError("chs_info not specified.")
|
|
325
364
|
return self._chs_info
|
|
@@ -59,16 +59,19 @@ class CTNet(EEGModuleMixin, nn.Module):
|
|
|
59
59
|
|
|
60
60
|
Parameters
|
|
61
61
|
----------
|
|
62
|
-
|
|
63
|
-
Activation function to use in the
|
|
62
|
+
activation_patch : nn.Module, default=nn.ELU
|
|
63
|
+
Activation function to use in the convolutional patch embedding.
|
|
64
|
+
activation_transformer : nn.Module, default=nn.GELU
|
|
65
|
+
Activation function to use in the Transformer encoder.
|
|
64
66
|
num_heads : int, default=4
|
|
65
67
|
Number of attention heads in the Transformer encoder.
|
|
66
|
-
embed_dim : int or None, default=
|
|
68
|
+
embed_dim : int or None, default=40
|
|
67
69
|
Embedding size (dimensionality) for the Transformer encoder.
|
|
68
70
|
num_layers : int, default=6
|
|
69
71
|
Number of encoder layers in the Transformer.
|
|
70
|
-
n_filters_time : int, default=
|
|
71
|
-
Number of temporal filters in the first convolutional layer.
|
|
72
|
+
n_filters_time : int or None, default=None
|
|
73
|
+
Number of temporal filters in the first convolutional layer. Inferred
|
|
74
|
+
from ``embed_dim`` and ``depth_multiplier`` when left at ``None``.
|
|
72
75
|
kernel_size : int, default=64
|
|
73
76
|
Kernel size for the temporal convolutional layer.
|
|
74
77
|
depth_multiplier : int, default=2
|
|
@@ -77,7 +80,7 @@ class CTNet(EEGModuleMixin, nn.Module):
|
|
|
77
80
|
Pooling size for the first average pooling layer.
|
|
78
81
|
pool_size_2 : int, default=8
|
|
79
82
|
Pooling size for the second average pooling layer.
|
|
80
|
-
|
|
83
|
+
cnn_drop_prob : float, default=0.3
|
|
81
84
|
Dropout probability after convolutional layers.
|
|
82
85
|
att_positional_drop_prob : float, default=0.1
|
|
83
86
|
Dropout probability for the positional encoding in the Transformer.
|