""" 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