EMS-YOLO:深度直接训练脉冲SNN实现目标检测的代码解析
2026/9/15 14:20:43 网站建设 项目流程

简介:深度直接训练脉冲神经网络EMS-YOLO的官方实现代码包源自ICCV 2023论文,面向脉冲视觉、低功耗目标检测以及神经形态计算方向的算法研究者与工程人员,可用于复现论文实验、开展算法对比或二次开发。压缩包体积约410KB,共60个文件,其中36个Python脚本构成训练、验证、检测与工具链主体,15个YAML文件提供模型结构、超参数及数据配置,另有shell脚本、依赖清单与说明文档,便于快速搭建环境。资源内覆盖COCO和Gen1数据集的处理与训练入口,包含预训练检查点自动下载、检测推理示例、能耗计算工具,并提供多种脉冲残差网络结构变体;借助environment.yml可一键复制conda环境,降低复现门槛。当前已有97人学习,适合具备一定深度学习基础并希望切入脉冲神经网络目标检测的读者深入研究。

1. 从 EMS-YOLO 看深度直接训练脉冲 SNN 的检测落地路径

对象检测领域里,脉冲神经网络(SNN)的大多数成果仍停留在分类任务,原因不外乎两层:直接训练深层 SNN 需要在时间维上展开 BPTT,计算代价远高于 ANN;检测任务又要求多尺度空间特征与密集预测头,这些结构与脉冲神经元基于离散事件的动力学天然不兼容。ICCV 2023 放出的这个 spiking-yolo 相关正式实现(EMS-YOLO)把问题拆成了两步——保留 YOLO 的检测头,把 backbone 换成可深度直接训练的脉冲 ResNet,并在 COCO 与 Gen1 事件相机数据集上给出了完整的 train/val/detect 流程。适合想复现深度直接训练 SNN 跑目标检测的算法工程师,也适合做事件相机感知的人,仓库里除了训练与推理代码,还有事件数据的组织脚本和脉冲发放率计算工具,改到自己的数据集上比从零搭管线要省事得多。

2. 代码库结构与脉冲化 backbone:resnet 系列 yaml 怎么选

2.1 根目录文件布局:训练管线一眼看穿

仓库根目录基本上沿用了 YOLOv5 的工程结构,改动集中在 models 下的脉冲 backbone 和 g1-resnet 子目录。先用一张表把主要文件跑一遍,后面找代码时能少走弯路。

