spikingjelly
PyTorch原生脉冲神经网络框架,支持ANN2SNN转换与神经形态硬件部署
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch原生脉冲神经网络框架,支持ANN2SNN转换与神经形态硬件部署
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下,当你用普通ANN(人工神经网络)识别一张猫的图片时,整个网络在同一时刻对所有像素做计算——无论这张图片是静态快照还是来自动态事件相机(Event Camera)拍摄的真实场景,ANN 都以相同的节拍处理。这种"同步驱动"的计算模式在面对高帧率视频、低功耗边缘设备或需要实时响应的场景时,往往力不从心。
脉冲神经网络(Spiking Neural Network,SNN)则完全不同。SNN 的神经元只有在收到足够强的刺激时才会"放电"(产生脉冲),这与真实生物神经元的工作方式高度相似。这种事件驱动(event-driven)的特性使得 SNN 在理论上具备极高的能效比——理论上,SNN 的功耗可以比同等精度的 ANN 低 100-1000 倍。正因如此,SNN 被认为是下一代神经形态计算的核心。
然而长期以来,SNN 的研究面临一个尴尬困境:没有好用的工具。研究者要么从零手写神经元模型,要么在通用深度学习框架里艰难地模拟脉冲行为。SpikingJelly 的出现,正是为了解决这个问题——它让 SNN 研究者能像使用 PyTorch 一样自然地构建、训练和部署脉冲神经网络。
SpikingJelly 由北京大学机器学习实验室(PKU MLG)和 PCL(鹏城实验室)联合开发维护。项目的核心作者来自北京大学智能科学系,最早在 2020 年发表相关研究论文(ICCV 2021),随后持续迭代至今。
项目自诞生起就坚持 PyTorch 原生的设计哲学——开发者不需要学习全新的 API 体系,只要熟悉 PyTorch,就可以直接上手 SpikingJelly。这种设计选择大大降低了 SNN 研究的入门门槛。框架目前已积累了 2000+ 引用的学术论文支撑,是 SNN 领域引用量最高的开源框架之一。
SpikingJelly 的设计围绕 SNN 研究者的全链路需求展开:
友好的神经元建模
开发者可以像写 PyTorch 代码一样定义脉冲神经元。框架提供了丰富的内置神经元类型:LIF(Leaky Integrate-and-Fire)、PLIF(Parametric LIF)、IF 等经典模型一应俱全。每个神经元都支持可学习的膜时间常数(membrane time constant),这正是 2021 年 ICCV 论文的核心贡献——让神经元自己学会最优的时间常数,而不是靠人工调参。
核心代码示例:定义一个 SNN 与定义 PyTorch 模型完全一致,通过 layer.Linear 和 neuron.LIFNode 快速构建网络:
from spikingjelly.activation_based import layer, neuron, surrogate
net = nn.Sequential(
layer.Flatten(),
layer.Linear(28 * 28, 10, bias=False),
neuron.LIFNode(tau=2.0, surrogate_function=surrogate.ATan())
)
多后端加速
训练大规模 SNN 的计算量不容小觑。SpikingJelly 支持三种计算后端:原生 torch、基于 CUDA 的 CuPy,以及英伟达的 Triton。用户可以在运行时切换后端,无需修改模型代码。在实际测试中,CuPy 后端相比纯 PyTorch 可获得数倍速度提升。此外,框架完全兼容 torch.compile,在最新 PyTorch 2.7 上验证通过。
ANN2SNN 转换
这是框架最实用的功能之一——可以将训练好的传统 ANN(卷积神经网络等)自动转换为等效的 SNN。这解决了 SNN 训练难度大的痛点:先在 ANN 体系下用成熟方法训练,再转换成能效更高的 SNN 部署到硬件。框架支持 CVPR、ICLR 等顶会论文中的前沿转换算法(如 QCFS 转换、膜电位残差补偿等)。
事件数据集与数据处理
SNN 天然适合事件相机(Event Camera)数据。SpikingJelly 内置了 DVS128 Gesture、CIFAR10-DVS 等常用事件数据集的数据加载和预处理流程,并支持通过 Neuromorphic Dataset 统一接口访问。框架还提供了完整的数据增强工具,帮助提升模型的泛化能力。
硬件部署支持
框架提供了 NIR(Neural Intermediate Representation)、Lava 和 Lynxi 芯片的导出接口,可以将训练好的 SNN 模型转换为面向神经形态硬件的中间表示,实现低功耗边缘部署。这是连接学术研究与实际硬件部署的关键桥梁。
SpikingJelly 的代码结构清晰,主模块包括:
核心训练逻辑完全兼容 PyTorch 生态:使用 torch.optim、torch.utils.data 和 torch.compile,支持分布式训练(DDP)和混合精度训练。框架还有专门的显存优化模块(memopt),通过梯度检查点(Gradient Checkpointing)和脉冲压缩技术,实现无损的低显存训练——这是训练深层 SNN 的关键技术。
SpikingJelly 支撑了 200+ 篇学术论文,覆盖 CVPR、NeurIPS、ICLR、ICCV、IJCAI 等顶会。核心算法包括:深度残差 SNN(Spiking ResNet)、SNN 剪枝(Gradient Rewiring)、并行脉冲神经元等。框架本身也在持续跟进 SNN 领域的前沿进展(如 Spikformer、SpikeGPT 等新型架构)。
SpikingJelly 的最大优点是 PyTorch 原生 API,零学习成本——任何熟悉 PyTorch 的开发者都可以无缝迁移。文档极为详尽,提供中英文双语 ReadTheDocs 和完整教程。缺点在于:纯命令行工具,没有 Web UI,对非编程用户不友好;需要 CUDA 环境才能发挥最佳性能;SNN 训练收敛比 ANN 更困难,需要一定的专业知识。
随着神经形态硬件(Intel Loihi、英伟达 SNN 芯片、国内的 Lynxi 等)逐渐走向成熟,SNN 的实际部署条件正在改善。SpikingJelly 作为最成熟的 PyTorch 原生 SNN 框架,既是学术研究的利器,也是未来边缘 AI 部署的重要基础设施。随着事件相机和低功耗 AI 芯片的普及,基于 SNN 的高效计算范式将越来越重要。