CNN-Transformer运动想象脑电分类:从预处理到注意力可视化的工程实践

发布时间:2026/10/6 8:16:21
CNN-Transformer运动想象脑电分类:从预处理到注意力可视化的工程实践 简介这份资源面向计算机、人工智能、通信工程、自动化等专业的高校学生、教师与科研人员提供一套基于CNN-Transformer框架的运动想象脑电信号分类完整实现可用于毕业设计、课程设计、科研实验或项目立项演示。压缩包共32个文件、约18.5MB以23个Python脚本为核心覆盖数据预处理、模型定义、训练与K折交叉验证、可视化分析等环节另含2个xlsx统计表、2个m脚本、1个pth权重文件、1个npy数据文件、1个xml配置、1个docx设计报告与1个md说明文档兼顾代码运行与文档参考。资源内含CNNTransformer、EEGNet、Conformer、空间注意力、Morlet变换等模型模块以及t-SNE、CAM热力图、AUC与箱线图等分析脚本便于理解脑电信号分类全流程。目前已有46人学习适合在此基础上修改扩展或直接用于课题实践。1. 从一段 4 秒脑电说起CNN-Transformer 到底在运动想象分类里干什么运动想象Motor Imagery, MI脑电分类这件事做过的人都知道它有多“玄学”。受试者只是在大脑里想象左手或右手运动头皮上那几微伏的电位变化就被 250Hz 采样记录下来信噪比低到令人发指。传统 CSPLDA 方案在 2 分类上勉强能到 70% 上下一旦跨受试者、跨 session精度立刻跳水。这几年 Transformer 在视觉、NLP 里大杀四方很多人第一反应就是把它搬到脑电上——但直接把原始 EEG 序列丢进标准 Transformer效果往往还不如 CSP因为脑电的局部节律特征mu 节律 8-13Hz、beta 节律 13-30Hz是卷积核最擅长抓的东西而全局通道间的依赖关系才是注意力机制的主场。CNN-Transformer 框架的核心思路就是让 CNN 先做局部时空滤波把每个通道的节律特征压成 token再交给 Transformer 建模通道间和长时依赖。这篇笔记就围绕这个框架把数据预处理、模型搭建、训练调参、踩坑排查整条链路讲清楚适合已经跑通过基础 EEG 分类、想上 Transformer 但被精度和显存反复折磨的从业者。2. 数据与预处理把 raw EEG 变成 Transformer 能吃的 token2.1 为什么不能直接把原始信号喂给 Transformer标准 Transformer 的 self-attention 复杂度是 O(n²)n 是序列长度。一段 4 秒、250Hz 的脑电单通道就是 1000 个点8 通道就是 8000 维输入。如果按时间点做 token序列长度直接爆炸而且每个时间点只包含单一时刻的电压值缺乏频域和局部节律信息注意力机制根本学不到东西。常见做法是先做带通滤波 滑动窗口切片把每个窗口内的多通道信号当作一个 token 的原始特征再用 CNN 提取局部时空特征最后把 CNN 输出的特征图按通道维度展平成 token 序列。这样序列长度从几千降到几十注意力才有意义。我一般会先把数据统一到 250Hz 采样、0.5-40Hz 带通含 50Hz 陷波再用 2 秒窗口、50% 重叠切片。BCI Competition IV 2a 数据集是这套流程最常用的验证集9 个受试者、4 分类左手、右手、双脚、舌头每受试者约 288 个 trial。下面这段预处理代码可以直接抄。import numpy as np from scipy.signal import butter, filtfilt, iirnotch def bandpass_notch(data, fs250, low0.5, high40, notch_freq50): # data shape: (n_trials, n_channels, n_times) b, a butter(4, [low/(fs/2), high/(fs/2)], btypeband) data filtfilt(b, a, data, axis-1) bn, an iirnotch(notch_freq/(fs/2), Q30) data filtfilt(bn, an, data, axis-1) return data def sliding_window(data, labels, window_sec2, overlap0.5, fs250): win int(window_sec * fs) step int(win * (1 - overlap)) X, y [], [] for trial, label in zip(data, labels): for start in range(0, trial.shape[1] - win 1, step): X.append(trial[:, start:startwin]) y.append(label) return np.array(X), np.array(y)butter(4, ...)里的 4 是滤波器阶数阶数越高过渡带越陡但相位失真越大脑电这种低幅信号建议不超过 4 阶。iirnotch的 Q30 控制陷波带宽Q 越大陷波越窄避免误伤 50Hz 附近的 beta 节律。滑窗的overlap0.5是经验值再高会导致训练集和验证集高度相关评估虚高。2.2 标准化与数据增强的三个必调参数脑电不同通道的幅值差异很大尤其靠近枕区的通道容易被眼电污染。逐通道 z-score 标准化是标配但要注意用训练集的均值和方差去标准化验证集否则就是数据泄漏。数据增强方面运动想象最有效的是通道随机丢弃和高斯噪声注入时间平移和缩放反而容易破坏节律结构。def normalize(X_train, X_val): mu X_train.mean(axis(0, 2), keepdimsTrue) sigma X_train.std(axis(0, 2), keepdimsTrue) 1e-8 return (X_train - mu) / sigma, (X_val - mu) / sigma def augment(X, drop_prob0.1, noise_std0.05): # 通道随机丢弃 mask np.random.binomial(1, 1-drop_prob, size(X.shape[0], X.shape[1], 1)) X X * mask # 高斯噪声 X X np.random.normal(0, noise_std, X.shape) return Xdrop_prob0.1是安全区间超过 0.2 会让某些 trial 的有效通道太少模型学不到稳定模式。noise_std0.05对应标准化后的数据相当于原始信号里几微伏的扰动再大就淹没真实特征了。这三个参数窗口长度、丢弃率、噪声强度是我调参时最先动的窗口从 2 秒改到 1.5 秒往往能涨 1-2 个点因为运动想象的 ERD/ERS 现象在 1-2 秒内最明显。3. CNN-Transformer 模型搭建从卷积提特征到注意力建全局依赖3.1 CNN 前端怎么设计才不浪费参数量CNN 前端的目标不是分类而是把(batch, channels, time)的输入压成(batch, tokens, d_model)。常见做法是先用一个时间卷积核如 25 点约 100ms抓局部波形再用一个空间卷积核如 8 通道全连接做通道加权这其实就是 EEGNet 的深度可分离卷积思路。我一般会堆两层时间卷积 一层空间卷积每层后面接 BatchNorm 和 ELU池化只沿时间轴做保留通道维度作为后续 token 的来源。import torch import torch.nn as nn class CNNFrontend(nn.Module): def __init__(self, n_channels8, d_model64): super().__init__() self.temporal nn.Sequential( nn.Conv2d(1, 16, (1, 25), padding(0, 12)), nn.BatchNorm2d(16), nn.ELU(), nn.Conv2d(16, 32, (1, 15), padding(0, 7)), nn.BatchNorm2d(32), nn.ELU(), nn.AvgPool2d((1, 4)) # 时间下采样4倍 ) self.spatial nn.Sequential( nn.Conv2d(32, d_model, (n_channels, 1)), nn.BatchNorm2d(d_model), nn.ELU() ) def forward(self, x): # x: (B, 1, C, T) x self.temporal(x) # (B, 32, C, T/4) x self.spatial(x) # (B, d_model, 1, T/4) x x.squeeze(2).permute(0, 2, 1) # (B, T/4, d_model) return x时间卷积核 25 点对应 100ms刚好覆盖一个 mu 节律周期第二层 15 点约 60ms抓 beta 节律。池化 4 倍后2 秒窗口的 500 点变成 125 个 token这个长度对 Transformer 很友好。d_model64是脑电任务的常用值再大容易过拟合再小注意力头数不好分。3.2 Transformer 编码器的层数、头数与位置编码选择CNN 输出的 token 已经带有时序信息但 Transformer 本身是置换不变的必须加位置编码。脑电这种准周期信号可学习位置编码比正弦编码更灵活我一般用nn.Parameter直接学一个(max_len, d_model)的表。编码器层数建议 2-4 层头数 4-8前馈维度是 d_model 的 2-4 倍。层数超过 4 层在 288 个 trial 的数据集上几乎必然过拟合除非你做跨受试者预训练。class EEGTransformer(nn.Module): def __init__(self, n_channels8, d_model64, nhead4, num_layers2, n_classes4, max_len200): super().__init__() self.cnn CNNFrontend(n_channels, d_model) self.pos_embed nn.Parameter(torch.randn(1, max_len, d_model) * 0.02) encoder_layer nn.TransformerEncoderLayer( d_modeld_model, nheadnhead, dim_feedforwardd_model*4, dropout0.3, batch_firstTrue, activationgelu ) self.encoder nn.TransformerEncoder(encoder_layer, num_layersnum_layers) self.cls_token nn.Parameter(torch.randn(1, 1, d_model) * 0.02) self.head nn.Sequential( nn.LayerNorm(d_model), nn.Linear(d_model, n_classes) ) def forward(self, x): x self.cnn(x) # (B, L, d_model) B, L, D x.shape cls self.cls_token.expand(B, -1, -1) x torch.cat([cls, x], dim1) # (B, L1, D) x x self.pos_embed[:, :L1, :] x self.encoder(x) return self.head(x[:, 0]) # 取 cls token 分类dropout0.3是脑电 Transformer 的保命参数注意力层和前馈层都要加。cls_token借鉴 ViT比平均池化更稳定。位置编码初始化用 0.02 倍标准差太大前期训练震荡太小等于没加。这套配置在 BCI IV 2a 单受试者上4 分类能到 75%-82%比纯 EEGNet 高 3-5 个点但训练时间翻倍显存占用约 1.5GBbatch32。3.3 训练循环里必须监控的三个量训练脚本本身不复杂但脑电任务有几个专属监控点每 epoch 的训练/验证损失比、混淆矩阵的类别偏斜、注意力权重的熵。损失比超过 1.5 说明过拟合混淆矩阵如果某一类 recall 低于 40% 说明该类的 ERD 模式没学到注意力熵持续下降说明注意力坍缩到少数 token。def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct, total 0, 0, 0 for X, y in loader: X, y X.to(device), y.to(device) optimizer.zero_grad() logits model(X) loss criterion(logits, y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() * X.size(0) correct (logits.argmax(1) y).sum().item() total X.size(0) return total_loss / total, correct / totalclip_grad_norm_的 max_norm1.0 是 Transformer 训练的标配脑电信号梯度容易爆。优化器用 AdamWlr1e-3weight_decay0.01配合 CosineAnnealing 到 1e-5。batch size 别超过 64小 batch 的梯度噪声反而有助于跳出局部最优。早停 patience 设 15-20 epoch脑电任务通常 60-100 epoch 收敛。4. 避坑与排查运动想象 Transformer 最容易翻车的五件事4.1 验证集精度比训练集还高现象训练 loss 正常下降但验证 accuracy 在第 10 个 epoch 就超过训练 accuracy且持续领先。原因滑窗切片时训练集和验证集来自同一 trial 的重叠窗口信息泄漏。或者标准化时用了全量数据的均值方差。解决按 trial 划分训练/验证不要按窗口划分标准化统计量只在训练集上算。如果做跨 session 评估验证集必须来自不同 session。4.2 注意力权重全是均匀分布现象可视化 attention map发现所有 token 的权重几乎一样模型退化成平均池化。原因位置编码没加或初始化太小或者 CNN 前端下采样太狠token 之间已经高度相似注意力无从区分。解决检查 pos_embed 是否真的加到了输入上把池化倍数从 4 降到 2让 token 保留更多差异适当提高位置编码初始化标准差到 0.05。4.3 某一类 recall 始终为 0现象4 分类里“舌头”或“双脚”类的 recall 长期低于 10%其他类正常。原因该类 trial 数量少类别不平衡或者该类运动想象的 ERD 空间分布与其它类高度重叠CNN 空间卷积没区分开。解决用带权重的 CrossEntropyLoss权重取类别频率的倒数或者改用 Focal Lossgamma2。同时检查空间卷积核是否覆盖了 C3/C4/Cz 这些关键通道如果数据里没有这些通道分类上限本身就很低。4.4 训练到一半 loss 变 NaN现象前 20 个 epoch 正常突然 loss 爆成 NaN梯度也变 NaN。原因脑电信号里有未滤除的工频干扰或电极脱落造成的尖峰标准化后仍有个别极大值或者学习率在后期没有衰减AdamW 的二阶矩估计被异常梯度带偏。解决预处理加一步幅值裁剪如 ±100μV训练时开clip_grad_norm_学习率用 warmup cosine前 5 个 epoch 从 1e-5 线性升到 1e-3。4.5 换受试者后精度断崖下跌现象在受试者 A 上训练到 80%直接用到受试者 B 上只剩 45%接近随机。原因不同受试者的脑电幅值、节律频率、电极阻抗差异巨大模型学到了受试者特定的模式而非通用运动想象特征。解决做跨受试者评估时必须用欧式对齐EA或 Riemannian 对齐把每个受试者的协方差矩阵对齐到单位矩阵或者用留一受试者交叉验证并在训练时混入多个受试者的数据做域泛化。单受试者模型不要指望直接迁移。5. 进阶技巧用注意力可视化反推电极位置与频带选择模型跑通之后最有价值的不是那零点几个点的精度而是注意力权重能告诉你哪些通道、哪些时间段真正在起作用。我习惯在验证集上把最后一层 Transformer 的 attention map 按 token 位置平均再映射回时间轴和原始信号的 ERD/ERS 曲线叠在一起看。如果注意力高峰出现在运动想象开始后 0.5-1.5 秒且集中在 C3/C4 通道对应的 token 上说明模型学到的和神经生理学一致这个模型才值得信任。具体做法是给nn.TransformerEncoderLayer加一个need_weightsTrue的 hook或者直接改用nn.MultiheadAttention手动前向。下面这段代码把注意力权重导出并做时间对齐。def extract_attention(model, x, layer_idx-1): model.eval() attns [] hooks [] def hook(module, inp, out): # out: (B, L, D)需要重新计算 attention 权重 pass # 更稳妥的方式手动跑 encoder 层 with torch.no_grad(): feat model.cnn(x) B, L, D feat.shape cls model.cls_token.expand(B, -1, -1) seq torch.cat([cls, feat], dim1) model.pos_embed[:, :L1, :] for i, layer in enumerate(model.encoder.layers): # 取 self-attention 的权重 attn_out, attn_w layer.self_attn(seq, seq, seq, need_weightsTrue) seq layer.norm1(seq layer.dropout1(attn_out)) seq layer.norm2(seq layer.dropout2(layer.linear2( layer.dropout(layer.activation(layer.linear1(seq)))))) if i layer_idx % len(model.encoder.layers): attns.append(attn_w.cpu()) return attns[0] # (B, L1, L1)拿到(B, L1, L1)的注意力矩阵后取 cls token 那一行attn[:, 0, 1:]对 batch 平均就得到每个时间 token 的重要性。再按 CNN 的池化倍数4 倍和窗口步长反推回原始时间轴。如果发现注意力集中在窗口边缘说明位置编码或窗口切分有问题如果集中在中间说明模型抓到了运动想象的中间段 ERD。这个分析我每次训练完都会做一遍比看 accuracy 曲线有用得多。几个实操习惯第一永远先跑通 EEGNet 或 CSPLDA 作为 baselineCNN-Transformer 涨点不超过 3 个点就别硬上性价比太低第二注意力可视化只信验证集上的训练集上的注意力图会骗人第三跨受试者场景优先考虑对齐和域适应模型结构再花哨也救不了分布偏移。这套框架我前后调了三个月最大的教训就是别一上来就堆 Transformer 层数先把 CNN 前端和预处理做扎实注意力机制才有发挥空间。希望帮到你。本文还有配套的精品资源点击获取