EEGNet脑电解码 入门教学与笔记
EEGNet的结构1*1*22*1000经过1*8*1*64的时间卷积核生成1*8*22*1000的时间特征图经过16*8*22*1的空间卷积核变成了1*16*1*1000的时间空间特征图再经过两轮14池化层和18池化层压缩时间变成1*16*1*311000/4/831的增强特征再经过分类卷积4*16*1*31分出1*4*1*1的任务分类结构弄明白下载BCIC-IV-2a数据集地址https://www.bbci.de/competition/iv/download/index.html?agreeyessubmit提交conda中配置好环境将数据集和项目文件放一起可以开始代码了GitHub链接GitHub - fahener/EEGNet-BCICIV-2a: EEGNet脑电解码-入门代码 · GitHub以下是全文注释代码环境mne numpy pytorch sklearn pandasimport mne import numpy as np import torch from torch.utils.data import Dataset, DataLoader from sklearn.model_selection import train_test_split from braindecode.models import EEGNet path /Users/fhe_047/Desktop/机器学习与深度学习/EEGnet/BCICIV_2a_gdf/A01T.gdf raw mne.io.read_raw_gdf( path,#告诉eeg数据在哪 preloadTrue#从硬盘中调用到内存中 ) #由于mne无法区分EEG与EOG所以我们人为将所有眼电标签改为eog raw.set_channel_types({ EOG-left: eog, EOG-central: eog, EOG-right: eog }) # print(raw) # print(raw.ch_names) # print(raw.get_channel_types())#可以检查是否把所有的25个通道全部标为eeg了 # print(raw.info)#获取raw的所有附属信息除时间序列之外 raw_eeg raw.copy().pick(eeg)#将原始数据复制一份然后将eeg标签提取出来不要EOG这样复制也不需要修改原数据 print(raw_eeg) #查看事件要开始切片eeg从一个事件的开始往后1000个采样点也就是4s events, event_id mne.events_from_annotations(raw) #这里面用raw而不是raw_eeg因为原始数据含有event有没有EOG不影响 print(event_id) # 5. 只保留四分类事件 event_id { 769: 7, # 左手 770: 8, # 右手 771: 9, # 双脚 772: 10 # 舌头 } # 6. 根据event切割脑电 epochs mne.Epochs( raw_eeg, # 22通道脑电 events, # 事件位置 event_id, # 四个类别 tmin0, # 从事件开始 tmax4, # 截取4秒 baselineNone, preloadTrue ) # print(epochs.events) # 7. 查看切出了多少个trial print(epochs) # 8. 转换成numpy X epochs.get_data() # 转float32 (修bug) X X.astype(np.float32) y epochs.events[:, -1] - 7 #将每一行数据采用但只保留每行的最后一列类别也就是事件每一行[事件位置, 0, 类别] print(X shape:, X.shape) print(y shape:, y.shape) # # 8. 调整时间长度 # EEGNet需要1000点 # X X[:, :, :1000] 样本数 通道数 时间点 原本 X (288, 22, 1001) print(调整后X:, X.shape) 样本数 通道数 时间点 原本 X (288, 22, 1000) # # 9. 划分训练集和测试集 # X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42, stratifyy #按照y里面类别比例切 ) class EEGDataset(Dataset): def __init__(self, X, y): 初始化 X: EEG数据 y: 标签 self.X X self.y y def __len__(self): 返回样本数量 return len(self.X) def __getitem__(self, index): 根据索引取一个样本 x self.X[index] label self.y[index] return x, label train_dataset EEGDataset( X_train, y_train ) test_dataset EEGDataset( X_test, y_test ) # # 11. DataLoader # train_loader DataLoader( train_dataset, batch_size16, shuffleTrue ) test_loader DataLoader( test_dataset, batch_size16, shuffleFalse ) # 测试一下数据 for x_batch, y_batch in train_loader: print(batch EEG:, x_batch.shape) print(batch label:, y_batch.shape) break # # 12. 创建EEGNet —— 把你学过的网络结构图用代码实例化成一个能运算的对象 # model EEGNet( # n_chans22: EEG通道数电极数量BCI IV 2a数据有22个EEG电极 n_chans22, # n_outputs4: 输出类别数四分类任务左手/右手/双脚/舌头 n_outputs4, # n_times1000: 每个trial的时间点数4秒 × 250Hz 1000个采样点 n_times1000 ) # model变量保存了这个EEGNet模型对象 # 此时模型的骨架已建好卷积层、BN层、激活层等的位置和大小都确定了 # 但参数权重、偏置还是随机初始化的还没经过训练 print(model) # 打印模型内部的所有层Conv2d、BatchNorm、Activation、DepthwiseConv等 # 可以对照你学过的EEGNet结构图验证代码里的层和结构图是否一一对应 # # 13. GPU/CPU —— 确保模型和数据在同一设备上才能计算 # device torch.device( # torch.cuda.is_available(): 检测电脑有没有可用的NVIDIA GPUCUDA # cuda if ... else cpu: 条件表达式三元运算符有GPU返回cuda否则返回cpu cuda if torch.cuda.is_available() else cpu ) # device变量保存了PyTorch的设备对象要么是GPU要么是CPU print(device:, device) # 打印确认用的是哪个设备方便调试 model.to(device) # .to(device): 把模型的所有参数权重、偏置搬到指定设备上 # 【关键】模型和数据必须在同一个设备上如果模型在GPU、数据在CPUPyTorch会报错 # 类比模型和数据要面对面交谈必须在同一个房间设备里 # # 14. 损失函数 —— 衡量模型预测与真实标签之间的差距 # criterion torch.nn.CrossEntropyLoss() # torch.nn.CrossEntropyLoss(): PyTorch内置的交叉熵损失函数类 # 交叉熵是分类任务最常用的损失函数它内部做了两件事 # 1. Softmax把模型输出转成概率四个类的概率加起来1 # 2. 计算负对数似然真实类别对应的概率越大损失越小概率越小损失越大 # 举例模型预测左手概率0.7 → loss -log(0.7) ≈ 0.36小猜得好 # 模型预测左手概率0.05 → loss -log(0.05) ≈ 3.0大猜错了惩罚重 # criterion变量保存了这个损失函数对象后面训练时会反复调用它 # # 15. 优化器 —— 决定怎么调整参数来缩小损失 # optimizer torch.optim.Adam( # model.parameters(): 返回模型中所有可训练参数权重、偏置的迭代器 # 告诉优化器你要管哪些参数优化器只更新这些参数 # EEGNet的卷积层权重、BatchNorm参数等都包含在其中 model.parameters(), # lr0.001: 学习率learning rate控制每次参数更新的步长 # 太大 → 震荡不收敛太小 → 收敛太慢0.001是Adam的默认值新手首选 lr0.001 ) # Adam全称 Adaptive Moment Estimation是最流行、最稳妥的优化器 # 它的聪明之处对每个参数自动调整学习率——梯度大的走小步梯度小的走大步 # 相比基础SGD随机梯度下降Adam收敛更快更稳定 # # 16. 开始训练 —— 最核心的代码反复调整参数让损失降到最低 # # epochs_num 70: epoch 把整个训练集从头到尾过一遍 # 70表示要把训练数据过70遍 # 不是1遍参数需要反复调整才能逐渐收敛到最优 # 不是1000遍太多会过拟合——死记硬背训练数据遇到新数据反而表现差 epochs_num 70 # 外层循环每个epoch模型把所有训练数据过一遍更新一次所有参数 # epoch从0开始计数所以实际是第0~69遍 for epoch in range(epochs_num): model.train() # model.train(): 把模型切换到训练模式 # PyTorch模型有两种模式训练模式(train) 和 评估模式(eval) # 有些层在训练和测试时行为不同 # BatchNorm训练时用当前batch的均值/方差测试时用训练时累积的全局均值/方差 # Dropout训练时随机丢弃一些神经元防过拟合测试时不丢弃全力输出 # model.train()告诉这些层现在是训练阶段按训练规则来 total_loss 0 # 初始化累加变量为0用来累加这一个epoch中所有batch的loss之和 # 最终打印出来看整体趋势loss逐epoch下降模型在学习不降甚至上升有问题 # 内层循环逐个batch处理数据 # train_loader每次吐出一个batch的数据16个样本一组 # 训练数据约230个样本batch_size16所以一个epoch约有230/16 ≈ 14~15个batch # 为什么不一次性处理全部数据因为内存/显存不够分成小batch每次只喂一小口 for X_batch, y_batch in train_loader: # # EEGNet需要 # [batch,1,22,1000] # X_batch X_batch.to(device) # .to(device): 把数据tensor搬到指定设备GPU或CPU # 前面已经把模型搬到了device上数据也要搬过去两者必须在同一设备 y_batch y_batch.to(device) # 标签也要搬到同一设备上和数据、模型保持一致 # 前向传播数据从输入层开始经过EEGNet内部的各个层一步步计算 # Conv2d → BatchNorm → Activation → DepthwiseConv → SeparableConv → ... # 最终输出形状为 [16, 4]16是batch大小4是类别数 # output[0] 第1个样本对4个类的原始分数还没转概率 # 这一步只算一遍不涉及参数更新 output model(X_batch) # 计算loss把模型输出和真实标签一起送给交叉熵损失函数 # 交叉熵内部先Softmax转概率 → 再计算负对数似然 → 对batch内16个样本取平均 # 返回一个标量单个数字就是这个batch的平均损失 # 这个loss值告诉模型你这批数据猜得怎么样差多少 loss criterion( output, y_batch ) # 清空旧梯度 —— 每个batch开始前必须清零 # PyTorch默认是梯度累加不清零的话新梯度会和上一轮的梯度加在一起 # 这会造成参数更新方向错误模型学不好 # 类比考试做错题老师说你错了3分但如果上次扣分没清零你会以为错了6分调整会过头 optimizer.zero_grad() # 反向传播 —— 整个训练的核心步骤之一 # 对loss调用.backward()计算损失函数对模型每个参数的梯度 # 梯度 参数变化一点点时loss变化多少 → 告诉我们往哪个方向调参数能让loss下降 # 使用链式法则从loss出发一层层往回算经过EEGNet的每一层 # 【注意】这一步只计算梯度不更新参数 loss.backward() # 参数更新 —— 整个训练的核心步骤之二 # optimizer.step(): 根据刚才backward()计算出来的梯度按Adam的规则更新每个参数 # 简化公式新参数 旧参数 - lr × 调整后的梯度 # 这一步之后模型的权重就变了——它学到了这个batch的知识 # 【完整训练流程就是三步循环】 # 1. zero_grad() —— 清零旧梯度 # 2. backward() —— 计算新梯度 # 3. step() —— 用梯度更新参数 optimizer.step() # loss.item(): 把loss tensor转成Python普通浮点数脱离PyTorch的计算图 # 为什么用.item()因为loss是PyTorch tensor带着梯度信息不能直接拿来累加 # 直接累加会干扰计算图导致内存泄漏 # .item()把它变成纯纯的数值脱离PyTorch的追踪 total_loss loss.item() # 每个epoch结束后打印loss让你能看到 # - loss是否在下降学习是否有效 # - 下降速度如何 # - 是否收敛loss不再下降 # epoch1是因为epoch从0开始但人类习惯从1开始数 # :.4f指定显示4位小数 print( fEpoch {epoch1}/{epochs_num}, fLoss:{total_loss:.4f} ) # # 17. 测试准确率 —— 训练完成后评估模型在测试数据上的表现 # model.eval() # model.eval(): 把模型切换到评估模式和model.train()对应 # eval模式下 # BatchNorm 用训练时累积的全局均值/方差不是当前batch的 # Dropout 不再丢弃神经元 # 为什么测试时我们想要模型稳定、确定的输出不需要训练时的随机行为 correct 0 # 初始化计数器累计预测正确的样本数 total 0 # 初始化计数器累计总样本数 # 最终 accuracy correct / total with torch.no_grad(): # torch.no_grad(): 告诉PyTorch在这个代码块内不要计算梯度 # 为什么测试时不需要梯度梯度是为了更新参数用的测试时不更新任何参数 # 不计算梯度有两个好处 # 1. 省内存梯度占大量显存不计算就能省下来 # 2. 加速少了反向传播的计算步骤 # 【好习惯】测试时一定要加 torch.no_grad() for X_batch, y_batch in test_loader: X_batch X_batch.to(device) # 测试数据也要搬到同一设备 y_batch y_batch.to(device) # 测试标签也要搬到同一设备 output model(X_batch) # 只做前向传播得到模型对测试数据的预测输出 # output形状: [16, 4]每行是4个类的分数 pred torch.argmax( output, # dim1: 沿着维度1类别维度找最大值的索引 # output形状是 [16, 4]dim1是类别那一维 # 返回形状 [16] 的tensor每个元素是0~3之间的整数 dim1 ) # argmax的意思output每行有4个类的分数比如 [2.1, -0.3, 0.5, -1.0] # argmax找到最大分数的位置 → 0因为2.1最大在第0个位置 # 这就是模型的预测类别我猜这个trial是左手运动 # 为什么argmax而不是softmax后再argmax # 因为argmax只看谁最大softmax只是把数值转成概率最大值的位置不会变 correct ( pred y_batch # pred y_batch: 逐元素比较返回一个布尔tensorTrue/False # 比如16个样本猜对了12个 → 有12个True、4个False ).sum().item() # .sum(): 把True(1)和False(0)加起来得到预测正确的数量 # .item(): 把PyTorch tensor转成Python普通整数 total y_batch.size(0) # y_batch.size(0): 获取tensor第0维的大小就是batch大小16 # 累计测试集的总样本数 accuracy correct / total # 除法得到浮点数比例比如0.45表示45%的测试样本预测正确 # 四分类随机猜测准确率是25%所以45%已经比随机好了但还有提升空间 # BCI IV 2a四分类的典型EEGNet准确率大约在60%~75%之间 # 只用单个受试者数据288个trial可能偏低多受试者数据会更好 print( Test Accuracy:, accuracy )

