-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdata.py
More file actions
150 lines (122 loc) · 4.39 KB
/
Copy pathdata.py
File metadata and controls
150 lines (122 loc) · 4.39 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
from pathlib import Path
from sklearn.preprocessing import LabelEncoder
import mne
from moabb.utils import set_log_level
from moabb.paradigms import MotorImagery
from skorch.dataset import Dataset
from moabb.datasets import BNCI2014_001, BNCI2014_004, Lee2019_MI, PhysionetMI, Schirrmeister2017
import logging
logger = logging.getLogger(__name__)
# Preprocessing steps in Xie2023:
# 1. Pick channels:
channels = ['C3', 'Cz', 'C4']
# 2. Re-reference using left mastoid
# ref = 'M1'
ref = None
# ref = 'average'
# -> canceled because not all datasets have this channel.
# -> instead, use average reference
# 3. Resampling
resample = 250
# 4. Epoching
tmin = -0.5
tmax = 4
# 5. Bandpass filtering
l_freq = 0
h_freq = 40
# 6. Channel-wise standardization
def _set_metadata(epochs, metadata, target=None, labels=None):
if labels is not None:
metadata['labels'] = labels
if target is not None:
metadata['target'] = target
epochs.metadata = metadata
def _epochs_to_dataset(epochs):
if epochs is None:
return None
X = epochs.get_data(units='uV').astype('float32')
# Channel-wise standardization:
# Slightly different from Xie2023, which used exponential moving channel-wise standardization,
# but they initialized its parameters on the first 4 seconds of each 4.5s epoch.
# So we are nearly equivalent.
mu = X.mean(axis=2, keepdims=True)
sigma = X.std(axis=2, keepdims=True)
X = (X - mu) / sigma
y = epochs.metadata['target'].values
return Dataset(X, y)
def preprocess_data(dataset) -> mne.Epochs: # tuple[mne.Epochs, mne.Epochs]:
paradigm = MotorImagery(channels=channels, resample=resample, tmin=tmin, tmax=tmax)
epochs, labels, metadata = paradigm.get_data(dataset, return_epochs=True)
if ref is not None:
epochs, _ = mne.set_eeg_reference(epochs, ref_channels=ref, copy=False)
# third-order Butterworth bandpass filter of 0–40 Hz applied on epochs as in Xie2023
epochs.filter(l_freq=l_freq, h_freq=h_freq, method="iir",
iir_params=dict(order=3, ftype='butter'))
le = LabelEncoder()
id_labels = le.fit_transform(labels)
_set_metadata(epochs, metadata, target=id_labels, labels=labels)
return epochs
def get_data(dataset, subjects: list[int] = None,
overwrite_data: bool = False,
data_dir=None, return_metadata=False):
set_log_level('info')
dataset_name = dataset.__class__.__name__
if data_dir is None:
data_dir = Path('~/') / 'data'
preprocessed_data_dir = Path(data_dir).expanduser() / 'preprocessed' / 'xie2023'
path = preprocessed_data_dir / f'{dataset_name}{"" if ref == "average" else "_no-ref"}-epo.fif'
if path.exists() and not overwrite_data:
logger.info(f'Loading pre-processed data from {path}')
epochs = mne.read_epochs(path, preload=False)
else:
logger.info('Pre-processing data')
epochs = preprocess_data(dataset)
logger.info(f'Saving pre-processed data to {path}')
preprocessed_data_dir.mkdir(parents=True, exist_ok=True)
assert preprocessed_data_dir.exists()
epochs.save(path, overwrite=True)
if subjects is not None:
epochs = epochs[epochs.metadata.subject.isin(subjects)]
Xy = _epochs_to_dataset(epochs)
if return_metadata:
return Xy, epochs.metadata
return Xy
pretrain_datasets = [
BNCI2014_001(),
BNCI2014_004(),
Lee2019_MI(),
PhysionetMI(imagined=True, executed=True),
Schirrmeister2017(),
]
finetune_datasets = [
BNCI2014_001(),
BNCI2014_004(),
Schirrmeister2017(),
]
class TestData:
def test_save_data(self, tmp_path):
subjects = [1]
dataset = pretrain_datasets[0]
# preprocess and save data:
out = get_data(
dataset,
subjects=subjects,
overwrite_data=False, # should still create the data
data_dir=tmp_path,
)
# load data:
out1 = get_data(
dataset,
subjects=subjects,
overwrite_data=False,
data_dir=tmp_path,
)
train_set = out
print(train_set.X.std())
for X, X1 in zip([out], [out1]):
assert X.X.shape == X1.X.shape
assert len(X.y) == len(X1.y)
assert X.X.dtype == X1.X.dtype
if __name__ == '__main__':
for dataset in pretrain_datasets:
_ = get_data(dataset, overwrite_data=False)