保姆级教程:用Python处理清华大学SSVEP脑电数据集(附完整代码与数据重塑技巧)
清华大学SSVEP脑电数据处理实战:从MATLAB到PyTorch的全流程解析
当你第一次打开清华大学的SSVEP脑电数据集时,那个神秘的4维矩阵(64, 1500, 40, 6)可能会让你感到无从下手。作为脑机接口研究中最常用的公开数据集之一,这份数据蕴含着宝贵的科研价值,但如何将它转化为深度学习模型可用的格式,却是许多初学者面临的第一个技术障碍。
1. 理解数据集结构与实验设计
在开始写代码之前,我们需要先弄清楚这个四维矩阵每个维度的实际含义。清华大学SSVEP数据集采集自35名健康受试者,采用64导联的EEG设备记录,采样率250Hz。实验设计有几个关键特点:
- 刺激范式:40个不同频率(8-15.8Hz,间隔0.2Hz)的视觉刺激
- 试验结构:每个受试者完成6个block,每个block包含40次试验(对应40个刺激频率)
- 数据分段:每次试验记录6秒EEG数据(刺激前0.5秒+刺激后5.5秒),共1500个时间点
因此,当我们加载S01.mat文件时,得到的data矩阵四个维度分别是:
| 维度 | 大小 | 含义 |
|---|---|---|
| 0 | 64 | EEG电极数量 |
| 1 | 1500 | 时间点(6秒×250Hz) |
| 2 | 40 | 刺激目标索引 |
| 3 | 6 | 试验block编号 |
提示:理解这个结构对后续的数据重塑至关重要,错误的维度理解会导致模型训练完全失效。
2. 数据加载与初步探索
首先我们需要解决MATLAB(.mat)文件的读取问题。Python中有几种常用方法可以加载.mat文件:
import scipy.io as sio # 方法1:使用scipy.io.loadmat data = sio.loadmat('S01.mat')['data'] # 返回字典,通过键名获取数据 print(f"数据形状: {data.shape}") # 方法2:使用h5py(适用于MATLAB v7.3及以上版本) import h5py with h5py.File('S01.mat', 'r') as f: data = f['data'][:]加载数据后,我们可以进行一些基本检查:
import numpy as np # 检查数据类型和范围 print(f"数据类型: {data.dtype}") print(f"最小值: {np.min(data):.2f}, 最大值: {np.max(data):.2f}") # 查看第一个电极的第一个block的第一个刺激 import matplotlib.pyplot as plt plt.plot(data[0, :, 0, 0]) plt.title("电极1-刺激1-block1的EEG信号") plt.xlabel("时间点") plt.ylabel("幅值(μV)") plt.show()常见问题及解决方案:
- 报错:Unknown file type":检查文件路径是否正确,确认文件未损坏
- 报错:Not a MATLAB v7.3 file:改用scipy.io.loadmat方法
- 数据值异常:检查是否需要乘以缩放因子(有些.mat文件会存储缩放信息)
3. 数据重塑与维度转换
原始四维数据需要转换为适合深度学习模型的格式。根据不同的框架和模型架构,有几种常见的转换方式:
3.1 转换为PyTorch张量
对于CNN模型,通常需要将数据转换为(batch, channel, height, width)格式:
# 原始维度:(64, 1500, 40, 6) -> (240, 1, 64, 1500) data_reshaped = data.transpose(2, 3, 0, 1) # 变为(40,6,64,1500) data_reshaped = data_reshaped.reshape(-1, 64, 1500) # (240,64,1500) data_reshaped = np.expand_dims(data_reshaped, 1) # 添加通道维 (240,1,64,1500) import torch tensor_data = torch.from_numpy(data_reshaped).float() print(f"PyTorch张量形状: {tensor_data.shape}")3.2 转换为TensorFlow/Keras格式
对于RNN/LSTM模型,可能需要(time_steps, features)格式:
# (240, 1500, 64) - 样本数×时间步长×特征数 data_reshaped = data.transpose(3, 2, 0, 1) # (6,40,64,1500) data_reshaped = data_reshaped.reshape(-1, 64, 1500) # (240,64,1500) data_reshaped = data_reshaped.transpose(0, 2, 1) # (240,1500,64) import tensorflow as tf tf_data = tf.convert_to_tensor(data_reshaped, dtype=tf.float32) print(f"TensorFlow张量形状: {tf_data.shape}")3.3 标签处理
刺激频率信息存储在单独的Freq_phase.mat文件中,我们需要将其转换为适合监督学习的标签:
freq_data = sio.loadmat('Freq_phase.mat') frequencies = freq_data['freqs'][0] # 40个刺激频率 # 创建标签 (240个样本,每个样本对应一个频率) labels = np.repeat(frequencies, 6) # 每个频率重复6次 # 分类任务:将频率转换为类别索引 (0-39) unique_freqs = np.unique(frequencies) label_indices = np.array([np.where(unique_freqs == f)[0][0] for f in labels]) # 独热编码 num_classes = len(unique_freqs) one_hot_labels = np.eye(num_classes)[label_indices]4. 数据预处理与增强技巧
原始EEG数据通常需要经过一系列预处理才能获得最佳模型性能:
4.1 滤波与降噪
from scipy import signal # 带通滤波 (过滤掉非SSVEP频段的噪声) def bandpass_filter(data, lowcut=7, highcut=30, fs=250, order=5): nyq = 0.5 * fs low = lowcut / nyq high = highcut / nyq b, a = signal.butter(order, [low, high], btype='band') return signal.filtfilt(b, a, data) # 应用滤波 filtered_data = np.apply_along_axis(bandpass_filter, 1, tensor_data.numpy())4.2 标准化方法比较
不同的标准化策略对模型性能有显著影响:
| 方法 | 公式 | 适用场景 | 代码实现 |
|---|---|---|---|
| 通道标准化 | (x - μ_ch) / σ_ch | 保留通道间差异 | (data - mean(axis=0)) / std(axis=0) |
| 样本标准化 | (x - μ_sample) / σ_sample | 独立处理每个样本 | (data - mean(axis=(1,2,3))) / std(axis=(1,2,3)) |
| 全局标准化 | (x - μ_all) / σ_all | 统一所有数据尺度 | (data - global_mean) / global_std |
# 通道标准化示例 channel_mean = tensor_data.mean(dim=(0, 2, 3), keepdim=True) channel_std = tensor_data.std(dim=(0, 2, 3), keepdim=True) normalized_data = (tensor_data - channel_mean) / (channel_std + 1e-8)4.3 数据增强技术
EEG数据增强需要特别注意保持信号的生理合理性:
- 时间扭曲:对时间轴进行轻微拉伸/压缩
- 通道丢弃:随机屏蔽部分电极通道
- 高斯噪声:添加微小的随机噪声
- 频段增强:特定SSVEP频段的幅度增强
class EEGAugmentation: def __init__(self, p=0.5): self.p = p def time_warp(self, x, max_warp=0.1): if np.random.rand() > self.p: return x length = x.shape[-1] warp_factor = 1 + np.random.uniform(-max_warp, max_warp) new_length = int(length * warp_factor) warped = torch.nn.functional.interpolate( x.unsqueeze(0), size=new_length, mode='linear', align_corners=False ).squeeze(0) if warp_factor > 1: # 截断 return warped[..., :length] else: # 填充 result = torch.zeros_like(x) result[..., :new_length] = warped return result def channel_dropout(self, x, max_drop=0.2): if np.random.rand() > self.p: return x num_drop = int(x.shape[1] * max_drop * np.random.rand()) drop_indices = np.random.choice(x.shape[1], num_drop, replace=False) x[:, drop_indices] = 0 return x5. 构建PyTorch数据管道
为了高效加载和预处理数据,我们应该实现一个完整的数据加载器:
from torch.utils.data import Dataset, DataLoader class SSVEPDataset(Dataset): def __init__(self, mat_files, label_file, transform=None): self.data_files = mat_files self.labels = self._load_labels(label_file) self.transform = transform def _load_labels(self, label_file): freq_data = sio.loadmat(label_file) frequencies = freq_data['freqs'][0] return np.repeat(frequencies, 6 * len(self.data_files)) def __len__(self): return len(self.labels) def __getitem__(self, idx): subj_idx = idx // 240 # 每个受试者240个样本 sample_idx = idx % 240 data = sio.loadmat(self.data_files[subj_idx])['data'] data = data.transpose(2, 3, 0, 1).reshape(-1, 64, 1500) sample = data[sample_idx] # 转换为PyTorch CNN输入格式 sample = torch.from_numpy(sample).float().unsqueeze(0) label = torch.tensor(self.labels[idx]).long() if self.transform: sample = self.transform(sample) return sample, label # 使用示例 mat_files = [f'S{i:02d}.mat' for i in range(1, 36)] # S01.mat到S35.mat dataset = SSVEPDataset(mat_files, 'Freq_phase.mat', transform=EEGAugmentation()) dataloader = DataLoader(dataset, batch_size=32, shuffle=True)6. 可视化分析与质量检查
在投入模型训练前,进行数据质量检查至关重要:
6.1 时频分析
def plot_time_frequency(data, fs=250, nperseg=250): f, t, Sxx = signal.spectrogram(data, fs=fs, nperseg=nperseg) plt.pcolormesh(t, f, 10*np.log10(Sxx), shading='gouraud') plt.ylabel('Frequency [Hz]') plt.xlabel('Time [sec]') plt.colorbar(label='Power (dB)') plt.show() # 查看不同刺激频率下的响应 for freq_idx in [0, 10, 20, 30]: # 查看第1、11、21、31个刺激 sample = data[0, :, freq_idx, 0] # 电极1,刺激freq_idx,block1 plot_time_frequency(sample)6.2 拓扑图可视化
from mne.viz import plot_topomap def plot_electrode_topology(data, times=[0.5, 1.0, 1.5], ch_names=None): """在不同时间点绘制电极拓扑图""" fig, axes = plt.subplots(1, len(times), figsize=(15, 5)) for i, t in enumerate(times): sample_idx = int(t * 250) # 转换为采样点索引 plot_topomap(data[:, sample_idx], pos, axes=ax[i], show=False) ax[i].set_title(f'{t}s') plt.show() # 需要先加载电极位置信息 pos = np.loadtxt('64channel.loc') # 假设已转换为MNE兼容格式 plot_electrode_topology(data.mean(axis=(2,3))) # 平均所有刺激和block7. 模型输入的最佳实践
根据不同的模型架构,SSVEP数据有多种组织方式:
CNN模型:
- 输入形状:(batch, 1, 64, 1500)
- 将EEG电极视为空间维度,时间序列作为时间维度
- 适合捕捉空间-时间联合特征
RNN/LSTM模型:
- 输入形状:(batch, 1500, 64)
- 每个时间步包含所有电极的特征
- 适合建模时间依赖性
混合模型:
- 先使用CNN处理每个电极的时间序列
- 再用RNN整合各电极输出
- 最后用全连接层分类
# 简单的CNN模型示例 import torch.nn as nn class SSVEP_CNN(nn.Module): def __init__(self, num_classes=40): super().__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=(3, 25), padding=(1, 12)) self.pool1 = nn.MaxPool2d(kernel_size=(1, 5)) self.conv2 = nn.Conv2d(32, 64, kernel_size=(3, 15), padding=(1, 7)) self.pool2 = nn.MaxPool2d(kernel_size=(1, 5)) self.fc = nn.Linear(64 * 64 * 60, num_classes) # 需要根据实际尺寸调整 def forward(self, x): x = torch.relu(self.conv1(x)) x = self.pool1(x) x = torch.relu(self.conv2(x)) x = self.pool2(x) x = x.view(x.size(0), -1) return self.fc(x)处理清华大学SSVEP数据集时,最常遇到的几个坑包括:维度顺序混淆、采样率误解、标签对齐错误。特别是在数据重塑阶段,一个常见的错误是直接使用reshape而不是transpose,这会打乱数据的时空关系。记住,reshape只改变数组视图而不改变数据顺序,而transpose会实际改变数据在内存中的排列方式。