相关新闻

qemu-img工具在KVM环境中的高级应用指南

qemu-img工具在KVM环境中的高级应用指南

qemu-img工具在KVM环境中的高级应用指南 在KVM虚拟化技术生态中,qemu-img作为核心磁盘管理工具,承担着虚拟磁盘创建、转换和维护的关键角色。本文将深入探讨该工具的高级功能及其在复杂虚拟化场景中的应用实践。 一、磁盘格式转换与优化 1.1 多格式支持与…

2026/7/23 21:29:58 阅读更多 →
报会计备考课有哪些坑?6个真实套路拆解,附省时上岸的工具选法

报会计备考课有哪些坑?6个真实套路拆解,附省时上岸的工具选法

会计职称考试的本质是合格性考试——60分万岁,不是选拔性高分竞赛。但恰恰是这种"看起来不难"的考试,让无数在职考生、宝妈、零基础小白栽在"选错课"上:钱花了大几千,课听了不到三分之一,最后要么…

2026/7/23 21:29:57 阅读更多 →
三亚商务办公室装修推荐

三亚商务办公室装修推荐

在三亚进行商务办公室装修,专业性至关重要。以下为大家分析相关痛点及优质推荐。痛点一:爆雷风险2026 年很多装修公司存在爆雷风险。一些小公司可能资金链不稳定,施工到一半就停工,让业主损失惨重。比如有的公司收了预付款后&…

