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 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
+ ]
@@ -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)
@@ -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)