meroni-api / eeg_processor.py
arrondavide's picture
Meroni API - EEG to fMRI translation pipeline
ba9143f
Raw History Blame Contribute Delete
10.6 kB
"""
Real EEG signal processing pipeline.
Parses EEG files and extracts features for fMRI translation.
"""
import numpy as np
from scipy import signal
from scipy.stats import kurtosis, skew
import io
import os
from typing import Optional
# Frequency bands (Hz)
BANDS = {
"delta": (1, 4),
"theta": (4, 8),
"alpha": (8, 13),
"beta": (13, 30),
"gamma": (30, 45),
}
BAND_NAMES = list(BANDS.keys())
def parse_eeg_file(file_bytes: bytes, filename: str) -> tuple[np.ndarray, float, list[str]]:
"""
Parse an EEG file and return (data, sfreq, channel_names).
data shape: (n_channels, n_samples)
Supports: .edf, .bdf, .csv, .txt
"""
ext = os.path.splitext(filename)[1].lower()
if ext in (".edf", ".bdf"):
return _parse_edf(file_bytes, ext)
elif ext in (".csv", ".txt"):
return _parse_csv(file_bytes)
else:
raise ValueError(f"Unsupported file format: {ext}")
def _parse_edf(file_bytes: bytes, ext: str) -> tuple[np.ndarray, float, list[str]]:
"""Parse EDF/BDF files using MNE."""
import mne
import tempfile
suffix = ext
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
tmp.write(file_bytes)
tmp_path = tmp.name
try:
if ext == ".edf":
raw = mne.io.read_raw_edf(tmp_path, preload=True, verbose=False)
else:
raw = mne.io.read_raw_bdf(tmp_path, preload=True, verbose=False)
# Pick only EEG channels
raw.pick(picks="eeg", exclude="bads")
data = raw.get_data() # (n_channels, n_samples)
sfreq = raw.info["sfreq"]
ch_names = raw.ch_names
return data, sfreq, ch_names
finally:
os.unlink(tmp_path)
def _parse_csv(file_bytes: bytes) -> tuple[np.ndarray, float, list[str]]:
"""
Parse CSV/TXT EEG files.
Expects: columns = channels, rows = time samples.
First row can be header (channel names).
Assumes 256 Hz sampling rate if not specified.
"""
import pandas as pd
text = file_bytes.decode("utf-8", errors="ignore")
df = pd.read_csv(io.StringIO(text))
# Try to detect if first column is timestamp
first_col = df.columns[0].lower()
if "time" in first_col or "stamp" in first_col or "sample" in first_col:
# Try to infer sample rate from timestamps
if df[df.columns[0]].dtype in [np.float64, np.float32, np.int64]:
timestamps = df[df.columns[0]].values
if len(timestamps) > 1:
dt = np.median(np.diff(timestamps))
if dt > 0:
sfreq = 1.0 / dt if dt < 1 else 1.0 / (dt / 1000.0)
else:
sfreq = 256.0
else:
sfreq = 256.0
else:
sfreq = 256.0
df = df.iloc[:, 1:] # Remove timestamp column
else:
sfreq = 256.0
ch_names = [str(c) for c in df.columns]
data = df.values.T.astype(np.float64) # (n_channels, n_samples)
return data, sfreq, ch_names
def preprocess(data: np.ndarray, sfreq: float) -> np.ndarray:
"""
Preprocess raw EEG data:
1. Bandpass filter 1-45 Hz
2. Notch filter at 50 Hz and 60 Hz (line noise)
3. Z-score normalization per channel
"""
n_channels, n_samples = data.shape
# Minimum samples check
min_samples = int(sfreq * 2) # Need at least 2 seconds
if n_samples < min_samples:
raise ValueError(
f"EEG recording too short: {n_samples/sfreq:.1f}s. Need at least 2 seconds."
)
# Bandpass filter 1-45 Hz
nyq = sfreq / 2.0
low = 1.0 / nyq
high = min(45.0 / nyq, 0.99) # Ensure below Nyquist
if high <= low:
raise ValueError(f"Sample rate too low ({sfreq} Hz) for 1-45 Hz bandpass filter.")
sos_bp = signal.butter(4, [low, high], btype="band", output="sos")
filtered = signal.sosfiltfilt(sos_bp, data, axis=1)
# Notch filters for line noise
for notch_freq in [50.0, 60.0]:
if notch_freq < nyq:
b_notch, a_notch = signal.iirnotch(notch_freq, Q=30.0, fs=sfreq)
filtered = signal.filtfilt(b_notch, a_notch, filtered, axis=1)
# Z-score normalize each channel
means = filtered.mean(axis=1, keepdims=True)
stds = filtered.std(axis=1, keepdims=True)
stds[stds < 1e-10] = 1.0 # Avoid division by zero
normalized = (filtered - means) / stds
return normalized
def extract_psd_features(data: np.ndarray, sfreq: float) -> np.ndarray:
"""
Extract Power Spectral Density features per channel per band.
Returns: (n_channels * n_bands,) = typically (8 * 5 = 40,) features
"""
n_channels = data.shape[0]
n_bands = len(BANDS)
features = np.zeros(n_channels * n_bands)
for ch in range(n_channels):
# Welch's method for PSD estimation
freqs, psd = signal.welch(
data[ch], fs=sfreq, nperseg=min(256, data.shape[1]), noverlap=128
)
for bi, (band_name, (fmin, fmax)) in enumerate(BANDS.items()):
idx = np.where((freqs >= fmin) & (freqs <= fmax))[0]
if len(idx) > 0:
# Relative band power (log-transformed)
band_power = np.trapz(psd[idx], freqs[idx])
total_power = np.trapz(psd, freqs)
if total_power > 0:
features[ch * n_bands + bi] = np.log1p(
band_power / total_power
)
return features
def extract_coherence_features(data: np.ndarray, sfreq: float) -> np.ndarray:
"""
Extract coherence between all channel pairs across frequency bands.
For n channels: n*(n-1)/2 pairs * n_bands features.
"""
n_channels = data.shape[0]
n_pairs = n_channels * (n_channels - 1) // 2
n_bands = len(BANDS)
features = np.zeros(n_pairs * n_bands)
pair_idx = 0
for i in range(n_channels):
for j in range(i + 1, n_channels):
# Magnitude squared coherence
freqs, coh = signal.coherence(
data[i], data[j], fs=sfreq,
nperseg=min(256, data.shape[1]),
noverlap=128,
)
for bi, (band_name, (fmin, fmax)) in enumerate(BANDS.items()):
idx = np.where((freqs >= fmin) & (freqs <= fmax))[0]
if len(idx) > 0:
features[pair_idx * n_bands + bi] = np.mean(coh[idx])
pair_idx += 1
return features
def extract_hjorth_features(data: np.ndarray) -> np.ndarray:
"""
Extract Hjorth parameters (Activity, Mobility, Complexity) per channel.
Returns: (n_channels * 3,) features
"""
n_channels = data.shape[0]
features = np.zeros(n_channels * 3)
for ch in range(n_channels):
x = data[ch]
# Activity = variance of the signal
activity = np.var(x)
# First derivative
dx = np.diff(x)
dx_var = np.var(dx)
# Second derivative
ddx = np.diff(dx)
ddx_var = np.var(ddx)
# Mobility = sqrt(var(dx) / var(x))
mobility = np.sqrt(dx_var / activity) if activity > 0 else 0
# Complexity = mobility(dx) / mobility(x)
mobility_dx = np.sqrt(ddx_var / dx_var) if dx_var > 0 else 0
complexity = mobility_dx / mobility if mobility > 0 else 0
features[ch * 3] = np.log1p(activity) # Log-transform activity
features[ch * 3 + 1] = mobility
features[ch * 3 + 2] = complexity
return features
def extract_statistical_features(data: np.ndarray) -> np.ndarray:
"""
Extract statistical features per channel: mean, std, skewness, kurtosis.
Returns: (n_channels * 4,) features
"""
n_channels = data.shape[0]
features = np.zeros(n_channels * 4)
for ch in range(n_channels):
x = data[ch]
features[ch * 4] = np.mean(np.abs(x))
features[ch * 4 + 1] = np.std(x)
features[ch * 4 + 2] = skew(x)
features[ch * 4 + 3] = kurtosis(x)
return features
def extract_all_features(
data: np.ndarray, sfreq: float, target_channels: int = 8
) -> np.ndarray:
"""
Extract all features from preprocessed EEG data.
Handles variable channel counts by padding or selecting channels.
Returns feature vector ready for the translation model.
"""
n_channels = data.shape[0]
# Handle channel count mismatch
if n_channels > target_channels:
# Select evenly spaced channels
indices = np.linspace(0, n_channels - 1, target_channels, dtype=int)
data = data[indices]
elif n_channels < target_channels:
# Pad with zeros
padding = np.zeros((target_channels - n_channels, data.shape[1]))
data = np.vstack([data, padding])
n_ch = data.shape[0] # Should be target_channels now
# Extract all feature types
psd = extract_psd_features(data, sfreq) # n_ch * 5 = 40
coherence = extract_coherence_features(data, sfreq) # n_ch*(n_ch-1)/2 * 5 = 140
hjorth = extract_hjorth_features(data) # n_ch * 3 = 24
stats = extract_statistical_features(data) # n_ch * 4 = 32
# Concatenate all features
all_features = np.concatenate([psd, coherence, hjorth, stats])
return all_features
def get_feature_dimension(n_channels: int = 8) -> int:
"""Calculate total feature dimension for given channel count."""
n_bands = 5
n_pairs = n_channels * (n_channels - 1) // 2
psd_dim = n_channels * n_bands # 40
coh_dim = n_pairs * n_bands # 140
hjorth_dim = n_channels * 3 # 24
stats_dim = n_channels * 4 # 32
return psd_dim + coh_dim + hjorth_dim + stats_dim # 236
def process_eeg_file(
file_bytes: bytes, filename: str, target_channels: int = 8
) -> tuple[np.ndarray, dict]:
"""
Full pipeline: parse → preprocess → extract features.
Returns (feature_vector, metadata).
"""
# Parse
data, sfreq, ch_names = parse_eeg_file(file_bytes, filename)
metadata = {
"filename": filename,
"original_channels": len(ch_names),
"channel_names": ch_names[:20], # Limit for response size
"sample_rate": sfreq,
"duration_seconds": round(data.shape[1] / sfreq, 2),
"n_samples": data.shape[1],
"target_channels": target_channels,
}
# Preprocess
preprocessed = preprocess(data, sfreq)
# Extract features
features = extract_all_features(preprocessed, sfreq, target_channels)
metadata["feature_dimension"] = len(features)
return features, metadata