braindecode 1.8.0.dev1128__tar.gz → 1.8.0.dev169056309__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.
Files changed (181) hide show
  1. braindecode-1.8.0.dev169056309/NOTICE.txt +65 -0
  2. {braindecode-1.8.0.dev1128/braindecode.egg-info → braindecode-1.8.0.dev169056309}/PKG-INFO +8 -5
  3. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/README.rst +6 -3
  4. braindecode-1.8.0.dev169056309/braindecode/datasets/_notebook_viewer.py +255 -0
  5. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/base.py +18 -0
  6. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/functional/functions.py +4 -3
  7. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/atcnet.py +2 -2
  8. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/attn_sleep.py +87 -28
  9. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/base.py +3 -3
  10. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/brainmodule.py +1 -1
  11. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/ctnet.py +10 -7
  12. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/dance.py +1 -1
  13. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegconformer.py +1 -1
  14. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegminer.py +2 -1
  15. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegsimpleconv.py +9 -7
  16. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/emg2qwerty.py +5 -4
  17. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/ifnet.py +10 -6
  18. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/luna.py +31 -5
  19. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/medformer.py +1 -1
  20. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/meta_neuromotor.py +1 -1
  21. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/mvpformer.py +1 -1
  22. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sleep_stager_blanco_2020.py +5 -9
  23. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sleep_stager_chambon_2018.py +1 -7
  24. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sparcnet.py +2 -2
  25. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sstdpn.py +1 -1
  26. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/summary.csv +1 -1
  27. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/tcformer.py +1 -1
  28. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/tidnet.py +1 -7
  29. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/usleep.py +3 -7
  30. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/zuna.py +26 -12
  31. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/windowers.py +126 -30
  32. braindecode-1.8.0.dev169056309/braindecode/version.py +1 -0
  33. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309/braindecode.egg-info}/PKG-INFO +8 -5
  34. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode.egg-info/SOURCES.txt +1 -0
  35. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode.egg-info/requires.txt +1 -1
  36. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/whats_new.rst +71 -2
  37. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/pyproject.toml +1 -1
  38. braindecode-1.8.0.dev1128/NOTICE.txt +0 -25
  39. braindecode-1.8.0.dev1128/braindecode/version.py +0 -1
  40. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/LICENSE.txt +0 -0
  41. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/MANIFEST.in +0 -0
  42. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/__init__.py +0 -0
  43. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/augmentation/__init__.py +0 -0
  44. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/augmentation/base.py +0 -0
  45. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/augmentation/functional.py +0 -0
  46. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/augmentation/transforms.py +0 -0
  47. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/classifier.py +0 -0
  48. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/__init__.py +0 -0
  49. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bbci.py +0 -0
  50. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bcicomp.py +0 -0
  51. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/__init__.py +0 -0
  52. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/datasets.py +0 -0
  53. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/format.py +0 -0
  54. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/hub.py +0 -0
  55. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/hub_format.py +0 -0
  56. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/hub_io.py +0 -0
  57. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/hub_validation.py +0 -0
  58. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/bids/iterable.py +0 -0
  59. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/chb_mit.py +0 -0
  60. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/collate.py +0 -0
  61. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/mne.py +0 -0
  62. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/moabb.py +0 -0
  63. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/nmt.py +0 -0
  64. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/registry.py +0 -0
  65. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/siena.py +0 -0
  66. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/sleep_physio_challe_18.py +0 -0
  67. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/sleep_physionet.py +0 -0
  68. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/tuh.py +0 -0
  69. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/utils.py +0 -0
  70. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datasets/xy.py +0 -0
  71. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datautil/__init__.py +0 -0
  72. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datautil/channel_utils.py +0 -0
  73. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datautil/hub_formats.py +0 -0
  74. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datautil/serialization.py +0 -0
  75. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/datautil/util.py +0 -0
  76. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/eegneuralnet.py +0 -0
  77. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/functional/__init__.py +0 -0
  78. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/functional/initialization.py +0 -0
  79. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/__init__.py +0 -0
  80. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/attentionbasenet.py +0 -0
  81. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/bendr.py +0 -0
  82. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/biot.py +0 -0
  83. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/cbramod.py +0 -0
  84. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/codebrain.py +0 -0
  85. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/config.py +0 -0
  86. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/contrawr.py +0 -0
  87. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/deep4.py +0 -0
  88. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/deepsleepnet.py +0 -0
  89. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/dgcnn.py +0 -0
  90. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegdino.py +0 -0
  91. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eeginception_erp.py +0 -0
  92. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eeginception_mi.py +0 -0
  93. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegitnet.py +0 -0
  94. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegnet.py +0 -0
  95. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegnex.py +0 -0
  96. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegpt.py +0 -0
  97. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegsym.py +0 -0
  98. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/eegtcnet.py +0 -0
  99. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/fbcnet.py +0 -0
  100. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/fblightconvnet.py +0 -0
  101. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/fbmsnet.py +0 -0
  102. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/hybrid.py +0 -0
  103. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/interpolated.py +0 -0
  104. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/labram.py +0 -0
  105. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/msvtnet.py +0 -0
  106. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/patchedtransformer.py +0 -0
  107. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/reve.py +0 -0
  108. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sccnet.py +0 -0
  109. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/shallow_fbcsp.py +0 -0
  110. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/signal_jepa.py +0 -0
  111. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/sinc_shallow.py +0 -0
  112. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/steegformer.py +0 -0
  113. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/syncnet.py +0 -0
  114. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/tcn.py +0 -0
  115. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/tsinception.py +0 -0
  116. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/models/util.py +0 -0
  117. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/__init__.py +0 -0
  118. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/activation.py +0 -0
  119. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/attention.py +0 -0
  120. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/blocks.py +0 -0
  121. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/convolution.py +0 -0
  122. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/dance_modules.py +0 -0
  123. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/filter.py +0 -0
  124. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/interpolation.py +0 -0
  125. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/layers.py +0 -0
  126. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/linear.py +0 -0
  127. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/parametrization.py +0 -0
  128. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/stats.py +0 -0
  129. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/util.py +0 -0
  130. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/modules/wrapper.py +0 -0
  131. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/__init__.py +0 -0
  132. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/eegprep_preprocess.py +0 -0
  133. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/mne_preprocess.py +0 -0
  134. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/preprocess.py +0 -0
  135. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/preprocessing/util.py +0 -0
  136. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/regressor.py +0 -0
  137. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/samplers/__init__.py +0 -0
  138. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/samplers/base.py +0 -0
  139. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/samplers/ssl.py +0 -0
  140. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/training/__init__.py +0 -0
  141. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/training/callbacks.py +0 -0
  142. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/training/losses.py +0 -0
  143. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/training/scoring.py +0 -0
  144. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/util.py +0 -0
  145. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/__init__.py +0 -0
  146. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/attribution.py +0 -0
  147. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/confusion_matrices.py +0 -0
  148. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/frequency.py +0 -0
  149. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/metrics.py +0 -0
  150. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/sanity.py +0 -0
  151. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode/visualization/topology.py +0 -0
  152. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode.egg-info/dependency_links.txt +0 -0
  153. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/braindecode.egg-info/top_level.txt +0 -0
  154. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/Makefile +0 -0
  155. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/_templates/autosummary/class.rst +0 -0
  156. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/_templates/autosummary/class_in_subdir.rst +0 -0
  157. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/_templates/autosummary/function.rst +0 -0
  158. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/_templates/autosummary/function_in_subdir.rst +0 -0
  159. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/api.rst +0 -0
  160. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/cite.rst +0 -0
  161. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/conf.py +0 -0
  162. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/help.rst +0 -0
  163. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/index.rst +0 -0
  164. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/install/install.rst +0 -0
  165. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/install/install_pip.rst +0 -0
  166. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/install/install_source.rst +0 -0
  167. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/attention.rst +0 -0
  168. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/channel.rst +0 -0
  169. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/convolution.rst +0 -0
  170. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/filterbank.rst +0 -0
  171. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/gnn.rst +0 -0
  172. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/interpretable.rst +0 -0
  173. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/lbm.rst +0 -0
  174. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/recurrent.rst +0 -0
  175. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/categorization/spd.rst +0 -0
  176. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/models.rst +0 -0
  177. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/models_categorization.rst +0 -0
  178. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/models_table.rst +0 -0
  179. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/models/models_visualization.rst +0 -0
  180. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/docs/sg_execution_times.rst +0 -0
  181. {braindecode-1.8.0.dev1128 → braindecode-1.8.0.dev169056309}/setup.cfg +0 -0