2026/7/23 21:28:57 阅读更多 →

最新新闻

【剪映AI配音实战指南】:20年音视频专家亲授,3步搞定专业级AI配音(附避坑清单)

【剪映AI配音实战指南】:20年音视频专家亲授,3步搞定专业级AI配音(附避坑清单)

更多请点击: https://intelliparadigm.com 第一章:剪映AI配音的核心原理与技术演进 剪映AI配音并非简单的语音合成(TTS)工具,而是融合了端到端深度学习、音色克隆、韵律建模与上下文感知语音生成的多模态语音系统。其…

2026/7/23 21:37:00 阅读更多 →
AI论文写作工具评测与学术诚信实践指南

AI论文写作工具评测与学术诚信实践指南

1. AI论文写作工具现状与选择逻辑学术写作正经历着从传统人工创作向智能辅助的转型期。当前市面上的AI论文工具主要分为三大技术流派:基于Transformer架构的通用大语言模型(如GPT系列)、针对学术领域优化的专业模型(如SciBERT&…

2026/7/23 21:37:00 阅读更多 →
文献综述自动化工具:AI如何提升学术研究效率

文献综述自动化工具:AI如何提升学术研究效率

1. 项目概述:文献综述自动化的痛点与突破写文献综述大概是每个研究生都经历过的噩梦。去年帮导师整理某个细分领域的研究进展时,我花了整整三周时间:先要检索上百篇论文,然后逐篇阅读摘要、筛选关键文献,再手动归类不同…

2026/7/23 21:37:00 阅读更多 →
非标准分带CAD数据如何与卫星影像叠加套合

非标准分带CAD数据如何与卫星影像叠加套合

自2008年7月1日起,国内测绘行业统一采用 2000 国家大地坐标系(CGCS2000),逐渐淘汰西安80和北京54坐标系。通常情况下,国家2000坐标系下的文件都基于标准的3度分带或6度分带进行投影。但为了提高坐标的精度,…

2026/7/23 21:37:00 阅读更多 →
Django毕设项目:基于 Django 的电子商务交易系统构建与安全防护方案研究 基于 Django 的商城系统设计与漏洞防范策略研究 (源码+文档,讲解、调试运行,定制等)

Django毕设项目:基于 Django 的电子商务交易系统构建与安全防护方案研究 基于 Django 的商城系统设计与漏洞防范策略研究 (源码+文档,讲解、调试运行,定制等)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围:&am…

2026/7/23 21:37:00 阅读更多 →
一次全站HTTPS证书过期故障的复盘:从应急响应到自动化证书管理体系的建立全过程

一次全站HTTPS证书过期故障的复盘:从应急响应到自动化证书管理体系的建立全过程

一次全站HTTPS证书过期故障的复盘:从应急响应到自动化证书管理体系的建立全过程 一、故障概述与影响评估 2025年8月12日凌晨02:47,某互联网企业的核心业务平台突然发生大规模服务不可用故障。用户访问官网、API网关、管理后台等所有HTTPS服务时&#xff…

2026/7/23 21:36:00 阅读更多 →

日新闻

从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表)

