braindecode 1.8.0.dev174836782__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.
Files changed (179) hide show
  1. {braindecode-1.8.0.dev174836782/braindecode.egg-info → braindecode-1.8.0.dev180935432}/PKG-INFO +1 -1
  2. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/base.py +25 -1
  3. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/functional.py +2 -4
  4. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/transforms.py +5 -3
  5. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/base.py +15 -0
  6. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub.py +90 -25
  7. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_io.py +67 -5
  8. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/sleep_physionet.py +9 -3
  9. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/atcnet.py +2 -2
  10. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/attn_sleep.py +87 -28
  11. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/base.py +40 -1
  12. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/ctnet.py +9 -6
  13. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegsimpleconv.py +8 -6
  14. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/ifnet.py +8 -5
  15. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sleep_stager_blanco_2020.py +5 -9
  16. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sleep_stager_chambon_2018.py +1 -7
  17. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sparcnet.py +2 -2
  18. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/tidnet.py +1 -7
  19. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/windowers.py +126 -30
  20. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/samplers/base.py +40 -2
  21. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/training/losses.py +7 -4
  22. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/util.py +16 -10
  23. braindecode-1.8.0.dev180935432/braindecode/version.py +1 -0
  24. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432/braindecode.egg-info}/PKG-INFO +1 -1
  25. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/whats_new.rst +115 -2
  26. braindecode-1.8.0.dev174836782/braindecode/version.py +0 -1
  27. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/LICENSE.txt +0 -0
  28. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/MANIFEST.in +0 -0
  29. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/NOTICE.txt +0 -0
  30. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/README.rst +0 -0
  31. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/__init__.py +0 -0
  32. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/augmentation/__init__.py +0 -0
  33. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/classifier.py +0 -0
  34. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/__init__.py +0 -0
  35. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bbci.py +0 -0
  36. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bcicomp.py +0 -0
  37. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/__init__.py +0 -0
  38. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/datasets.py +0 -0
  39. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/format.py +0 -0
  40. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_format.py +0 -0
  41. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/hub_validation.py +0 -0
  42. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/bids/iterable.py +0 -0
  43. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/chb_mit.py +0 -0
  44. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/collate.py +0 -0
  45. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/mne.py +0 -0
  46. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/moabb.py +0 -0
  47. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/nmt.py +0 -0
  48. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/registry.py +0 -0
  49. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/siena.py +0 -0
  50. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/sleep_physio_challe_18.py +0 -0
  51. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/tuh.py +0 -0
  52. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/utils.py +0 -0
  53. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datasets/xy.py +0 -0
  54. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datautil/__init__.py +0 -0
  55. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datautil/channel_utils.py +0 -0
  56. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datautil/hub_formats.py +0 -0
  57. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datautil/serialization.py +0 -0
  58. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/datautil/util.py +0 -0
  59. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/eegneuralnet.py +0 -0
  60. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/functional/__init__.py +0 -0
  61. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/functional/functions.py +0 -0
  62. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/functional/initialization.py +0 -0
  63. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/__init__.py +0 -0
  64. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/attentionbasenet.py +0 -0
  65. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/bendr.py +0 -0
  66. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/biot.py +0 -0
  67. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/brainmodule.py +0 -0
  68. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/cbramod.py +0 -0
  69. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/codebrain.py +0 -0
  70. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/config.py +0 -0
  71. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/contrawr.py +0 -0
  72. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/dance.py +0 -0
  73. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/deep4.py +0 -0
  74. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/deepsleepnet.py +0 -0
  75. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/dgcnn.py +0 -0
  76. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegconformer.py +0 -0
  77. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegdino.py +0 -0
  78. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eeginception_erp.py +0 -0
  79. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eeginception_mi.py +0 -0
  80. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegitnet.py +0 -0
  81. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegminer.py +0 -0
  82. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegnet.py +0 -0
  83. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegnex.py +0 -0
  84. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegpt.py +0 -0
  85. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegsym.py +0 -0
  86. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/eegtcnet.py +0 -0
  87. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/emg2qwerty.py +0 -0
  88. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/fbcnet.py +0 -0
  89. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/fblightconvnet.py +0 -0
  90. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/fbmsnet.py +0 -0
  91. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/hybrid.py +0 -0
  92. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/interpolated.py +0 -0
  93. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/labram.py +0 -0
  94. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/luna.py +0 -0
  95. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/medformer.py +0 -0
  96. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/meta_neuromotor.py +0 -0
  97. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/msvtnet.py +0 -0
  98. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/mvpformer.py +0 -0
  99. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/patchedtransformer.py +0 -0
  100. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/reve.py +0 -0
  101. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sccnet.py +0 -0
  102. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/shallow_fbcsp.py +0 -0
  103. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/signal_jepa.py +0 -0
  104. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sinc_shallow.py +0 -0
  105. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/sstdpn.py +0 -0
  106. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/steegformer.py +0 -0
  107. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/summary.csv +0 -0
  108. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/syncnet.py +0 -0
  109. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/tcformer.py +0 -0
  110. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/tcn.py +0 -0
  111. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/tsinception.py +0 -0
  112. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/usleep.py +0 -0
  113. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/util.py +0 -0
  114. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/models/zuna.py +0 -0
  115. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/__init__.py +0 -0
  116. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/activation.py +0 -0
  117. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/attention.py +0 -0
  118. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/blocks.py +0 -0
  119. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/convolution.py +0 -0
  120. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/dance_modules.py +0 -0
  121. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/filter.py +0 -0
  122. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/interpolation.py +0 -0
  123. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/layers.py +0 -0
  124. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/linear.py +0 -0
  125. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/parametrization.py +0 -0
  126. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/stats.py +0 -0
  127. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/util.py +0 -0
  128. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/modules/wrapper.py +0 -0
  129. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/__init__.py +0 -0
  130. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/eegprep_preprocess.py +0 -0
  131. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/mne_preprocess.py +0 -0
  132. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/preprocess.py +0 -0
  133. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/preprocessing/util.py +0 -0
  134. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/regressor.py +0 -0
  135. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/samplers/__init__.py +0 -0
  136. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/samplers/ssl.py +0 -0
  137. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/training/__init__.py +0 -0
  138. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/training/callbacks.py +0 -0
  139. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/training/scoring.py +0 -0
  140. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/__init__.py +0 -0
  141. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/attribution.py +0 -0
  142. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/confusion_matrices.py +0 -0
  143. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/frequency.py +0 -0
  144. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/metrics.py +0 -0
  145. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/sanity.py +0 -0
  146. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode/visualization/topology.py +0 -0
  147. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/SOURCES.txt +0 -0
  148. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/dependency_links.txt +0 -0
  149. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/requires.txt +0 -0
  150. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/braindecode.egg-info/top_level.txt +0 -0
  151. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/Makefile +0 -0
  152. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/class.rst +0 -0
  153. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/class_in_subdir.rst +0 -0
  154. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/function.rst +0 -0
  155. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/_templates/autosummary/function_in_subdir.rst +0 -0
  156. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/api.rst +0 -0
  157. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/cite.rst +0 -0
  158. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/conf.py +0 -0
  159. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/help.rst +0 -0
  160. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/index.rst +0 -0
  161. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/install/install.rst +0 -0
  162. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/install/install_pip.rst +0 -0
  163. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/install/install_source.rst +0 -0
  164. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/attention.rst +0 -0
  165. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/channel.rst +0 -0
  166. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/convolution.rst +0 -0
  167. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/filterbank.rst +0 -0
  168. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/gnn.rst +0 -0
  169. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/interpretable.rst +0 -0
  170. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/lbm.rst +0 -0
  171. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/recurrent.rst +0 -0
  172. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/categorization/spd.rst +0 -0
  173. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/models.rst +0 -0
  174. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/models_categorization.rst +0 -0
  175. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/models_table.rst +0 -0
  176. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/models/models_visualization.rst +0 -0
  177. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/docs/sg_execution_times.rst +0 -0
  178. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/pyproject.toml +0 -0
  179. {braindecode-1.8.0.dev174836782 → braindecode-1.8.0.dev180935432}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: braindecode