文件/目录职责
train.py / val.py / detect.py训练、验证、推理三个主入口
data/coco.yaml、data/gen1.yamlCOCO 与 Gen1 数据集路径、类别、下载配置
models/*.yamlbackbone 结构定义,resnet10/18/34 等变体
models/yolo.py、common.py模型组装逻辑与脉冲算子实现
utils/loss.py、datasets.py检测损失、数据加载与增强
runs/train训练输出目录,opt.yaml、hyp.yaml、results.png 都在这里
scripts/get_coco.shCOCO 数据下载与解压脚本
g1-resnetGen1 事件数据集训练与帧率分析补充代码

这段布局的关键在于:检测头(YOLO head)没有动,动的是 backbone 的每个卷积块。也就是说,YOLO 原有的 anchor、损失、NMS 逻辑可以直接复用,脉冲化改造被约束在models/内,这对想要做 SNN 检测算法对比的人来说是非常友好的实验框架。

2.2 五个 backbone yaml 的差异:SEW、EE 和基础 ResNet

models 目录下除了 resnet10.yaml、resnet18.yaml、resnet34.yaml 之外,还有 res18-sew.yaml、res18-ee.yaml、res18-eebk.yaml、resnet34-cat.yaml。第一次打开仓库的人很容易被这些命名绕晕,实际从名字能看出三种不同的脉冲化策略。

基础 resnet18.yaml / resnet34.yaml 是用最简单的脉冲卷积替代普通卷积,每个残差块里的 ReLU 变成脉冲神经元发放。SEW(Spike-Element-Wise)来自另一条技术路线:对膜电位做逐元素操作时,如果两个输入分支都是脉冲信号,加法会让脉冲幅值膨胀,SEW 用 element-wise 操作替代加法来缓解脉冲信息在残差路径上的衰减。EE 通常指输入侧的编码策略调整,事件流被组织成适合脉冲网络处理的时间面片(time surface),eebk 从命名习惯上看是 EE 策略在 backbone 上的配套变体。resnet34-cat 则是把特征融合从逐元素相加改为通道拼接,类似 ResNet 里 1x1 卷积降维后的 concat 路径。

这些 yaml 结构上都是 A 类配置,用一个简单脚本就能对比 backbone 长度和输入通道差异:

import yaml files = [ "models/resnet18.yaml", "models/resnet18.yaml", "models/res18-sew.yaml", "models/res18-ee.yaml", ] for path in files: with open(path) as fp: cfg = yaml.safe_load(fp) backbone = cfg["backbone"] # 每个元素是 [卷积核数, 输出通道数, 模块类型, 参数列表] 这类结构 print(path, "stem通道:", backbone[0][1], "block数:", len(backbone))

这里打印的 stem 通道数和 block 数只做结构校验,真正决定 SEW 与 EE 差异的是 common.py 里脉冲残差块的实现。yaml 里通常不会显式写“SEW”字样,而是复用同一个 block 类型名,然后在 yaml 的深层参数里标记不同的神经元类型、膜电位时间常数或编码方式,因此改结构前最好先打开对应 yaml 确认 backbone 最后一个元素里的参数列表。

2.3 yolo.py 如何把脉冲 backbone 与 YOLO 头组装

models/yolo.py 里的 Model 类读入 yaml 后,通过 parse_model 把字符串映射到 common.py 或 experimental.py 的类。SNN 版本的关键变化是前向传播:普通 CNN 一次前向得到特征图,SNN 需要对 T 个时间步循环模拟脉冲发放,再把时间维上的特征聚合出来。

# 伪代码:SNN backbone 多时间步前向的常见写法 def forward_spiking(net, x, T=4): # x: [B, T, C, H, W] 事件帧,或把静态图重复 T 次 batch, steps, _, _, _ = x.shape feats = [] for t in range(steps): out = net.backbone(x[:, t]) feats.append(out) # 可选:按时间求平均,或做注意力加权 return sum(feats) / steps

T 是时间步长,静态 COCO 图像可以设 T=1 到 2,Gen1 事件数据通常用 T=4 到 8。runs/train 里的 opt.yaml 会记录训练时的 T、batch、img size 等参数,复现实验结果前先看这个文件比翻 README 更靠谱。如果 T 设得过大,训练显存占用会等比上涨,batch size 就要相应下调,否则 BN 统计不稳定。

3. 环境复现与 COCO/Gen1 数据准备

3.1 conda 环境精准复刻:environment.yml 与 requirements.txt 的分工

原仓库在 environment.yml 里锁定了 pytorch=1.10.1、python=3.8、cuda=11.3、cudnn=8.2.0_0,也就是说 conda 会创建一个完整可用的 PyTorch 环境。创建命令很直接:

git clone https://github.com/BICLab/EMS-YOLO.git cd EMS-YOLO conda env create -f environment.yml # 创建完成后先确认环境名 conda env list conda activate <环境名> pip install -r requirements.txt

environment.yml 负责带 cuda/cudnn 的编译型依赖,requirements.txt 负责剩余纯 Python 包。手动复刻时如果本机 CUDA 版本不同,我一般会把 environment.yml 里的 torch 安装行替换成对应 CUDA 版本的安装命令,但这么做的风险是 cudnn 版本与训练脚本里隐式调用的算子可能对不上,出现“能 import 但前向就崩”的情况。先跑一段 2 batch 的 train 命令验证环境,比直接启动完整训练要快得多。

文件负责内容常见包举例
environment.yml锁定 Python、PyTorch、CUDA、cuDNNpython=3.8, pytorch=1.10.1, cudnn=8.2.0_0
requirements.txt其余运行依赖numpy, opencv-python, matplotlib, seaborn 等

3.2 下载 COCO:get_coco.sh 会做什么

bash scripts/get_coco.sh

脚本通常会从 COCO 官方地址下载 train2017、val2017 图片与 annotations 标注,并解压到脚本内预设的 data/coco 目录。运行前先确认磁盘剩余空间在 20GB 以上,否则下到一半中断后,较常见的 wget 断点续传问题会浪费很多时间。下载完成后打开 data/coco.yaml,确认 path 字段指向实际数据根目录,类别数 80 与 nc 字段一致。

COCO 图片是静态帧,输入给 SNN 时只算单时间步的重复输入;但如果原实现里对静态图做了多步重复输入,训练成本会成倍增加。跑 COCO 之前先看一眼 detect.py 或 train.py 里对输入张量形状的预处理,确认是否把 [B, C, H, W] 扩展成了 [B, T, C, H, W]。

3.3 Gen1 事件数据:从事件流到可训练样本

Gen1 是 Prophesee 发布的汽车、行人、自行车检测数据集,原始数据是事件流,不能直接交给卷积层。g1-resnet 目录里的 give_g1_data.py 负责把事件累计成固定时间窗口下的 2D 帧,calculate_fr.py 负责统计脉冲发放率。事件按极性分成两通道,与普通 RGB 三通道的输入维度不同,所以 data/gen1.yaml 里的 ch 字段需要单独配置。

我一般会先用 give_g1_data.py 处理小段数据验证输出尺寸,再用 datasets.py 的加载逻辑检查标签坐标是否与帧分辨率对应。Gen1 标签的时间戳是毫秒级,事件帧窗口的时间跨度要与标签对齐,否则训练时会出现标签明显错位而不收敛的现象。这部分细节在 g1-resnet 的 val.py 里也有对应校验逻辑。

4. 训练-推理-评估的关键环节:hyp.yaml、opt.yaml、loss 与 detect

4.1 训练入口与超参配置文件

训练主入口沿用了 YOLO 系列的写法,常见启动命令是:

python train.py \ --data data/coco.yaml \ --cfg models/resnet34.yaml \ --weights "" \ --batch-size 32 \ --img 640 \ --epochs 120

train.py 启动后会把本次启动参数写入 runs/train/exp/opt.yaml,这个文件是回看实验配置的第一手资料。hyp 超参则在 hyp.scratch.yaml 这类文件里维护,包括 lr0、momentum、weight_decay、mosaic 增强开关等。SNN 直接训练对 batch size 比普通 CNN 更敏感,底层脉冲计算方差大,我一般要求 batch size 不低于 16,否则 BN 统计噪声会让 loss 曲线出现明显震荡。

运行中查看结果用 tensorboard 或直接看 results.png:

tensorboard --logdir runs/train

results.png 每轮更新,包含 box_loss、cls_loss、mAP 等曲线。如果 box_loss 下降但 mAP 平稳不涨,先检查 anchor 是否重新计算过,仓库里 utils/autoanchor.py 有自动锚框的逻辑,数据集的标签尺寸分布与 COCO 差异大时要在训练前跑一次。

4.2 损失函数与深度直接训练的代理梯度

utils/loss.py 里保留了 YOLO 的完整损失结构,包含 box 回归损失、分类损失和对象置信度损失。SNN 深度直接训练的关键区别不在损失函数,而在反向传播路径上:脉冲神经元“发放或不发放”的阶跃函数不可导,直接训练必须用代理梯度(surrogate gradient)来近似。

# 通用做法:矩形代理梯度替代阶跃函数导数 def surrogate_gradient(x): # x 为膜电位减去阈值的差值 return (x.abs() < 0.5).float()

这个梯度函数在反向传播时替代真实的脉冲导数,让梯度能够穿过时间步回传。实现里通常会把它封装在神经元的 backward 逻辑中,训练时无需手动干预。判断代理梯度是否生效的最快方法:在 train.py 里打印第一层骨干的梯度范数,如果所有梯度严格为 0,说明代理梯度没有被正确挂到脉冲神经元的 forward-backward 上。

loss 数值的合理性也可以快速验证:

python train.py --weights "" --batch-size 2 --epochs 1

batch size 设为 2 不是正常训练,但能验证 loss 是 nan 还是有限值。如果第一步就出现 nan,优先检查学习率初始值和学习率预热逻辑,SNN 直接训练的梯度幅度比 ANN 小,但偶尔会出现突刺。

4.3 推理与验证:detect.py 自动下载权重,val.py 出指标

detect.py 支持图片、视频、目录等来源,并会自动下载 COCO 上训练的 EMS-ResNet34 权重:

python detect.py \ --weights COCO_EMS-ResNet34.pt \ --source data/images \ --img 640 \ --conf-thres 0.25

val.py 用于评估:

python val.py \ --data data/coco.yaml \ --weights runs/train/exp/weights/best.pt \ --img 640

detect.py 的常用参数及其作用如下表,参数命名与仓库内脚本保持一致,方便直接照抄。

参数作用
--weights权重路径,可接收仓库自动下载的 .pt 文件
--source推理来源,支持图片路径、目录、视频文件
--img推理分辨率,需与训练一致,否则检测性能下降
--conf-thres置信度阈值,调低会增多漏检变少但误检变多
--iou-thresNMS 的 IoU 阈值,目标密集场景适当调高
--device指定 CPU 或 GPU 编号,不设则自动选择
--save-txt输出 txt 格式检测结果,便于脚本处理

5. 迁移到 Gen1 的实战技巧:give_g1_data 与 calculate_fr

5.1 把 Gen1 数据接进训练管线

gen1 数据集训练需要把 g1-resnet 里的文件按 README 说明替换或添加到根目录对应位置。官方给出的命令是:

python path/to/train_g1.py --weights ***.pt --img 640

train_g1.py 的入口参数比 train.py 少一些,weights 可以填 COCO 预训练模型作为初始化,也可以填空字符串从头训练。迁移训练时我建议从 COCO 权重开始,因为 SNN backbone 的低层特征对边缘和运动方向的响应是通用的,事件帧里的极性与灰度图的边缘有强相关性,用预训练初始化能显著减少 Gen1 上收敛所需 epoch。

give_g1_data.py 负责原始事件文件的解析与帧化,处理完的数据路径要与 data/gen1.yaml 里的 train/val 字段对应。Gen1 的标签是 RLE 或 bounding box 格式,转成 YOLO 的归一化 xywh 时要注意坐标除以的分辨率是事件帧的分辨率,不是原始事件流的分辨率,二者不一致是迁移训练最常踩的坑。

5.2 用 calculate_fr 验证脉冲网络是否退化

calculate_fr.py 统计各层脉冲发放率 FR(firing rate)。FR 过低意味着神经元长期不发放,梯度信号在时间维上迅速衰减;FR 过高则脉冲近似于连续激活,SNN 的时间特性被弱化。通常当 FR 落在 0.05~0.5 区间时网络处于较健康的发放状态。

# 通用做法:逐层统计平均发放率 def layer_fr(model, sample, T): fr = {} name = "spiking_conv" for t in range(T): out = model(sample, timestep=t, save_spikes=True) # 假设模型内部把每层脉冲计数写入 output_spikes for module_name, spikes in model.output_spikes.items(): if name in module_name: fr[module_name] = spikes.float().mean().item() return fr

如果某层 FR 长期接近 0,先降低膜电位阈值或提高初始化权重幅度,常见做法是把该层的时间常数调大;如果 FR 持续大于 0.8,检查输入事件帧是否过密、时间窗口是否过长。窗口长度直接影响每帧事件数量,窗口越大事件越密,FR 越高,适当缩短窗口能同时提升推理帧率。以 calculate_fr.py 输出的均值 FR 作为是否重采样时间窗口的依据,比肉眼调参靠谱得多。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询