braindecode 1.8.0.dev169208086__tar.gz → 1.8.0.dev180935432__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.dev169208086/braindecode.egg-info → braindecode-1.8.0.dev180935432}/PKG-INFO +1 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/base.py +25 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/functional.py +2 -4
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/transforms.py +5 -3
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/classifier.py +2 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/base.py +15 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub.py +90 -25
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_io.py +67 -5
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/sleep_physionet.py +9 -3
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/eegneuralnet.py +4 -2
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/__init__.py +2 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/atcnet.py +2 -2
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/attn_sleep.py +87 -28
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/base.py +40 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/ctnet.py +9 -6
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegsimpleconv.py +8 -6
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/ifnet.py +8 -5
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sleep_stager_blanco_2020.py +5 -9
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sleep_stager_chambon_2018.py +1 -7
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sparcnet.py +2 -2
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/summary.csv +1 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/tidnet.py +1 -7
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/util.py +9 -0
- braindecode-1.8.0.dev180935432/braindecode/models/zuna.py +612 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/eegprep_preprocess.py +1 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/windowers.py +126 -30
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/regressor.py +2 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/samplers/base.py +40 -2
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/training/losses.py +7 -4
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/util.py +16 -10
- braindecode-1.8.0.dev180935432/braindecode/version.py +1 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432/braindecode.egg-info}/PKG-INFO +1 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/SOURCES.txt +1 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/api.rst +1 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/whats_new.rst +134 -3
- braindecode-1.8.0.dev169208086/braindecode/version.py +0 -1
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/LICENSE.txt +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/MANIFEST.in +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/NOTICE.txt +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/README.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bbci.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bcicomp.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/datasets.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/format.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_format.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_validation.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/iterable.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/chb_mit.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/collate.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/mne.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/moabb.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/nmt.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/registry.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/siena.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/sleep_physio_challe_18.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/tuh.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/utils.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/xy.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datautil/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datautil/channel_utils.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datautil/hub_formats.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datautil/serialization.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datautil/util.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/functional/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/functional/functions.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/functional/initialization.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/attentionbasenet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/bendr.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/biot.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/brainmodule.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/cbramod.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/codebrain.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/config.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/contrawr.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/dance.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/deep4.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/deepsleepnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/dgcnn.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegconformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegdino.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eeginception_erp.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eeginception_mi.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegitnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegminer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegnex.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegpt.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegsym.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/eegtcnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/emg2qwerty.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/fbcnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/fblightconvnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/fbmsnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/hybrid.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/interpolated.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/labram.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/luna.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/medformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/meta_neuromotor.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/msvtnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/mvpformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/patchedtransformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/reve.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sccnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/shallow_fbcsp.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/signal_jepa.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sinc_shallow.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/sstdpn.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/steegformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/syncnet.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/tcformer.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/tcn.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/tsinception.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/usleep.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/activation.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/attention.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/blocks.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/convolution.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/dance_modules.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/filter.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/interpolation.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/layers.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/linear.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/parametrization.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/stats.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/util.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/modules/wrapper.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/mne_preprocess.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/preprocess.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/util.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/samplers/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/samplers/ssl.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/training/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/training/callbacks.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/training/scoring.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/__init__.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/attribution.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/confusion_matrices.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/frequency.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/metrics.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/sanity.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/visualization/topology.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/dependency_links.txt +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/requires.txt +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/top_level.txt +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/Makefile +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/class.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/class_in_subdir.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/function.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/function_in_subdir.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/cite.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/conf.py +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/help.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/index.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/install/install.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/install/install_pip.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/install/install_source.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/attention.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/channel.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/convolution.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/filterbank.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/gnn.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/interpretable.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/lbm.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/recurrent.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/categorization/spd.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/models.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/models_categorization.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/models_table.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/models/models_visualization.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/docs/sg_execution_times.rst +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/pyproject.toml +0 -0
- {braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/setup.cfg +0 -0
{braindecode-1.8.0.dev169208086/braindecode.egg-info → braindecode-1.8.0.dev180935432}/PKG-INFO
RENAMED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: braindecode
|
|
3
|
-
Version: 1.8.0.
|
|
3
|
+
Version: 1.8.0.dev180935432
|
|
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>
|
{braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/base.py
RENAMED
|
@@ -3,6 +3,7 @@
|
|
|
3
3
|
# Bruno Aristimunha <b.aristimunha@gmail.com>
|
|
4
4
|
# Martin Wimpff <martin.wimpff@iss.uni-stuttgart.de>
|
|
5
5
|
# Valentin Iovene <val@too.gy>
|
|
6
|
+
# Sarthak Tayal <sarthaktayal2@gmail.com>
|
|
6
7
|
# License: BSD (3-clause)
|
|
7
8
|
|
|
8
9
|
from numbers import Real
|
|
@@ -176,6 +177,14 @@ class Compose(Transform):
|
|
|
176
177
|
return X, y
|
|
177
178
|
|
|
178
179
|
|
|
180
|
+
def _as_mixed_target(y, lam_dtype):
|
|
181
|
+
# an untouched target is the same as being mixed with itself with lam of one
|
|
182
|
+
if isinstance(y, (tuple, list)) and len(y) == 3:
|
|
183
|
+
return tuple(y)
|
|
184
|
+
lam = torch.ones(y.shape[0], device=y.device, dtype=lam_dtype)
|
|
185
|
+
return y, y, lam
|
|
186
|
+
|
|
187
|
+
|
|
179
188
|
class _AugmentationCollate:
|
|
180
189
|
"""Collate that applies a transform to each batch, with optional expansion.
|
|
181
190
|
|
|
@@ -193,7 +202,10 @@ class _AugmentationCollate:
|
|
|
193
202
|
``0`` (default) applies the transform in place (batch size unchanged).
|
|
194
203
|
``> 0`` keeps the clean originals and appends ``n_augmentation``
|
|
195
204
|
independently transformed copies, returning ``(X, y)`` of
|
|
196
|
-
``(1 + n_augmentation)`` times the original size.
|
|
205
|
+
``(1 + n_augmentation)`` times the original size. When the transform
|
|
206
|
+
mixes targets, as :class:`braindecode.augmentation.Mixup` does, ``y``
|
|
207
|
+
stays the ``(y_a, y_b, lam)`` triple and the clean originals get a
|
|
208
|
+
mixing coefficient of one.
|
|
197
209
|
"""
|
|
198
210
|
|
|
199
211
|
def __init__(self, transform, device=None, n_augmentation=0):
|
|
@@ -216,6 +228,18 @@ class _AugmentationCollate:
|
|
|
216
228
|
aug_X, aug_y = self.transform(X, y)
|
|
217
229
|
xs.append(aug_X)
|
|
218
230
|
ys.append(aug_y)
|
|
231
|
+
mixed_ys = [
|
|
232
|
+
aug_y
|
|
233
|
+
for aug_y in ys
|
|
234
|
+
if isinstance(aug_y, (tuple, list)) and len(aug_y) == 3
|
|
235
|
+
]
|
|
236
|
+
if mixed_ys:
|
|
237
|
+
# a target-mixing transform such as Mixup returns (y_a, y_b, lam),
|
|
238
|
+
# so the parts are concatenated one by one and the untouched copies
|
|
239
|
+
# are given a mixing coefficient of one
|
|
240
|
+
lam_dtype = mixed_ys[0][2].dtype
|
|
241
|
+
ys = [_as_mixed_target(aug_y, lam_dtype) for aug_y in ys]
|
|
242
|
+
return torch.cat(xs), tuple(torch.cat(part) for part in zip(*ys))
|
|
219
243
|
return torch.cat(xs), torch.cat(ys)
|
|
220
244
|
|
|
221
245
|
|
|
@@ -1065,13 +1065,11 @@ def mixup(
|
|
|
1065
1065
|
batch_size, n_channels, n_times = X.shape
|
|
1066
1066
|
|
|
1067
1067
|
X_mix = torch.zeros((batch_size, n_channels, n_times)).to(device)
|
|
1068
|
-
y_a =
|
|
1069
|
-
y_b =
|
|
1068
|
+
y_a = y.clone()
|
|
1069
|
+
y_b = y[idx_perm].clone()
|
|
1070
1070
|
|
|
1071
1071
|
for idx in range(batch_size):
|
|
1072
1072
|
X_mix[idx] = lam[idx] * X[idx] + (1 - lam[idx]) * X[idx_perm[idx]]
|
|
1073
|
-
y_a[idx] = y[idx]
|
|
1074
|
-
y_b[idx] = y[idx_perm[idx]]
|
|
1075
1073
|
|
|
1076
1074
|
return X_mix, (y_a, y_b, lam)
|
|
1077
1075
|
|
|
@@ -1068,16 +1068,18 @@ class Mixup(Transform):
|
|
|
1068
1068
|
device = X.device
|
|
1069
1069
|
batch_size, _, _ = X.shape
|
|
1070
1070
|
|
|
1071
|
+
# lam follows the dtype of X, numpy draws float64 and that would leak
|
|
1072
|
+
# into the mixed signal and into the loss returned by mixup_criterion
|
|
1071
1073
|
if self.alpha > 0:
|
|
1072
1074
|
if self.beta_per_sample:
|
|
1073
1075
|
lam = torch.as_tensor(
|
|
1074
1076
|
self.rng.beta(self.alpha, self.alpha, batch_size)
|
|
1075
|
-
).to(device)
|
|
1077
|
+
).to(device=device, dtype=X.dtype)
|
|
1076
1078
|
else:
|
|
1077
|
-
lam = torch.ones(batch_size).to(device)
|
|
1079
|
+
lam = torch.ones(batch_size, dtype=X.dtype).to(device)
|
|
1078
1080
|
lam *= self.rng.beta(self.alpha, self.alpha)
|
|
1079
1081
|
else:
|
|
1080
|
-
lam = torch.ones(batch_size).to(device)
|
|
1082
|
+
lam = torch.ones(batch_size, dtype=X.dtype).to(device)
|
|
1081
1083
|
|
|
1082
1084
|
idx_perm = torch.as_tensor(
|
|
1083
1085
|
self.rng.permutation(
|
|
@@ -228,8 +228,9 @@ class EEGClassifier(_EEGNeuralNet, NeuralNetClassifier):
|
|
|
228
228
|
if return_targets:
|
|
229
229
|
return preds, X.get_metadata()["target"].to_numpy()
|
|
230
230
|
return preds
|
|
231
|
+
self.check_is_fitted()
|
|
231
232
|
return predict_trials(
|
|
232
|
-
module=self.
|
|
233
|
+
module=self.module_,
|
|
233
234
|
dataset=X,
|
|
234
235
|
return_targets=return_targets,
|
|
235
236
|
batch_size=self.batch_size,
|
{braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/datasets/base.py
RENAMED
|
@@ -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.dev169208086 → braindecode-1.8.0.dev180935432}/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
|
|
|
@@ -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
|
|
@@ -117,9 +117,15 @@ class SleepPhysionet(BaseConcatDataset):
|
|
|
117
117
|
sleep_event_inds = np.where(mask)[0]
|
|
118
118
|
|
|
119
119
|
# Crop raw
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
120
|
+
a_tmin = annots[sleep_event_inds[0]]
|
|
121
|
+
a_tmax = annots[sleep_event_inds[-1]]
|
|
122
|
+
tmin = a_tmin["onset"] - crop_wake_mins * 60
|
|
123
|
+
tmax = a_tmax["onset"] + a_tmax["duration"] + crop_wake_mins * 60
|
|
124
|
+
raw.crop(
|
|
125
|
+
tmin=max(tmin, raw.times[0]),
|
|
126
|
+
tmax=min(tmax, raw.times[-1] + 1 / raw.info["sfreq"]),
|
|
127
|
+
include_tmax=False,
|
|
128
|
+
)
|
|
123
129
|
|
|
124
130
|
# Rename EEG channels
|
|
125
131
|
ch_names = {i: i.replace("EEG ", "") for i in raw.ch_names if "EEG" in i}
|
{braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/eegneuralnet.py
RENAMED
|
@@ -1,5 +1,6 @@
|
|
|
1
1
|
# Authors: Bruno Aristimunha <b.aristimunha@gmail.com>
|
|
2
2
|
# Pierre Guetschel <pierre.guetschel@gmail.com>
|
|
3
|
+
# Sarthak Tayal <sarthaktayal2@gmail.com>
|
|
3
4
|
#
|
|
4
5
|
# License: BSD (3-clause)
|
|
5
6
|
|
|
@@ -126,7 +127,8 @@ class _EEGNeuralNet(NeuralNet, abc.ABC):
|
|
|
126
127
|
self._last_window_inds_ = None
|
|
127
128
|
|
|
128
129
|
def predict_with_window_inds_and_ys(self, dataset):
|
|
129
|
-
self.module.
|
|
130
|
+
# self.module can still be a name or a class, self.module_ is the built one
|
|
131
|
+
self.module_.eval()
|
|
130
132
|
preds = []
|
|
131
133
|
i_window_in_trials = []
|
|
132
134
|
i_window_stops = []
|
|
@@ -142,7 +144,7 @@ class _EEGNeuralNet(NeuralNet, abc.ABC):
|
|
|
142
144
|
i_window_in_trials.append(i[0].cpu().numpy())
|
|
143
145
|
i_window_stops.append(i[2].cpu().numpy())
|
|
144
146
|
with torch.no_grad():
|
|
145
|
-
preds.append(to_numpy(self.
|
|
147
|
+
preds.append(to_numpy(self.module_.forward(X.to(self.device))))
|
|
146
148
|
window_ys.append(y.cpu().numpy())
|
|
147
149
|
preds = np.concatenate(preds)
|
|
148
150
|
i_window_in_trials = np.concatenate(i_window_in_trials)
|
{braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/__init__.py
RENAMED
|
@@ -71,6 +71,7 @@ from .util import (
|
|
|
71
71
|
models_mandatory_parameters,
|
|
72
72
|
positions_from_chs_info,
|
|
73
73
|
)
|
|
74
|
+
from .zuna import ZUNA
|
|
74
75
|
|
|
75
76
|
# Call this last in order to make sure the dataset list is populated with
|
|
76
77
|
# the models imported in this file.
|
|
@@ -145,6 +146,7 @@ __all__ = [
|
|
|
145
146
|
"TIDNet",
|
|
146
147
|
"TSception",
|
|
147
148
|
"USleep",
|
|
149
|
+
"ZUNA",
|
|
148
150
|
"build_model_config",
|
|
149
151
|
"_init_models_dict",
|
|
150
152
|
"models_mandatory_parameters",
|
{braindecode-1.8.0.dev169208086 → braindecode-1.8.0.dev180935432}/braindecode/models/atcnet.py
RENAMED
|
@@ -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
|