从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表)

更多请点击: https://intelliparadigm.com 第一章:从单点好评到指数级传播:AI副业主理人必须掌握的4层口碑渗透模型(含ROI测算表) 当AI副业主理人不再仅满足于单次服务交付,而是主动构建可复用、可裂变、可…

2026/7/23 0:00:25 阅读更多 →
AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析

AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析

更多请点击: https://codechina.net 第一章:AI写作开头钩子设计:为什么你的AI文案完读率不足18%?——基于2,346篇A/B测试报告的归因分析 在对2,346篇跨行业AI生成文案的A/B测试数据进行聚类分析后,我们发现&#xff1…

2026/7/23 0:01:26 阅读更多 →
Chitchatter完整指南:免费开源的终极点对点安全聊天工具

Chitchatter完整指南:免费开源的终极点对点安全聊天工具

Chitchatter完整指南:免费开源的终极点对点安全聊天工具 【免费下载链接】chitchatter Secure peer-to-peer chat that is serverless, decentralized, and ephemeral 项目地址: https://gitcode.com/gh_mirrors/ch/chitchatter Chitchatter是一款革命性的安…

2026/7/23 0:01:26 阅读更多 →

周新闻

Go语言静态资源打包方案对比与实践指南

Go语言静态资源打包方案对比与实践指南

1. 项目背景与核心需求在Go语言开发中,我们经常需要处理静态资源文件的打包问题。无论是Web应用的模板文件、前端资源,还是配置文件、证书等,都需要随程序一起分发。传统做法是将这些文件与编译后的二进制文件放在同一目录下,但这…

2026/7/22 8:58:19 阅读更多 →
Go语言实现高性能LDAP认证服务的架构与实践

Go语言实现高性能LDAP认证服务的架构与实践

1. 项目背景与核心价值LDAP(轻量级目录访问协议)作为企业级身份认证的黄金标准,已经服务了超过80%的财富500强公司。我在金融科技领域实施统一认证体系时,发现传统Java方案存在启动慢、内存占用高等痛点。而Go语言凭借其协程并发模…

2026/7/22 19:43:43 阅读更多 →
【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

更多请点击: https://intelliparadigm.com 第一章:AI面试官实战指南的核心价值与适用场景 AI面试官并非替代人类HR的“黑箱工具”,而是以可解释、可审计、可迭代的方式,赋能招聘全链路的关键基础设施。其核心价值在于将主观经验沉…

2026/7/23 17:49:47 阅读更多 →

月新闻