Bonn数据集入门指南:从数据加载到模型训练的全流程解析

1次阅读
没有评论

共计 1873 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

数据集概述

Bonn 数据集是癫痫脑电信号 (EEG) 分析领域的经典数据集,包含健康人群和癫痫患者发作期 / 间歇期的 5 种状态记录。该数据集被广泛用于癫痫预测、睡眠分期等研究,特点是采样规范、标注明确,非常适合初学者理解 EEG 信号特性。

Bonn 数据集入门指南:从数据加载到模型训练的全流程解析

环境准备

推荐使用 Python 3.8+ 环境,主要依赖库可通过以下命令安装:

pip install numpy scipy matplotlib scikit-learn

数据加载与探索

  1. 下载数据集
  2. 官方地址:https://epileptologie-bonn.de/cms/front_content.php?idcat=193
  3. 数据集包含 5 个子集(Z/S/O/N/F),每个子集 100 个单通道 EEG 片段

  4. 读取.mat 文件

    import scipy.io as sio
    import numpy as np
    
    # 示例加载健康人闭眼状态数据(子集 O)data = sio.loadmat('O.mat')
    eeg_signals = data['data']  # 获取 EEG 信号数据
    print(f"数据形状:{eeg_signals.shape}")  # 典型输出:(100, 4097)

预处理流程

  1. 带通滤波(0.5-40Hz):

    from scipy import signal
    
    def bandpass_filter(data, low=0.5, high=40, fs=173.61):
        b, a = signal.butter(4, [low, high], btype='bandpass', fs=fs)
        return signal.filtfilt(b, a, data)
    
    filtered = np.apply_along_axis(bandpass_filter, 1, eeg_signals[:, :-1])  # 最后 1 列是标签

  2. 数据分段

    # 将每个 4096 点样本分成 8 个 512 点片段
    segments = np.array([filtered[:, i*512:(i+1)*512] 
                         for i in range(8)]).transpose(1,0,2)
    print(f"分段后形状:{segments.shape}")  # (100,8,512)

特征工程

  1. 时域特征(均值 / 方差):

    time_features = np.array([[np.mean(seg), np.std(seg)] 
        for sample in segments for seg in sample
    ]).reshape(-1, 16)  # 每个样本 16 维特征

  2. 频域特征(功率谱密度):

    freq_features = []
    for sample in segments:
        for seg in sample:
            f, psd = signal.welch(seg, fs=173.61)
            freq_features.append(psd[:30])  # 取前 30 个频点

模型训练与评估

  1. 数据准备

    from sklearn.model_selection import train_test_split
    
    X = np.hstack([time_features, np.array(freq_features)])
    y = eeg_signals[:, -1]  # 标签列
    X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)

  2. SVM 分类器

    from sklearn.svm import SVC
    from sklearn.metrics import accuracy_score
    
    model = SVC(kernel='rbf')
    model.fit(X_train, y_train)
    preds = model.predict(X_test)
    print(f"测试集准确率:{accuracy_score(y_test, preds):.2f}")

优化建议

  1. 预处理优化
  2. 尝试不同滤波器阶数(4- 8 阶)
  3. 比较 Butterworth 与 Chebyshev 滤波器效果

  4. 特征选择

  5. 加入非线性特征(近似熵、样本熵)
  6. 使用 PCA 降维消除冗余特征

常见问题解答

Q1:为什么我的准确率始终低于 50%?
A1:检查标签是否正确对齐,Bonn 数据集最后 1 列是标签(0- 4 对应不同状态)

Q2:如何处理内存不足问题?
A2:可先进行下采样(如从 4096 点降至 1024 点),或改用生成器逐步加载数据

Q3:特征提取耗时太长怎么办?
A3:使用 numba 加速计算,或提前提取特征保存为.npy 文件

结语

通过本指南,我们完成了从数据加载到模型训练的全流程。建议初学者先用子集 O 和 Z(健康人数据)练手,逐步扩展到其他类别。后续可尝试 LSTM 等时序模型,比较不同方法的性能差异。

正文完
 0
评论(没有评论)