3
- Version: 1.8.0.dev174836782
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>
@@ -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 = torch.arange(batch_size).to(device)
1069
- y_b = torch.arange(batch_size).to(device)
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(
@@ -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:
@@ -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
- # Keep reference to first dataset for preprocessing kwargs
695
- first_ds = self.datasets[0]
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 (check first dataset, assuming uniform preprocessing)
710
- # These are typically set by windowing functions on individual datasets
711
- for kwarg_name in [
712
- "raw_preproc_kwargs",
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 = ds.windows.info.to_json_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 = ds.raw.info.to_json_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 = ds.raw.info.to_json_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 = json.loads(root.attrs[kwarg_name])
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 legacy zarr stores.
81
+ """Restore NaN values from None in JSON-loaded attributes.
22
82
 
23
- Datasets saved before zarr v3 native NaN support used
24
- ``_sanitize_for_json`` to convert NaN/Inf → None. This restores them
25
- on load so ``mne.Info.from_json_dict`` gets proper NaN arrays.
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(isinstance(x, (int, float, type(None))) for x in obj):
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
- tmin = annots[int(sleep_event_inds[0])]["onset"] - crop_wake_mins * 60
121
- tmax = annots[int(sleep_event_inds[-1])]["onset"] + crop_wake_mins * 60
122
- raw.crop(tmin=max(tmin, raw.times[0]), tmax=min(tmax, raw.times[-1]))
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}
@@ -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,