eegproc 1.0.0__py3-none-any.whl
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.
- eegproc/__init__.py +7 -0
- eegproc/featurization.py +546 -0
- eegproc/preprocessing.py +125 -0
- eegproc-1.0.0.dist-info/METADATA +441 -0
- eegproc-1.0.0.dist-info/RECORD +8 -0
- eegproc-1.0.0.dist-info/WHEEL +5 -0
- eegproc-1.0.0.dist-info/licenses/LICENSE +339 -0
- eegproc-1.0.0.dist-info/top_level.txt +1 -0
eegproc/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
from .preprocessing import bandpass_filter, apply_detrend, FREQUENCY_BANDS
|
|
2
|
+
from .featurization import hjorth_params, psd_bandpowers, wavelet_band_energy, wavelet_entropy, shannons_entropy, imf_band_energy, imf_entropy
|
|
3
|
+
|
|
4
|
+
__all__ = [
|
|
5
|
+
"bandpass_filter", "apply_detrend", "FREQUENCY_BANDS",
|
|
6
|
+
"hjorth_params", "psd_bandpowers", "shannons_entropy", "wavelet_band_energy", "wavelet_entropy", "imf_band_energy", "imf_entropy"
|
|
7
|
+
]
|
eegproc/featurization.py
ADDED
|
@@ -0,0 +1,546 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pandas as pd
|
|
3
|
+
import pywt
|
|
4
|
+
from math import log2, floor
|
|
5
|
+
from scipy.signal import welch
|
|
6
|
+
from eegproc import bandpass_filter, apply_detrend, FREQUENCY_BANDS
|
|
7
|
+
from PyEMD import EMD
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
'''SPECTRAL ENTROPY'''
|
|
11
|
+
def psd_bandpowers(
|
|
12
|
+
df: pd.DataFrame,
|
|
13
|
+
fs: float,
|
|
14
|
+
bands: dict[str, tuple[float, float]] = FREQUENCY_BANDS,
|
|
15
|
+
window_sec: float = 4.0,
|
|
16
|
+
overlap: float = 0.5,
|
|
17
|
+
detrend: str | None = "constant",
|
|
18
|
+
) -> pd.DataFrame:
|
|
19
|
+
df = apply_detrend(detrend, df)
|
|
20
|
+
|
|
21
|
+
band_keys = set(bands.keys())
|
|
22
|
+
col_band, col_chan = {}, {}
|
|
23
|
+
for col in df.columns:
|
|
24
|
+
parts = col.rsplit("_", 1)
|
|
25
|
+
if len(parts) == 2 and parts[1] in band_keys:
|
|
26
|
+
col_band[col] = parts[1]
|
|
27
|
+
col_chan[col] = parts[0]
|
|
28
|
+
if not col_band:
|
|
29
|
+
raise ValueError("No columns named like '{channel}_{band}' with band in FREQUENCY_BANDS.")
|
|
30
|
+
df = df[list(col_band.keys())]
|
|
31
|
+
|
|
32
|
+
data = df.to_numpy(dtype=float, copy=False)
|
|
33
|
+
n_samples, n_cols = data.shape
|
|
34
|
+
nperseg = int(round(window_sec * fs))
|
|
35
|
+
if nperseg <= 8:
|
|
36
|
+
raise ValueError("window_sec too small for given fs; increase window_sec.")
|
|
37
|
+
if not (0.0 <= overlap < 1.0):
|
|
38
|
+
raise ValueError("overlap must be in [0.0, 1.0).")
|
|
39
|
+
hop = int(round(nperseg * (1.0 - overlap)))
|
|
40
|
+
if hop <= 0:
|
|
41
|
+
raise ValueError("overlap too large; hop size must be >= 1 sample.")
|
|
42
|
+
if nperseg > n_samples:
|
|
43
|
+
return pd.DataFrame(columns=list(df.columns))
|
|
44
|
+
|
|
45
|
+
band_to_idx = {}
|
|
46
|
+
for i, col in enumerate(df.columns):
|
|
47
|
+
band_to_idx.setdefault(col_band[col], []).append(i)
|
|
48
|
+
|
|
49
|
+
rows = []
|
|
50
|
+
for start in range(0, n_samples - nperseg + 1, hop):
|
|
51
|
+
seg = data[start:start + nperseg, :]
|
|
52
|
+
|
|
53
|
+
f, psd = welch(
|
|
54
|
+
seg,
|
|
55
|
+
fs=fs,
|
|
56
|
+
window="hann",
|
|
57
|
+
nperseg=nperseg,
|
|
58
|
+
noverlap=0,
|
|
59
|
+
detrend=False,
|
|
60
|
+
scaling="density",
|
|
61
|
+
return_onesided=True,
|
|
62
|
+
axis=0,
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
row = {}
|
|
66
|
+
for band, idxs in band_to_idx.items():
|
|
67
|
+
lo, hi = bands[band]
|
|
68
|
+
m = (f >= lo) & (f <= hi)
|
|
69
|
+
if not m.any():
|
|
70
|
+
for j in idxs:
|
|
71
|
+
row[df.columns[j]] = 0.0
|
|
72
|
+
continue
|
|
73
|
+
|
|
74
|
+
band_power = np.trapezoid(psd[m][:, idxs], f[m], axis=0)
|
|
75
|
+
for k, j in enumerate(idxs):
|
|
76
|
+
row[df.columns[j]] = float(band_power[k])
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
rows.append(row)
|
|
80
|
+
|
|
81
|
+
return pd.DataFrame(rows, columns=list(df.columns))
|
|
82
|
+
|
|
83
|
+
def shannons_entropy(
|
|
84
|
+
df: pd.DataFrame,
|
|
85
|
+
fs: float,
|
|
86
|
+
bands: dict[str, tuple[float, float]] = FREQUENCY_BANDS,
|
|
87
|
+
window_sec: float = 4.0,
|
|
88
|
+
overlap: float = 0.5,
|
|
89
|
+
eps: float = 1e-300, # avoids 0 denominators and log(0)
|
|
90
|
+
detrend: str | None = "constant",
|
|
91
|
+
) -> pd.DataFrame:
|
|
92
|
+
df = apply_detrend(detrend, df)
|
|
93
|
+
band_keys = set(bands.keys())
|
|
94
|
+
col_band = {}
|
|
95
|
+
for col in df.columns:
|
|
96
|
+
parts = col.rsplit("_", 1)
|
|
97
|
+
if len(parts) == 2 and parts[1] in band_keys:
|
|
98
|
+
col_band[col] = parts[1]
|
|
99
|
+
if not col_band:
|
|
100
|
+
raise ValueError("No columns named like '{channel}_{band}' with band in FREQUENCY_BANDS.")
|
|
101
|
+
df = df[list(col_band.keys())]
|
|
102
|
+
|
|
103
|
+
data = df.to_numpy(dtype=float, copy=False)
|
|
104
|
+
n_samples, n_cols = data.shape
|
|
105
|
+
nperseg = int(round(window_sec * fs))
|
|
106
|
+
if nperseg <= 8:
|
|
107
|
+
raise ValueError("window_sec too small for given fs; increase window_sec.")
|
|
108
|
+
if not (0.0 <= overlap < 1.0):
|
|
109
|
+
raise ValueError("overlap must be in [0.0, 1.0).")
|
|
110
|
+
hop = int(round(nperseg * (1.0 - overlap)))
|
|
111
|
+
if hop <= 0:
|
|
112
|
+
raise ValueError("overlap too large; hop size must be >= 1 sample.")
|
|
113
|
+
if nperseg > n_samples:
|
|
114
|
+
return pd.DataFrame(columns=[f"{c}_entropy" for c in df.columns])
|
|
115
|
+
|
|
116
|
+
band_to_idx = {}
|
|
117
|
+
for i, col in enumerate(df.columns):
|
|
118
|
+
band_to_idx.setdefault(col_band[col], []).append(i)
|
|
119
|
+
|
|
120
|
+
rows = []
|
|
121
|
+
for start in range(0, n_samples - nperseg + 1, hop):
|
|
122
|
+
seg = data[start:start + nperseg, :]
|
|
123
|
+
|
|
124
|
+
f, psd = welch(
|
|
125
|
+
seg,
|
|
126
|
+
fs=fs,
|
|
127
|
+
window="hann",
|
|
128
|
+
nperseg=nperseg,
|
|
129
|
+
noverlap=0,
|
|
130
|
+
detrend=False,
|
|
131
|
+
scaling="density",
|
|
132
|
+
return_onesided=True,
|
|
133
|
+
axis=0,
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
row = {}
|
|
137
|
+
for band, idxs in band_to_idx.items():
|
|
138
|
+
lo, hi = bands[band]
|
|
139
|
+
m = (f >= lo) & (f <= hi)
|
|
140
|
+
count = int(np.count_nonzero(m))
|
|
141
|
+
if count < 2:
|
|
142
|
+
for j in idxs:
|
|
143
|
+
row[f"{df.columns[j]}_entropy"] = np.nan
|
|
144
|
+
continue
|
|
145
|
+
|
|
146
|
+
band_power = psd[m][:, idxs]
|
|
147
|
+
totals = np.sum(band_power, axis=0)
|
|
148
|
+
valid = np.isfinite(totals) & (totals > 0)
|
|
149
|
+
|
|
150
|
+
p = np.empty_like(band_power)
|
|
151
|
+
p[:, valid] = band_power[:, valid] / totals[valid]
|
|
152
|
+
p[:, ~valid] = np.nan
|
|
153
|
+
p = np.clip(p, eps, 1.0)
|
|
154
|
+
|
|
155
|
+
H = -np.nansum(p * np.log2(p), axis=0)
|
|
156
|
+
H /= np.log2(count)
|
|
157
|
+
|
|
158
|
+
for k, j in enumerate(idxs):
|
|
159
|
+
row[f"{df.columns[j]}_entropy"] = float(H[k]) if np.isfinite(H[k]) else np.nan
|
|
160
|
+
|
|
161
|
+
rows.append(row)
|
|
162
|
+
|
|
163
|
+
return pd.DataFrame(rows, columns=[f"{c}_entropy" for c in df.columns])
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
'''HJORTH PARAMETRIZATION'''
|
|
167
|
+
def hjorth_params(
|
|
168
|
+
df: pd.DataFrame,
|
|
169
|
+
fs: float,
|
|
170
|
+
window_sec: float = 4.0,
|
|
171
|
+
overlap: float = 0.5,
|
|
172
|
+
detrend: str | None = "constant",
|
|
173
|
+
eps: float = 1e-300,
|
|
174
|
+
) -> pd.DataFrame:
|
|
175
|
+
"""
|
|
176
|
+
Compute Hjorth parameters per window for each numeric column in a band-passed EEG DataFrame.
|
|
177
|
+
|
|
178
|
+
Returns a DataFrame with multiple rows (one per window):
|
|
179
|
+
<col>_activity, <col>_mobility, <col>_complexity
|
|
180
|
+
|
|
181
|
+
Parameters
|
|
182
|
+
df : DataFrame (samples df channels/bands)
|
|
183
|
+
fs : float sampling rate in Hz
|
|
184
|
+
window_sec : float window length in seconds
|
|
185
|
+
step_sec : float step between windows in seconds; defaults to window_sec (no overlap)
|
|
186
|
+
detrend : {"constant", "linear" ,None}
|
|
187
|
+
eps : float numerical guard
|
|
188
|
+
"""
|
|
189
|
+
df = apply_detrend(detrend, df)
|
|
190
|
+
|
|
191
|
+
cols = list(df.columns)
|
|
192
|
+
data = df.to_numpy(dtype=float)
|
|
193
|
+
n_samples, n_cols = data.shape
|
|
194
|
+
|
|
195
|
+
win = int(round(window_sec * fs))
|
|
196
|
+
if win < 3:
|
|
197
|
+
raise ValueError("window_sec too small (need >= 3 samples for second differences).")
|
|
198
|
+
if not (0.0 <= overlap < 1.0):
|
|
199
|
+
raise ValueError("overlap must be in [0.0, 1.0).")
|
|
200
|
+
|
|
201
|
+
hop = int(round(win * (1.0 - overlap)))
|
|
202
|
+
if hop <= 0:
|
|
203
|
+
raise ValueError("overlap too large; hop size must be >= 1 sample.")
|
|
204
|
+
|
|
205
|
+
rows = []
|
|
206
|
+
starts = range(0, n_samples - win + 1, hop)
|
|
207
|
+
for i0 in starts:
|
|
208
|
+
i1 = i0 + win
|
|
209
|
+
seg = data[i0:i1, :]
|
|
210
|
+
if seg.shape[0] < 3:
|
|
211
|
+
continue
|
|
212
|
+
|
|
213
|
+
act = np.nanvar(seg, axis=0, ddof=0)
|
|
214
|
+
|
|
215
|
+
dx = np.diff(seg, n=1, axis=0)
|
|
216
|
+
ddx = np.diff(seg, n=2, axis=0)
|
|
217
|
+
|
|
218
|
+
var_dx = np.nanvar(dx, axis=0, ddof=0)
|
|
219
|
+
var_ddx = np.nanvar(ddx, axis=0, ddof=0)
|
|
220
|
+
|
|
221
|
+
mob = np.sqrt((var_dx + eps) / (act + eps))
|
|
222
|
+
mob_dx = np.sqrt((var_ddx + eps) / (var_dx + eps))
|
|
223
|
+
comp = mob_dx / (mob + eps)
|
|
224
|
+
|
|
225
|
+
row = {}
|
|
226
|
+
|
|
227
|
+
for k, c in enumerate(cols):
|
|
228
|
+
row[f"{c}_activity"] = float(act[k]) if np.isfinite(act[k]) else np.nan
|
|
229
|
+
row[f"{c}_mobility"] = float(mob[k]) if np.isfinite(mob[k]) else np.nan
|
|
230
|
+
row[f"{c}_complexity"] = float(comp[k]) if np.isfinite(comp[k]) else np.nan
|
|
231
|
+
rows.append(row)
|
|
232
|
+
|
|
233
|
+
return pd.DataFrame(rows)
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
'''WAVELET FEATURES'''
|
|
237
|
+
def _choose_dwt_level(n_samples: int, fs: float, wavelet: str, min_freq: float) -> int:
|
|
238
|
+
max_lvl = pywt.dwt_max_level(n_samples, pywt.Wavelet(wavelet).dec_len)
|
|
239
|
+
target = max(1, floor(log2(fs / max(min_freq, 1e-6)) - 1))
|
|
240
|
+
return max(1, min(max_lvl, target))
|
|
241
|
+
|
|
242
|
+
def _dwt_subband_ranges(fs: float, level: int) -> dict[str, tuple[float, float]]:
|
|
243
|
+
bands: dict[str, tuple[float, float]] = {}
|
|
244
|
+
for j in range(1, level + 1):
|
|
245
|
+
f_hi = fs / (2 ** j)
|
|
246
|
+
f_lo = fs / (2 ** (j + 1))
|
|
247
|
+
bands[f"D{j}"] = (f_lo, f_hi)
|
|
248
|
+
bands[f"A{level}"] = (0.0, fs / (2 ** (level + 1)))
|
|
249
|
+
return bands
|
|
250
|
+
|
|
251
|
+
def _overlap(a: tuple[float, float], b: tuple[float, float]) -> float:
|
|
252
|
+
lo = max(a[0], b[0]); hi = min(a[1], b[1])
|
|
253
|
+
return max(0.0, hi - lo)
|
|
254
|
+
|
|
255
|
+
def wavelet_band_energy(
|
|
256
|
+
df: pd.DataFrame,
|
|
257
|
+
fs: float,
|
|
258
|
+
bands: dict[str, tuple[float, float]],
|
|
259
|
+
wavelet: str = "db4",
|
|
260
|
+
mode: str = "periodization",
|
|
261
|
+
window_sec: float = 4.0,
|
|
262
|
+
overlap: float = 0.5,
|
|
263
|
+
) -> pd.DataFrame:
|
|
264
|
+
df = df.select_dtypes(include=[np.number])
|
|
265
|
+
|
|
266
|
+
n_samples = len(df)
|
|
267
|
+
nperseg = int(round(window_sec * fs))
|
|
268
|
+
if nperseg <= 8:
|
|
269
|
+
raise ValueError("window_sec too small for given fs; increase window_sec.")
|
|
270
|
+
if not (0.0 <= overlap < 1.0):
|
|
271
|
+
raise ValueError("overlap must be in [0.0, 1.0).")
|
|
272
|
+
hop = int(round(nperseg * (1.0 - overlap)))
|
|
273
|
+
if hop <= 0:
|
|
274
|
+
raise ValueError("overlap too large; hop size must be >= 1 sample.")
|
|
275
|
+
if nperseg > n_samples:
|
|
276
|
+
return pd.DataFrame(columns=[f"{ch}_{b}_wenergy" for ch in df.columns for b in bands])
|
|
277
|
+
|
|
278
|
+
min_band_lo = min(lo for lo, _ in bands.values())
|
|
279
|
+
L = _choose_dwt_level(n_samples=nperseg, fs=fs, wavelet=wavelet, min_freq=min_band_lo)
|
|
280
|
+
sub_ranges = _dwt_subband_ranges(fs, L)
|
|
281
|
+
|
|
282
|
+
cols = [f"{ch}_{b}_wenergy" for ch in df.columns for b in bands]
|
|
283
|
+
rows = []
|
|
284
|
+
|
|
285
|
+
for start in range(0, n_samples - nperseg + 1, hop):
|
|
286
|
+
win = df.iloc[start:start + nperseg]
|
|
287
|
+
row: dict[str, float] = {}
|
|
288
|
+
|
|
289
|
+
for ch in df.columns:
|
|
290
|
+
y = win[ch].to_numpy(dtype=float, copy=False)
|
|
291
|
+
|
|
292
|
+
coeffs = pywt.wavedec(y, wavelet=wavelet, level=L, mode=mode) # applies wavelet transform function
|
|
293
|
+
wv_coeff_approx = coeffs[0]
|
|
294
|
+
wv_coeff_details = coeffs[1:]
|
|
295
|
+
|
|
296
|
+
sub_eng: dict[str, float] = {}
|
|
297
|
+
for idx, c in enumerate(wv_coeff_details):
|
|
298
|
+
j = L - idx
|
|
299
|
+
sub_eng[f"D{j}"] = float(np.sum(c.astype(float) ** 2))
|
|
300
|
+
sub_eng[f"A{L}"] = float(np.sum(wv_coeff_approx.astype(float) ** 2))
|
|
301
|
+
|
|
302
|
+
band_energy = {name: 0.0 for name in bands}
|
|
303
|
+
for sub_name, e_sub in sub_eng.items():
|
|
304
|
+
f_lo, f_hi = sub_ranges[sub_name]
|
|
305
|
+
width = (f_hi - f_lo) or 1.0
|
|
306
|
+
if width <= 0:
|
|
307
|
+
continue
|
|
308
|
+
for band_name, (blo, bhi) in bands.items():
|
|
309
|
+
olap = _overlap((f_lo, f_hi), (blo, bhi))
|
|
310
|
+
if olap > 0:
|
|
311
|
+
band_energy[band_name] += e_sub * (olap / width)
|
|
312
|
+
|
|
313
|
+
for band_name, e in band_energy.items():
|
|
314
|
+
row[f"{ch}_{band_name}_wenergy"] = float(e)
|
|
315
|
+
|
|
316
|
+
rows.append(row)
|
|
317
|
+
|
|
318
|
+
return pd.DataFrame(rows, columns=cols)
|
|
319
|
+
|
|
320
|
+
def wavelet_entropy(
|
|
321
|
+
wv_band_energy_df: pd.DataFrame,
|
|
322
|
+
bands: dict[str, tuple[float, float]],
|
|
323
|
+
normalize: bool = True,
|
|
324
|
+
eps: float = 1e-300,
|
|
325
|
+
) -> pd.DataFrame:
|
|
326
|
+
df = wv_band_energy_df.select_dtypes(include=[np.number]).copy()
|
|
327
|
+
|
|
328
|
+
band_list = list(bands.keys())
|
|
329
|
+
K = len(band_list)
|
|
330
|
+
norm = (np.log(K) if (normalize and K > 1) else 1.0)
|
|
331
|
+
|
|
332
|
+
channel_to_cols: dict[str, list[str]] = {}
|
|
333
|
+
for col in df.columns:
|
|
334
|
+
if not col.endswith("_wenergy"):
|
|
335
|
+
continue
|
|
336
|
+
core = col[:-8]
|
|
337
|
+
if "_" not in core:
|
|
338
|
+
continue
|
|
339
|
+
ch, b = core.rsplit("_", 1)
|
|
340
|
+
if b in bands:
|
|
341
|
+
channel_to_cols.setdefault(ch, [None] * K)
|
|
342
|
+
|
|
343
|
+
if not channel_to_cols:
|
|
344
|
+
raise ValueError("No columns with pattern '{channel}_{band}_wenergy' matching provided bands.")
|
|
345
|
+
|
|
346
|
+
for col in df.columns:
|
|
347
|
+
if not col.endswith("_wenergy"):
|
|
348
|
+
continue
|
|
349
|
+
core = col[:-8]
|
|
350
|
+
if "_" not in core:
|
|
351
|
+
continue
|
|
352
|
+
ch, b = core.rsplit("_", 1)
|
|
353
|
+
if ch in channel_to_cols and b in bands:
|
|
354
|
+
idx = band_list.index(b)
|
|
355
|
+
channel_to_cols[ch][idx] = col
|
|
356
|
+
|
|
357
|
+
out_cols = [f"{ch}_wentropy" for ch in channel_to_cols.keys()]
|
|
358
|
+
rows = []
|
|
359
|
+
|
|
360
|
+
for i in range(len(df)):
|
|
361
|
+
row_out = {}
|
|
362
|
+
for ch, cols_in_order in channel_to_cols.items():
|
|
363
|
+
vals = []
|
|
364
|
+
for c in cols_in_order:
|
|
365
|
+
if c is None:
|
|
366
|
+
vals.append(0.0)
|
|
367
|
+
else:
|
|
368
|
+
v = df.iat[i, df.columns.get_loc(c)]
|
|
369
|
+
vals.append(float(v) if np.isfinite(v) else 0.0)
|
|
370
|
+
|
|
371
|
+
total = float(np.nansum(vals))
|
|
372
|
+
total = total if (np.isfinite(total) and total > 0) else eps
|
|
373
|
+
|
|
374
|
+
p = np.asarray(vals, dtype=float) / total
|
|
375
|
+
p = np.clip(p, eps, 1.0)
|
|
376
|
+
p /= p.sum()
|
|
377
|
+
|
|
378
|
+
H = -np.sum(p * np.log(p))
|
|
379
|
+
row_out[f"{ch}_wentropy"] = float(H / (norm or 1.0))
|
|
380
|
+
rows.append(row_out)
|
|
381
|
+
|
|
382
|
+
return pd.DataFrame(rows, columns=out_cols)
|
|
383
|
+
|
|
384
|
+
|
|
385
|
+
'''IMF FEATURES'''
|
|
386
|
+
def imf_band_energy(
|
|
387
|
+
df: pd.DataFrame,
|
|
388
|
+
fs: float,
|
|
389
|
+
imf_to_band: list[str] = ["delta", "theta", "alpha", "betaL", "betaH", "gamma"],
|
|
390
|
+
window_sec: float = 4.0,
|
|
391
|
+
overlap: float = 0.5,
|
|
392
|
+
EMD_kwargs: dict = {},
|
|
393
|
+
) -> pd.DataFrame:
|
|
394
|
+
|
|
395
|
+
df = df.select_dtypes(include=[np.number])
|
|
396
|
+
|
|
397
|
+
n_samples = len(df)
|
|
398
|
+
nperseg = int(round(window_sec * fs))
|
|
399
|
+
|
|
400
|
+
if nperseg <= 8:
|
|
401
|
+
raise ValueError("window_sec too small for given fs; increase window_sec.")
|
|
402
|
+
if not (0.0 <= overlap < 1.0):
|
|
403
|
+
raise ValueError("overlap must be in [0.0, 1.0).")
|
|
404
|
+
hop = int(round(nperseg * (1.0 - overlap)))
|
|
405
|
+
if hop <= 0:
|
|
406
|
+
raise ValueError("overlap too large; hop size must be >= 1 sample.")
|
|
407
|
+
if nperseg > n_samples:
|
|
408
|
+
cols = [f"{ch}_{band}_imfenergy" for ch in df.columns for band in imf_to_band]
|
|
409
|
+
return pd.DataFrame(columns=cols)
|
|
410
|
+
|
|
411
|
+
max_imf_needed = int(len(imf_to_band))
|
|
412
|
+
cols = [f"{ch}_{band}_imfenergy" for ch in df.columns for band in imf_to_band]
|
|
413
|
+
|
|
414
|
+
emd = EMD(**(EMD_kwargs))
|
|
415
|
+
|
|
416
|
+
rows = []
|
|
417
|
+
row: dict[str, float] = {}
|
|
418
|
+
|
|
419
|
+
emd._imf_cumsums = {}
|
|
420
|
+
|
|
421
|
+
for ch in df.columns:
|
|
422
|
+
y_full = df[ch].to_numpy(dtype=float, copy=False).astype(float, copy=False)
|
|
423
|
+
imfs_full = emd.emd(y_full, max_imf=max_imf_needed)
|
|
424
|
+
sq = imfs_full ** 2
|
|
425
|
+
cumsq = np.hstack([np.zeros((sq.shape[0], 1), dtype=sq.dtype),
|
|
426
|
+
np.cumsum(sq, axis=1)])
|
|
427
|
+
|
|
428
|
+
emd._imf_cumsums[ch] = (imfs_full, cumsq)
|
|
429
|
+
|
|
430
|
+
|
|
431
|
+
for start in range(0, n_samples - nperseg + 1, hop):
|
|
432
|
+
end = start + nperseg
|
|
433
|
+
row: dict[str, float] = {}
|
|
434
|
+
|
|
435
|
+
for ch in df.columns:
|
|
436
|
+
e_win = emd._imf_cumsums[ch][1][:, end] - emd._imf_cumsums[ch][1][:, start]
|
|
437
|
+
|
|
438
|
+
for imf_idx in range(len(imf_to_band)):
|
|
439
|
+
e = float(e_win[imf_idx]) if imf_idx < e_win.shape[0] else 0.0
|
|
440
|
+
row[f"{ch}_{imf_to_band[imf_idx]}_imfenergy"] = e
|
|
441
|
+
|
|
442
|
+
rows.append(row)
|
|
443
|
+
|
|
444
|
+
|
|
445
|
+
return pd.DataFrame(rows, columns=cols)
|
|
446
|
+
|
|
447
|
+
def imf_entropy(
|
|
448
|
+
imf_energy_df: pd.DataFrame,
|
|
449
|
+
bands: list[str] = ["delta", "theta", "alpha", "betaL", "betaH", "gamma"],
|
|
450
|
+
normalize: bool = True,
|
|
451
|
+
eps: float = 1e-300,
|
|
452
|
+
) -> pd.DataFrame:
|
|
453
|
+
df = imf_energy_df.select_dtypes(include=[np.number])
|
|
454
|
+
|
|
455
|
+
k = len(bands)
|
|
456
|
+
norm = (np.log(k) if (normalize and k > 1) else 1.0)
|
|
457
|
+
|
|
458
|
+
channel_to_cols: dict[str, list[str]] = {}
|
|
459
|
+
suffix = "_imfenergy"
|
|
460
|
+
|
|
461
|
+
for col in df.columns:
|
|
462
|
+
if not col.endswith(suffix):
|
|
463
|
+
continue
|
|
464
|
+
ch_band = col[:-len(suffix)]
|
|
465
|
+
if "_" not in ch_band:
|
|
466
|
+
continue
|
|
467
|
+
ch, band = ch_band.rsplit("_", 1)
|
|
468
|
+
if band in bands:
|
|
469
|
+
channel_to_cols.setdefault(ch, [None] * k)
|
|
470
|
+
|
|
471
|
+
for col in df.columns:
|
|
472
|
+
if not col.endswith(suffix):
|
|
473
|
+
continue
|
|
474
|
+
core = col[:-len(suffix)]
|
|
475
|
+
if "_" not in core:
|
|
476
|
+
continue
|
|
477
|
+
ch, band = core.rsplit("_", 1)
|
|
478
|
+
if ch in channel_to_cols and band in bands:
|
|
479
|
+
channel_to_cols[ch][bands.index(band)] = col
|
|
480
|
+
|
|
481
|
+
out_cols = [f"{ch}_imfentropy" for ch in channel_to_cols.keys()]
|
|
482
|
+
rows: list[dict[str, float]] = []
|
|
483
|
+
|
|
484
|
+
for i in range(len(df)):
|
|
485
|
+
row_out: dict[str, float] = {}
|
|
486
|
+
for ch, cols_in_order in channel_to_cols.items():
|
|
487
|
+
vals = []
|
|
488
|
+
for c in cols_in_order:
|
|
489
|
+
v = df.iat[i, df.columns.get_loc(c)]
|
|
490
|
+
vals.append(float(v) if np.isfinite(v) else 0.0)
|
|
491
|
+
|
|
492
|
+
total = float(np.nansum(vals))
|
|
493
|
+
if not (np.isfinite(total) and total > 0):
|
|
494
|
+
row_out[f"{ch}_imfentropy"] = np.nan
|
|
495
|
+
continue
|
|
496
|
+
|
|
497
|
+
p = np.asarray(vals, dtype=float) / total
|
|
498
|
+
p = np.clip(p, eps, 1.0)
|
|
499
|
+
p /= p.sum()
|
|
500
|
+
H = -np.sum(p * np.log(p))
|
|
501
|
+
row_out[f"{ch}_imfentropy"] = float(H / (norm or 1.0))
|
|
502
|
+
|
|
503
|
+
rows.append(row_out)
|
|
504
|
+
|
|
505
|
+
return pd.DataFrame(rows, columns=out_cols)
|
|
506
|
+
|
|
507
|
+
|
|
508
|
+
if __name__ == "__main__":
|
|
509
|
+
FS = 128
|
|
510
|
+
csv_path = "DREAMER.csv"
|
|
511
|
+
chunk_iter = pd.read_csv(csv_path, chunksize=1)
|
|
512
|
+
first_chunk = next(chunk_iter)
|
|
513
|
+
sensor_columns = [col for col in first_chunk.columns if col[len(col)-1].isdigit()]
|
|
514
|
+
print(f"Detected sensor columns: {sensor_columns}")
|
|
515
|
+
|
|
516
|
+
dreamer_df = []
|
|
517
|
+
|
|
518
|
+
for chunk in pd.read_csv(csv_path, chunksize=10000):
|
|
519
|
+
sensor_df = chunk[sensor_columns]
|
|
520
|
+
dreamer_df.append(sensor_df)
|
|
521
|
+
|
|
522
|
+
dreamer_df = pd.concat(dreamer_df, ignore_index=True)
|
|
523
|
+
|
|
524
|
+
clean = bandpass_filter(dreamer_df, FS, bands=FREQUENCY_BANDS, low=0.5, high=45.0, notch_hz=60)
|
|
525
|
+
print("Bandpass filtering\n", clean)
|
|
526
|
+
|
|
527
|
+
hj = hjorth_params(clean, FS)
|
|
528
|
+
print("Hjorth Parameters\n", hj)
|
|
529
|
+
|
|
530
|
+
psd_df = psd_bandpowers(clean, FS, bands=FREQUENCY_BANDS)
|
|
531
|
+
print("PSD\n", psd_df)
|
|
532
|
+
|
|
533
|
+
shannons_df = shannons_entropy(clean, FS, bands=FREQUENCY_BANDS)
|
|
534
|
+
print("Shannons\n", shannons_df)
|
|
535
|
+
|
|
536
|
+
wt_df = wavelet_band_energy(dreamer_df, FS, bands=FREQUENCY_BANDS)
|
|
537
|
+
print("WT Energy\n", wt_df)
|
|
538
|
+
|
|
539
|
+
wt_df = wavelet_entropy(wt_df, bands=FREQUENCY_BANDS)
|
|
540
|
+
print("WT Entropy\n", wt_df)
|
|
541
|
+
|
|
542
|
+
imf_df = imf_band_energy(dreamer_df, FS)
|
|
543
|
+
print("IMF Energy\n", imf_df)
|
|
544
|
+
|
|
545
|
+
imf_df = imf_entropy(imf_df)
|
|
546
|
+
print("IMF Entropy\n", imf_df)
|
eegproc/preprocessing.py
ADDED
|
@@ -0,0 +1,125 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pandas as pd
|
|
3
|
+
from scipy.signal import butter, sosfiltfilt, iirnotch, filtfilt
|
|
4
|
+
from scipy.signal import detrend as scipy_detrend
|
|
5
|
+
|
|
6
|
+
FREQUENCY_BANDS = {
|
|
7
|
+
"delta": (0.5, 4.0),
|
|
8
|
+
"theta": (4.0, 8.0),
|
|
9
|
+
"alpha": (8.0, 13.0),
|
|
10
|
+
"betaL": (13.0, 20.0),
|
|
11
|
+
"betaH": (20.0, 30.0),
|
|
12
|
+
"gamma": (30.0, 45.0),
|
|
13
|
+
}
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def apply_detrend(detrend: str | None, df: pd.DataFrame) -> pd.DataFrame:
|
|
17
|
+
if detrend in {"constant", "linear"}:
|
|
18
|
+
df = detrend_df(df, kind=detrend)
|
|
19
|
+
elif detrend is None:
|
|
20
|
+
df = _numeric_interp(df)
|
|
21
|
+
else:
|
|
22
|
+
raise ValueError("detrend must be 'constant', 'linear', or None")
|
|
23
|
+
|
|
24
|
+
return df
|
|
25
|
+
|
|
26
|
+
def _numeric_interp(df: pd.DataFrame) -> pd.DataFrame:
|
|
27
|
+
df = df.select_dtypes(include=[np.number]).astype(float).copy()
|
|
28
|
+
return df.apply(lambda s: s.interpolate(limit_direction="both"))
|
|
29
|
+
|
|
30
|
+
def detrend_df(df: pd.DataFrame, kind: str = "linear") -> pd.DataFrame:
|
|
31
|
+
df = _numeric_interp(df)
|
|
32
|
+
arr = df.to_numpy(copy=False)
|
|
33
|
+
arr = scipy_detrend(arr, type=kind, axis=0)
|
|
34
|
+
return pd.DataFrame(arr, index=df.index, columns=df.columns)
|
|
35
|
+
|
|
36
|
+
def bandpass_filter(
|
|
37
|
+
df: pd.DataFrame,
|
|
38
|
+
fs: float,
|
|
39
|
+
bands: dict[str, tuple[float, float]] = FREQUENCY_BANDS,
|
|
40
|
+
low: float | None = None,
|
|
41
|
+
high: float | None = None,
|
|
42
|
+
*,
|
|
43
|
+
order: int = 4,
|
|
44
|
+
notch_hz: float | int | list | tuple | None = None,
|
|
45
|
+
notch_q: float = 30.0,
|
|
46
|
+
reref: bool = True,
|
|
47
|
+
detrend: bool = True,
|
|
48
|
+
) -> pd.DataFrame:
|
|
49
|
+
df = _numeric_interp(df).apply(pd.to_numeric, errors="coerce")
|
|
50
|
+
df = df.astype("float64")
|
|
51
|
+
nyq = fs / 2.0
|
|
52
|
+
cols = list(df.columns)
|
|
53
|
+
|
|
54
|
+
def _sosfiltfilt_safe(sos, y):
|
|
55
|
+
if np.all(np.isnan(y)):
|
|
56
|
+
return y
|
|
57
|
+
y = y.copy()
|
|
58
|
+
nans = np.isnan(y)
|
|
59
|
+
if nans.any():
|
|
60
|
+
iddf = np.where(~nans)[0]
|
|
61
|
+
if iddf.size >= 2:
|
|
62
|
+
y[nans] = np.interp(np.flatnonzero(nans), iddf, y[iddf])
|
|
63
|
+
else:
|
|
64
|
+
y[nans] = 0.0
|
|
65
|
+
if y.size < 15:
|
|
66
|
+
return y
|
|
67
|
+
return sosfiltfilt(sos, y)
|
|
68
|
+
|
|
69
|
+
def _apply_notch_once(dfin: pd.DataFrame) -> pd.DataFrame:
|
|
70
|
+
if notch_hz is None:
|
|
71
|
+
return dfin
|
|
72
|
+
freqs = notch_hz if isinstance(notch_hz, (list, tuple)) else [notch_hz]
|
|
73
|
+
edfpanded = []
|
|
74
|
+
for f0 in freqs:
|
|
75
|
+
edfpanded.append(float(f0))
|
|
76
|
+
if 2 * f0 < nyq - 1.0:
|
|
77
|
+
edfpanded.append(float(2 * f0))
|
|
78
|
+
out = dfin.copy()
|
|
79
|
+
for f0 in edfpanded:
|
|
80
|
+
w0 = f0 / nyq
|
|
81
|
+
if 0 < w0 < 1:
|
|
82
|
+
b, a = iirnotch(w0, notch_q)
|
|
83
|
+
for c in cols:
|
|
84
|
+
y = out[c].to_numpy()
|
|
85
|
+
if y.size >= max(15, 3 * max(len(a), len(b))):
|
|
86
|
+
out[c] = filtfilt(b, a, y, method="gust")
|
|
87
|
+
return out
|
|
88
|
+
|
|
89
|
+
if reref:
|
|
90
|
+
car = df.mean(axis=1)
|
|
91
|
+
for c in cols:
|
|
92
|
+
df[c] = df[c] - car
|
|
93
|
+
|
|
94
|
+
df = _apply_notch_once(df)
|
|
95
|
+
|
|
96
|
+
if bands is None:
|
|
97
|
+
if low is None or high is None:
|
|
98
|
+
raise ValueError("Provide (low, high) or a bands dict.")
|
|
99
|
+
if not (0 < low < high < nyq):
|
|
100
|
+
raise ValueError(f"Cutoffs must satisfy 0 < low < high < fs/2={nyq:.3f}.")
|
|
101
|
+
sos = butter(order, [low/nyq, high/nyq], btype="bandpass", output="sos")
|
|
102
|
+
Y = pd.DataFrame(index=df.index)
|
|
103
|
+
for c in cols:
|
|
104
|
+
Y[c] = _sosfiltfilt_safe(sos, df[c].to_numpy())
|
|
105
|
+
if detrend:
|
|
106
|
+
for c in cols:
|
|
107
|
+
Y[c] = scipy_detrend(Y[c].to_numpy(), type="constant")
|
|
108
|
+
return Y
|
|
109
|
+
|
|
110
|
+
out = {}
|
|
111
|
+
band_sos = {}
|
|
112
|
+
for name, (lo, hi) in bands.items():
|
|
113
|
+
if not (0 < lo < hi < nyq):
|
|
114
|
+
raise ValueError(f"Bad band {name}: {lo}-{hi} vs fs/2={nyq:.3f}")
|
|
115
|
+
band_sos[name] = butter(order, [lo/nyq, hi/nyq], btype="bandpass", output="sos")
|
|
116
|
+
|
|
117
|
+
for c in cols:
|
|
118
|
+
y0 = df[c].to_numpy()
|
|
119
|
+
for name, sos in band_sos.items():
|
|
120
|
+
yb = _sosfiltfilt_safe(sos, y0)
|
|
121
|
+
if detrend:
|
|
122
|
+
yb = scipy_detrend(yb, type="constant")
|
|
123
|
+
out[f"{c}_{name}"] = yb
|
|
124
|
+
|
|
125
|
+
return pd.DataFrame(out, index=df.index)
|