Spaces:
Sleeping
Sleeping
Download eeg_processor.py from arrondavide/meroni-api: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/spaces/arrondavide/meroni-api/resolve/main/eeg_processor.py
- Command line
-
hf download hf://spaces/arrondavide/meroni-api/eeg_processor.py
-
curl -L -o eeg_processor.py https://huggingface.co/spaces/arrondavide/meroni-api/resolve/main/eeg_processor.py
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 | |