@@ -0,0 +1,65 @@
1
+ # BRAINDECODE Notice
2
+
3
+ ## Licensed Components
4
+
5
+ ### BSD-3-Clause Licensed Files
6
+
7
+ All files within the `braindecode/` package are licensed under the BSD-3-Clause
8
+ License, except for those listed in the sections below.
9
+
10
+ ### CC BY-NC 4.0 Licensed Files
11
+
12
+ The following components are licensed under the Creative Commons Attribution-NonCommercial 4.0 International License:
13
+
14
+ - `braindecode/models/eegminer.py`
15
+ - `braindecode/models/meta_neuromotor.py`
16
+ - `braindecode/models/brainmodule.py`
17
+
18
+ As well as class later imported into the `braindecode.models.module` named as GeneralizedGaussianFilter.
19
+
20
+ The `meta_neuromotor.py` file is a derivative of
21
+ `facebookresearch/generic-neuromotor-interface`, released by Meta Platforms,
22
+ Inc. under CC BY-NC 4.0, and inherits the same noncommercial terms.
23
+
24
+ ### CC BY-NC-SA 4.0 Licensed Files
25
+
26
+ The following components are licensed under the Creative Commons
27
+ Attribution-NonCommercial-ShareAlike 4.0 International License:
28
+
29
+ - `braindecode/models/emg2qwerty.py`
30
+
31
+ The `emg2qwerty.py` file is a derivative of `facebookresearch/emg2qwerty`,
32
+ released by Meta Platforms, Inc. under CC BY-NC-SA 4.0, and inherits the
33
+ same noncommercial ShareAlike terms.
34
+
35
+ ### MIT Licensed Files
36
+
37
+ The following components are licensed under the MIT License:
38
+
39
+ - `braindecode/models/ctnet.py`
40
+ - `braindecode/models/dance.py`
41
+ - `braindecode/models/medformer.py`
42
+ - `braindecode/models/tcformer.py`
43
+ - `braindecode/models/ifnet.py`
44
+
45
+ ### Apache-2.0 Licensed Files
46
+
47
+ The following components are licensed under the Apache License 2.0:
48
+
49
+ - `braindecode/models/mvpformer.py`
50
+ - `braindecode/models/zuna.py`
51
+ - `braindecode/models/luna.py`
52
+
53
+ ## License Links
54
+
55
+ - [BSD-3-Clause License](https://opensource.org/licenses/BSD-3-Clause)
56
+ - [CC BY-NC 4.0 License](https://creativecommons.org/licenses/by-nc/4.0/)
57
+ - [CC BY-NC-SA 4.0 License](https://creativecommons.org/licenses/by-nc-sa/4.0/)
58
+ - [MIT License](https://opensource.org/licenses/MIT)
59
+ - [Apache-2.0 License](https://www.apache.org/licenses/LICENSE-2.0)
60
+
61
+ ## Note
62
+
63
+ This list covers files whose own headers declare a license other than
64
+ BSD-3-Clause. Per-file provenance review of the remaining models is
65
+ tracked separately.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: braindecode
3
- Version: 1.8.0.dev1128
3
+ Version: 1.8.0.dev169056309
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>
@@ -34,7 +34,6 @@ Requires-Dist: pandas
34
34
  Requires-Dist: wfdb>=4.3.1
35
35
  Requires-Dist: linear_attention_transformer
36
36
  Requires-Dist: docstring_inheritance
37
- Requires-Dist: rotary_embedding_torch
38
37
  Requires-Dist: pydantic>=2.0
39
38
  Provides-Extra: moabb
40
39
  Requires-Dist: moabb>=1.4.3; extra == "moabb"
@@ -51,6 +50,7 @@ Requires-Dist: pytest-cov; extra == "tests"
51
50
  Requires-Dist: codecov; extra == "tests"
52
51
  Requires-Dist: pytest_cases; extra == "tests"
53
52
  Requires-Dist: mypy; extra == "tests"
53
+ Requires-Dist: ipython; extra == "tests"
54
54
  Requires-Dist: transformers>=4.57.0; extra == "tests"
55
55
  Requires-Dist: bids_validator; extra == "tests"
56
56
  Provides-Extra: typing
@@ -260,7 +260,10 @@ This project is primarily licensed under the BSD-3-Clause License.
260
260
  Additional Components
261
261
  =====================
262
262
 
263
- Some components within this repository are licensed under the Creative Commons
264
- Attribution-NonCommercial 4.0 International License.
263
+ Some components within this repository are licensed under other licenses, including
264
+ Creative Commons Attribution-NonCommercial 4.0 International (CC BY-NC 4.0), Creative
265
+ Commons Attribution-NonCommercial-ShareAlike 4.0 International (CC BY-NC-SA 4.0), MIT
266
+ and Apache-2.0.
265
267
 
266
- Please refer to the ``LICENSE`` and ``NOTICE`` files for more detailed information.
268
+ Please refer to the ``LICENSE`` and ``NOTICE`` files for the per-file list and more
269
+ detailed information.
@@ -170,7 +170,10 @@ This project is primarily licensed under the BSD-3-Clause License.
170
170
  Additional Components
171
171
  =====================
172
172
 
173
- Some components within this repository are licensed under the Creative Commons
174
- Attribution-NonCommercial 4.0 International License.
173
+ Some components within this repository are licensed under other licenses, including
174
+ Creative Commons Attribution-NonCommercial 4.0 International (CC BY-NC 4.0), Creative
175
+ Commons Attribution-NonCommercial-ShareAlike 4.0 International (CC BY-NC-SA 4.0), MIT
176
+ and Apache-2.0.
175
177
 
176
- Please refer to the ``LICENSE`` and ``NOTICE`` files for more detailed information.
178
+ Please refer to the ``LICENSE`` and ``NOTICE`` files for the per-file list and more
179
+ detailed information.
@@ -0,0 +1,255 @@
1
+ # Authors: Bruno Aristimunha <b.aristimunha@gmail.com>
2
+ #
3
+ # License: BSD (3-clause)
4
+ """Serverless in-notebook viewer for file-backed recordings.
5
+
6
+ The bytes on disk behind a dataset element are inlined in the cell output
7
+ as base64 and handed to the deployed eegdash-viewer over its ``postMessage``
8
+ bridge (``docs/embedding.md`` in https://github.com/eegdash/eegdash-viewer):
9
+ no server, no CORS. The output is an iframe plus an inline script, so it
10
+ renders when the cell ran in your session or the saved notebook is trusted
11
+ (``jupyter trust``).
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import base64
17
+ import json
18
+ import os
19
+ import uuid
20
+ from pathlib import Path
21
+ from urllib.parse import urlsplit
22
+
23
+ import mne_bids
24
+ from mne.utils import _soft_import
25
+
26
+ CDN = "https://eegdash.github.io/eegdash-viewer"
27
+ MAX_BYTES = 64 * 2**20 # base64 output per call; it is saved with the notebook
28
+ EXTENSIONS = {
29
+ ".set",
30
+ ".edf",
31
+ ".bdf",
32
+ ".vhdr",
33
+ ".fif",
34
+ ".snirf",
35
+ ".nwb",
36
+ } # viewer readers
37
+ _BIDS = {"eeg", "ieeg", "emg", "meg", "nirs"}
38
+ _SIBLINGS = {".vhdr": (".eeg", ".vmrk"), ".set": (".fdt",)} # travel with the header
39
+ _HEADER = {s: h for h, ss in _SIBLINGS.items() for s in ss} # data file -> header
40
+
41
+ _SCRIPT = """<iframe id=%(id)s title="eegdash trace viewer" style="width:100%%;height:%(height)spx;
42
+ border:1px solid var(--jp-border-color1,#d9dce1);border-radius:6px;background:transparent"></iframe>
43
+ <script>
44
+ (function () {
45
+ var self = document.currentScript, id = %(id)s; // Lab re-runs scripts in place; VS Code/nbclassic elsewhere
46
+ var frame = (self && self.previousElementSibling && self.previousElementSibling.tagName === "IFRAME")
47
+ ? self.previousElementSibling : document.getElementById(id);
48
+ if (!frame) { console.error("eegdash viewer: output iframe " + id + " not found"); return; }
49
+ var payload = %(payload)s, origin = %(origin)s, files = null, pose = null;
50
+ function decode(b64) {
51
+ if (Uint8Array.fromBase64) return Uint8Array.fromBase64(b64);
52
+ var bin = atob(b64), out = new Uint8Array(bin.length);
53
+ for (var i = 0; i < bin.length; i++) out[i] = bin.charCodeAt(i);
54
+ return out;
55
+ }
56
+ function send(target) {
57
+ try {
58
+ if (!files) {
59
+ files = payload.files.map(function (f) { return new File([decode(f.b64)], f.name); });
60
+ pose = payload.pose ? "data:application/json;base64," + payload.pose : null;
61
+ payload = null;
62
+ }
63
+ frame.contentWindow.postMessage({ type: "eegdash-viewer:open", files: files, pose: pose }, target || origin);
64
+ } catch (err) {
65
+ frame.insertAdjacentHTML("afterend", '<div style="font:12px system-ui;color:#b3261e">eegdash viewer: '
66
+ + String(err.message).replace(/</g, "&lt;") + "</div>");
67
+ }
68
+ }
69
+ window.addEventListener("message", function onMessage(e) {
70
+ if (e.source === frame.contentWindow && e.data && e.data.type === "eegdash-viewer:ready") send(e.origin);
71
+ else if (!frame.isConnected) window.removeEventListener("message", onMessage);
72
+ });
73
+ frame.src = %(src)s; // after the listener, so "ready" can never precede it
74
+ })();
75
+ </script>"""
76
+
77
+
78
+ def recording_files(recording: Path) -> tuple[list[Path], Path | None]:
79
+ """``(files, pose)`` to inline for a recording: the header first, then the
80
+ split-format siblings, the BIDS-inherited ``_channels.tsv``/``_events.tsv``
81
+ (``mne_bids`` parses the name; plain names get none) and, separately, the
82
+ ``<prefix>_desc-pose.json`` hand-pose sidecar next to it. Symlinks keep
83
+ their name (git-annex/datalad)."""
84
+ rec = Path(recording)
85
+ if not rec.exists():
86
+ raise ValueError(
87
+ f"{rec.name}: file not found (a datalad symlink may need `datalad get`)"
88
+ )
89
+ if (
90
+ rec.is_dir()
91
+ or rec.suffix.lower() not in EXTENSIONS
92
+ or rec.stem.endswith("_epo")
93
+ ):
94
+ raise ValueError(
95
+ f"{rec.name}: the viewer opens raw recordings in {' '.join(sorted(EXTENSIONS))}"
96
+ )
97
+ sidecars: list[Path | None] = []
98
+ try:
99
+ bids = mne_bids.get_bids_path_from_fname(rec, check=False)
100
+ if bids.subject is not None: # hyphen-free names parse, but are not BIDS
101
+ sidecars = [
102
+ bids.find_matching_sidecar(
103
+ suffix=s, extension=".tsv", on_error="ignore"
104
+ )
105
+ for s in ("channels", "events")
106
+ ]
107
+ except (
108
+ KeyError,
109
+ ValueError,
110
+ ): # not a BIDS name / unknown entity somewhere in the tree
111
+ pass
112
+ files = [rec]
113
+ for p in [
114
+ rec.with_suffix(e) for e in _SIBLINGS.get(rec.suffix.lower(), ())
115
+ ] + sidecars:
116
+ if p and p.is_file():
117
+ if p not in files:
118
+ files.append(p)
119
+ elif p and p.is_symlink(): # dangling: git-annex/datalad content not fetched
120
+ raise ValueError(f"{p.name}: dangling symlink (try `datalad get`)")
121
+ stem, _, token = rec.stem.rpartition("_")
122
+ pose = rec.with_name((stem if token in _BIDS else rec.stem) + "_desc-pose.json")
123
+ return files, pose if pose.is_file() else None
124
+
125
+
126
+ def build_viewer_html(
127
+ recording: Path,
128
+ *,
129
+ height: int = 520,
130
+ cdn_url: str = CDN,
131
+ max_bytes: int = MAX_BYTES,
132
+ ) -> str:
133
+ """Viewer iframe + inlined bytes + bridge script for one recording."""
134
+ url = urlsplit(cdn_url)
135
+ if (
136
+ url.scheme not in ("http", "https")
137
+ or not url.netloc
138
+ or url.username
139
+ or url.query
140
+ or url.fragment
141
+ or url.path.endswith(("index.html", "index.htm"))
142
+ ):
143
+ raise ValueError(
144
+ f"cdn_url must be the viewer's base http(s) URL, got {cdn_url!r}"
145
+ )
146
+ files, pose = recording_files(recording)
147
+ encoded = sum(
148
+ 4 * -(-p.stat().st_size // 3) for p in files + ([pose] if pose else [])
149
+ )
150
+ if encoded > max_bytes:
151
+ raise ValueError(
152
+ f"{files[0].name}: {encoded / 2**20:.1f} MiB of base64 would be inlined into the "
153
+ f"notebook output (max_bytes={max_bytes / 2**20:.1f} MiB); crop/downsample or raise it"
154
+ )
155
+ rec = files[0]
156
+ # The viewer picks the recording by its *_<datatype>.<ext> name: plain names
157
+ # are posted as <stem>_eeg<ext>, an EEGLAB .fdt next to the posted .set name.
158
+ head = (
159
+ rec.name
160
+ if rec.stem.rpartition("_")[2] in _BIDS
161
+ else f"{rec.stem}_eeg{rec.suffix.lower()}"
162
+ )
163
+ names = [head] + [
164
+ head.rsplit("_", 1)[0] + "_eeg.fdt" if p.suffix.lower() == ".fdt" else p.name
165
+ for p in files[1:]
166
+ ]
167
+ b64 = [base64.b64encode(p.read_bytes()).decode() for p in files]
168
+ literals = {
169
+ "id": f"eegdash-viewer-{uuid.uuid4().hex[:8]}",
170
+ "height": int(height),
171
+ "payload": {
172
+ "files": [{"name": n, "b64": b} for n, b in zip(names, b64)],
173
+ "pose": base64.b64encode(pose.read_bytes()).decode() if pose else None,
174
+ },
175
+ "origin": f"{url.scheme}://{url.netloc}",
176
+ "src": f"{url.geturl().rstrip('/')}/index.html?embed=1",
177
+ }
178
+ return _SCRIPT % {
179
+ k: json.dumps(v).replace("<", "\\u003c") for k, v in literals.items()
180
+ }
181
+
182
+
183
+ def _recording(dataset, index: int) -> Path:
184
+ """File behind ``dataset.datasets[index]``: the one its mne ``raw`` reads (a
185
+ data file maps back to its header; lazily downloading datasets fetch it when
186
+ ``raw`` is accessed) or the recorded ``description["path"]``."""
187
+ ds = dataset.datasets[index]
188
+ names = [
189
+ Path(f)
190
+ for f in getattr(getattr(ds, "raw", None), "filenames", None) or ()
191
+ if isinstance(f, (str, os.PathLike))
192
+ ]
193
+ if len(names) > 1:
194
+ raise ValueError(
195
+ f"{type(dataset).__name__}[{index}]: split recordings are not supported"
196
+ )
197
+ desc = getattr(ds, "description", None) # dict or pandas Series
198
+ recorded = desc.get("path") if desc is not None else None
199
+ path = (
200
+ names[0]
201
+ if names
202
+ else Path(recorded)
203
+ if isinstance(recorded, (str, os.PathLike))
204
+ else None
205
+ )
206
+ if path is None:
207
+ raise ValueError(
208
+ f"{type(dataset).__name__}[{index}] is not backed by a recording file"
209
+ )
210
+ header = path.with_suffix(_HEADER.get(path.suffix.lower(), path.suffix))
211
+ return header if header.is_file() else path
212
+
213
+
214
+ def plot(
215
+ dataset,
216
+ index: int = 0,
217
+ *,
218
+ height: int = 520,
219
+ cdn_url: str = CDN,
220
+ max_bytes: int = MAX_BYTES,
221
+ ):
222
+ """Show one recording in the eegdash-viewer inside a Jupyter cell.
223
+
224
+ Serverless: the recording bytes (as on disk) are inlined in the output and
225
+ pushed to the viewer at ``cdn_url`` over ``postMessage``; the output renders
226
+ when the cell ran in your session or the notebook is trusted. A
227
+ ``*_desc-pose.json`` sidecar next to the recording adds the synchronized
228
+ hand-pose panel. Needs IPython (soft dependency).
229
+
230
+ Parameters
231
+ ----------
232
+ index : int
233
+ Recording to display.
234
+ height : int
235
+ Viewer height in pixels.
236
+ cdn_url : str
237
+ Base URL of a deployed eegdash-viewer.
238
+ max_bytes : int
239
+ Refuse to inline more than this much base64 (default 64 MiB); the
240
+ payload is saved with the notebook and, like any cell output, stays
241
+ referenced by IPython's ``Out`` history for the session.
242
+
243
+ Returns
244
+ -------
245
+ IPython.display.HTML
246
+ """
247
+ ipython = _soft_import("IPython", purpose=f"{type(dataset).__name__}.plot()")
248
+ return ipython.display.HTML(
249
+ build_viewer_html(
250
+ _recording(dataset, index),
251
+ height=height,
252
+ cdn_url=cdn_url,
253
+ max_bytes=max_bytes,
254
+ )
255
+ )
@@ -32,6 +32,7 @@ from mne.utils.docs import deprecated
32
32
  from torch.utils.data import ConcatDataset, Dataset, IterableDataset
33
33
  from typing_extensions import TypeVar
34
34
 
35
+ from ._notebook_viewer import plot as _viewer_plot
35
36
  from .bids.hub import HubDatasetMixin
36
37
  from .bids.hub_io import _restore_nan_from_json
37
38
  from .registry import register_dataset
@@ -59,6 +60,7 @@ def _html_row(label, value):
59
60
 
60
61
  _METADATA_INTERNAL_COLS = {
61
62
  "i_window_in_trial",
63
+ "i_trial_in_dataset",
62
64
  "i_start_in_trial",
63
65
  "i_stop_in_trial",
64
66
  "target",
@@ -1173,6 +1175,8 @@ class BaseConcatDataset(ConcatDataset, HubDatasetMixin, Generic[T]):
1173
1175
  If True, defer computing cumulative sizes until length or item access.
1174
1176
  """
1175
1177
 
1178
+ plot = _viewer_plot # eegdash-viewer embed; a direct member so the API docs list it
1179
+
1176
1180
  datasets: list[T]
1177
1181
 
1178
1182
  def __init__(
@@ -1374,6 +1378,20 @@ class BaseConcatDataset(ConcatDataset, HubDatasetMixin, Generic[T]):
1374
1378
  "datasets are WindowsDataset."
1375
1379
  )
1376
1380
 
1381
+ for ds in self.datasets:
1382
+ if hasattr(ds, "_windows") and ds._windows is not None:
1383
+ df = ds._windows.metadata
1384
+ else:
1385
+ df = ds.metadata
1386
+ if (
1387
+ "i_trial_in_dataset" in df.columns
1388
+ and "i_trial_in_dataset" in ds.description
1389
+ ):
1390
+ raise ValueError(
1391
+ "Dataset descriptions cannot contain the reserved window "
1392
+ "metadata key 'i_trial_in_dataset'."
1393
+ )
1394
+
1377
1395
  all_dfs = list()
1378
1396
  for ds in self.datasets:
1379
1397
  if hasattr(ds, "_windows") and ds._windows is not None:
@@ -40,7 +40,7 @@ def drop_path(
40
40
  x : torch.Tensor
41
41
  input tensor
42
42
  drop_prob : float, optional
43
- survival rate (i.e. probability of being kept), by default 0.0
43
+ probability of dropping a path, by default 0.0
44
44
  training : bool, optional
45
45
  whether the model is in training mode, by default False
46
46
  scale_by_keep : bool, optional
@@ -68,9 +68,10 @@ def drop_path(
68
68
  shape = (x.shape[0],) + (1,) * (
69
69
  x.ndim - 1
70
70
  ) # work with diff dim tensors, not just 2D ConvNets
71
- random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
71
+ probabilities = torch.full(shape, keep_prob, dtype=torch.float32, device=x.device)
72
+ random_tensor = torch.bernoulli(probabilities).to(dtype=x.dtype)
72
73
  if keep_prob > 0.0 and scale_by_keep:
73
- random_tensor.div_(keep_prob)
74
+ random_tensor = random_tensor / keep_prob
74
75
  return x * random_tensor
75
76
 
76
77
 
@@ -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
- att_dropout : float
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
- tcn_dropout : float
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
@@ -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. Should be a PyTorch activation
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.ReLU
69
- Activation function class to apply in the Mask R-CNN layer.
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 = self._len_last_layer(self.n_times)
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
- d_model * after_reduced_cnn_size, self.n_outputs
162
- )
163
-
164
- def _len_last_layer(self, input_size):
165
- self.feature_extractor.eval()
166
- with torch.no_grad():
167
- out = self.feature_extractor(torch.Tensor(1, 1, input_size))
168
- self.feature_extractor.train()
169
- return len(out.flatten())
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(inplace=True),
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(inplace=True)
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, activate=activation_se
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, activate: type[nn.Module] = nn.ReLU
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(block(self.inplanes, planes, stride, downsample))
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, activate=activate))
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
- assert d_model % h == 0
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,