☰
PyTorch从零实现脉冲神经网络SNN与STDP学习规则
2026/9/29 18:46:49 网站建设 项目流程

1. 先弄明白SNN和STDP在解决什么问题

1.1 脉冲神经网络与传统人工神经网络的本质区别

脉冲神经网络(Spiking Neural Network, SNN)这几年在我视野里出现的频率明显变高了。原因其实不复杂:传统的深度神经网络虽然能力强,但反向传播需要全局梯度、需要精确的张量运算,这在GPU上是优势,到了低功耗场景就成了负担。而SNN走的是另一条路——信息不是用连续浮点数表示的,而是用“有没有脉冲”、“脉冲在什么时间出现”来表达,神经元也只有发放和不发放两种状态。这种离散事件驱动的计算方式,天然更接近生物神经系统,也更容易映射到脉冲芯片上,功耗可以做到极低。

但SNN的难点也很直接:网络里的操作很多是不可导的,传统那套BP训练方法不能直接套用。于是就有了一批替代方案,比如代理梯度(Surrogate Gradient)、ANN2SNN转换、无监督学习方法等。而STDP(Spike-Timing-Dependent Plasticity,脉冲时间依赖可塑性)是其中最有生物依据、也最优雅的一条无监督学习路线。

STDP的核心思想很朴素:突触前神经元A的脉冲如果总是比突触后神经元B的脉冲早到一点点,那A对B的连接就会变强;反过来,如果A的脉冲总是落在B的脉冲之后,那这条连接就会变弱。一句话,“紧跟其后的输入会被记住,姗姗来迟的输入会被遗忘”。这个规则不需要任何标签,不需要梯度回传,只需要局部的脉冲时序信息就能更新权重。

1.2 为什么值得用PyTorch手写一遍STDP

有人可能会问,现在有SpikingJelly、BindsNET、Norse这些现成的SNN框架,直接用不就行了?我当时也是这么想的,后来被一个特殊需求卡住了,才意识到框架封装的便利性有时候反而会挡住原理。写这个实战项目的目的,就是为了把LIF神经元模型、STDP更新公式、事件驱动的权重调整过程全部打开揉碎,用最基础的PyTorch张量操作从零实现一遍。

这个项目包含三个层次:第一层是用PyTorch实现LIF(Leaky Integrate-and-Fire)神经元,这是SNN最常用的计算单元;第二层是实现了基于Trace机制的STDP学习规则,这是整个项目的核心;第三层是把它组装成一个完整的小网络,让它学会区分两类不同的脉冲模式。整个项目CPU就能跑,代码全部贴在正文里,基本上照着抄一遍就能跑通。

1.3 这套方案适合谁,能学到什么

如果你是刚刚接触SNN,想搞清楚“脉冲神经网络到底是怎么训练的”,这个项目是一个非常合适的切入点。如果你已经用过SpikingJelly这类框架,但不知道底层权重是怎么更新的,这篇文章也能帮你补上这块拼图。如果你是做事件相机数据、神经形态芯片、时序信号处理相关方向的,STDP这套无监督局部学习机制对你后续理解相关工作会有直接帮助。

我写代码的时候会刻意把每个超参数、每个矩阵维度的来源都讲清楚,方便你改成自己的数据格式。掌握这套东西之后,再去扩展卷积SNN、侧抑制、奖励调节STDP(如R-STDP)等复杂机制,你会轻松很多。

2. 环境准备与整体设计:把项目拆成三层

2.1 PyTorch环境怎么搭最省事

这个项目只需要PyTorch,不需要任何特殊硬件。我自己就是在普通笔记本的CPU上跑的,200个回合训练几十秒就完事。如果你还没装PyTorch,用Anaconda是最省事的方式:

conda create -n snn python=3.10 conda activate snn conda install pytorch cpuonly -c pytorch

GPU版本的同学就换成官网给出的对应命令,本质上一样。装完验证一下:

import torch print(torch.__version__)

能输出版本号就行。这个项目用到的东西极其基础,torch.Tensor、torch.nn.Module、torch.clamp,没有分布式、没有自动求导层面的花活,所以对PyTorch版本没有强制要求。

2.2 整体架构:编码层、LIF神经元层、STDP更新器

这个项目我把它拆成三个模块,每个模块只干一件事:

  • 输入编码模块:把某种外部信息编码成脉冲序列,形状是(T, N_in),T是模拟的时间步数,N_in是输入神经元个数。本项目里我直接构造了手工设计的脉冲模式,用来模拟两类不同的时序事件。
  • LIF神经元层:接收输入脉冲,维护膜电位状态,超过阈值就发放脉冲。它的代码要完成“积分-发放-重置”的完整闭环。
  • STDP更新器:在每个时间步根据突触前和突触后脉冲的时间关系更新权重。

这样拆的好处是,后续如果你想换成更复杂的神经元模型,或者想换一种学习规则,只需要替换其中一层,其他部分不用动。

2.3 要提前定好的关键超参数

在写代码前,超参数必须想清楚,SNN里这些参数的耦合关系比DNN要强很多,随意设一个数很可能导致网络“死掉”或者“刷屏”。这个项目最终用到的主参数如下:

参数取值含义
T50单个样本模拟的时间步数
N_in40输入神经元数量
N_out2输出神经元数量
tau20膜电位时间常数,单位为时间步
V_th1.0发放阈值
dt1.0模拟步长
tau_pre20突触前Trace的衰减时间常数
tau_post20突触后Trace的衰减时间常数
A_plus0.1STDP增强幅度
A_minus0.12STDP抑制幅度
lr0.01权重更新学习率

后面我会逐个解释这些参数的物理意义和调参方向,你不需要死记硬背,但建议先照着跑通,再改着玩。

3. 从零实现LIF神经元:SNN的计算核心

3.1 从RC电路理解LIF模型

LIF模型的全称是Leaky Integrate-and-Fire,翻译过来就是“带漏电的积分-发放模型”。你可以把神经元想象成一个底部有洞的水桶:

  • 输入脉冲相当于往桶里倒水;
  • 膜电位就是桶里的水位;
  • “漏电”指的是这个桶会持续漏水,漏水速度和水位成正比;
  • 当水位超过某个高度(阈值),神经元就发一次脉冲,然后桶被清空(重置)。

这个物理过程对应一条微分方程:

τ_m * dV/dt = -(V - V_rest) + R * I(t)

其中τ_m是膜时间常数,V_rest是静息电位,I(t)是输入电流。电脑里没法直接算微分方程,所以要离散化。最常用的做法是用指数衰减来近似:

V_t = V_{t-1} * exp(-dt / τ_m) + I_t

这里exp(-dt/τ_m)就是代码里的decay。注意时间常数τ越大,decay越接近1,相当于“漏水越慢,记忆越长”。用时间步做单位时,τ=20意味着膜电位约在20个时间步后衰减到原来的1/e左右。

3.2 离散化推导:为什么是 V = decay * V + input

我把上式的来源展开一下,这样你后面调参时心里有底。对RC电路模型做欧拉法展开可以得到:

V_t = V_{t-1} + (dt / τ_m) * (-(V_{t-1} - V_rest) + R * I)

整理后令decay = (1 - dt/τ_m),在dt足够小时,这个形式和指数形式几乎是等价的。但我会用exp(-dt/τ_m)作为decay,原因是指数形式在dt较大时数值更稳定,不会出现dt/τ_m > 1导致decay变成负数这种离谱情况。

代码实现很简单:

decay = torch.exp(-dt / tau) # 约等于 0.9512 V = decay * V + input_current

然后把超阈值部分变成脉冲,把发放过的神经元电压归零。生物神经元还有个“不应期”,也就是发放后短时间内无法再次发放,这里先不加,后面在常见问题里会单独说。

3.3 PyTorch代码实现LIF并踩掉两个坑

下面是我实际使用的LIF网络层代码,我把权重、膜电位、Trace都放在一个类里,这样训练逻辑可以集中处理:

import torch import torch.nn as nn class SNNLayer(nn.Module): def __init__( self, n_in, n_out, tau=20.0, V_th=1.0, dt=1.0, lr=0.01, A_plus=0.1, A_minus=0.12, tau_pre=20.0, tau_post=20.0 ): super().__init__() self.n_in = n_in self.n_out = n_out self.tau = tau self.V_th = V_th self.dt = dt self.lr = lr self.A_plus = A_plus self.A_minus = A_minus self.tau_pre = tau_pre self.tau_post = tau_post self.decay = torch.exp(-dt / tau) self.decay_pre = torch.exp(-dt / tau_pre) self.decay_post = torch.exp(-dt / tau_post) # 权重设计为 (n_out, n_in),不需要梯度,STDP直接手动更新 self.weight = torch.rand(n_out, n_in) * 0.15 + 0.05 def reset(self): self.V = torch.zeros(self.n_out) self.pre_trace = torch.zeros(self.n_out, self.n_in) self.post_trace = torch.zeros(self.n_out) def forward(self, input_spikes): """ input_spikes: (T, n_in) 的张量,值为0或1 返回: (T, n_out) 的输出脉冲序列 """ T = input_spikes.size(0) out_spikes = torch.zeros(T, self.n_out) for t in range(T): x = input_spikes[t] # (n_in,) I = self.weight @ x # (n_out,),电流注入 self.V = self.decay * self.V + I # 膜电位更新 spk = (self.V >= self.V_th).float() # 发放判断 self.V = self.V * (1 - spk) # 发放后重置为0 out_spikes[t] = spk return out_spikes

这里有两个坑,我在第一次实现时都踩过。

第一个坑是self.V = self.V * (1 - spk)。很多人习惯写成self.V = self.V - spk * self.V_th,这两种写法在spk=1时效果一样,都是把V降为0;但V_reset不是0的时候语义就不同了。spk=1时V = V * (1-spk)一定是0,而V - spk * V_th是V - V_th,这可能还有残余正电压,会让神经元更容易再次发放。所以我建议根据你的重置机制选择,如果设置重置电位为0,用前者最稳。

第二个坑是权重不要注册成nn.Parameter。因为STDP更新是完全手工的,不经过PyTorch的自动求导,注册成Parameter反而可能在后面调用.backward()时引发意想不到的报错。这里用普通torch.Tensor就够了,更新的时候通过.data赋值并做clamp。

4. 实现STDP学习规则:最核心的几十行代码

4.1 STDP公式拆解:因果增强、反因果抑制

STDP更新量是脉冲时间差Δt的函数,把这个函数画出来是典型的左右不对称形态。当突触前脉冲先于突触后脉冲出现时,Δt = t_post - t_pre > 0,权重增加;当突触后脉冲先于突触前脉冲出现时,Δt < 0,权重减少:

ΔW = A+ * exp(-Δt / τ+) 当 Δt > 0 ΔW = -A- * exp(Δt / τ-) 当 Δt < 0

直觉上这是合理的:如果输入A的脉冲总是出现在输出脉冲之前的20毫秒左右,那A很大概率是引发输出脉冲的原因,连接应该加强;反之,A在输出之后才来,说明它没起作用,连接应该削弱。

但在软件里,我们不能“回溯”去看每个历史脉冲的精确时间,更常见的做法是用一个随时间指数衰减的“Trace”记录神经元最近的放电历史。一个神经元发放的瞬间,它的Trace会跳到1,然后每个时间步按比例衰减:

trace = trace * decay + spike

这样当某个事件发生时,只需要读取当前Trace,就能知道目标事件在多久之前发生过。Trace值越大,说明历史脉冲越“新鲜”。

4.2 用Trace机制让STDP可在线计算

我会在forward循环里维护两个Trace变量:

  • pre_trace,形状(n_out, n_in),记录每个突触前神经元的放电历史。为什么是二维而不是一维?因为这个Trace需要被广播到每个突触后神经元,后面更新权重时它要和输出神经元的脉冲做外积,提前展开成二维方便直接相乘。
  • post_trace,形状(n_out,),记录每个突触后神经元的放电历史。

权重更新的逻辑分两个方向:

  1. 当突触后神经元i在当前时间步发放脉冲时,读取pre_trace[i, :],表示“过去一段时间,哪些突触前神经元频繁放电”。如果pre_trace很大,则它们到神经元i的连接增强。这对应因果增强。
  2. 当突触前神经元j在当前时间步发放脉冲时,读取post_trace[:],表示“过去一段时间,哪些突触后神经元刚发放过”。如果post_trace[i]很大,那么来自神经元j到i的连接削弱。这对应反因果抑制。

用矩阵运算一次性写完:

delta_W = torch.zeros_like(self.weight) if spk.sum() > 0: delta_W += self.A_plus * spk.unsqueeze(1) * self.pre_trace if x.sum() > 0: delta_W -= self.A_minus * x.unsqueeze(0) * self.post_trace.unsqueeze(1) self.weight.data += self.lr * delta_W self.weight.data.clamp_(0.0, 1.0)

注意一个细节:我是在当前时间步脉冲判断完之后、Trace更新之前执行权重更新的。这意味着同一时间步的pre spike和post spike不会互相“看见”,避免同时到达的脉冲对权重产生无意义的更新。这在生物上也有对应解释——完全同时的脉冲通常不产生可塑性变化。

4.3 完整STDP代码与权重约束

前面SNNLayer的代码里其实已经包含了STDP的核心,这里把前向函数完整展开写一遍,方便你直接对比查看:

def forward(self, input_spikes): T = input_spikes.size(0) out_spikes = torch.zeros(T, self.n_out) for t in range(T): x = input_spikes[t] I = self.weight @ x self.V = self.decay * self.V + I spk = (self.V >= self.V_th).float() self.V = self.V * (1 - spk) out_spikes[t] = spk # ---------- STDP 更新 ---------- delta_W = torch.zeros_like(self.weight) if spk.sum() > 0: # 突触后发放:增强那些近期活跃过的突触前连接 delta_W += self.A_plus * spk.unsqueeze(1) * self.pre_trace if x.sum() > 0: # 突触前发放:削弱那些近期活跃过的突触后连接 delta_W -= self.A_minus * x.unsqueeze(0) * self.post_trace.unsqueeze(1) self.weight.data += self.lr * delta_W self.weight.data.clamp_(0.0, 1.0) # ------------------------------- # ---------- Trace 更新 ---------- self.pre_trace = self.decay_pre * self.pre_trace self.pre_trace += x.unsqueeze(0) # 广播到 (n_out, n_in) self.post_trace = self.decay_post * self.post_trace self.post_trace += spk # ------------------------------- return out_spikes

权重约束这里我用的是clamp_(0.0, 1.0)。STDP没有天然的权重范围,如果不加约束,最容易出现的情况就是增强的权重一直涨,反复增强的路径最终全部变成1,失去区分度。把权重限制在0到1之间,从神经科学角度也说得通,突触连接强度本来就有物理上限。

还有一个细节:A_minus我设置得比A_plus大(0.12 vs 0.1)。这不是随手写的。由于在输入脉冲模式下,每轮实验中突触前脉冲的数量往往多于突触后脉冲,削弱的机会比增强多,如果两者一样大,权重会整体往下掉。让A_minus略微大一点其实是保证削弱不要过猛,配合clamp的下界0,这个比例用下来是稳的。

5. 完整实验:让SNN学会区分两类脉冲模式

5.1 构造实验数据:生成两组不同的脉冲序列

这个项目不打算用MNIST这种大路货,而是构造一个能体现STDP“时序依赖”特性的小型实验。我先让40个输入神经元分成两组:

  • 模式A:输入神经元0到19在t=15到t=20之间分批发放脉冲;
  • 模式B:输入神经元20到39在t=30到t=35之间分批发放脉冲。

两组模式在空间和时间上都被分开了,方便我们观察STDP是否真的能让两个输出神经元分别“锁定”一种模式。

代码是这样写的:

def generate_pattern(pattern_id, n_in=40, T=50): spikes = torch.zeros(T, n_in) if pattern_id == 0: idx = list(range(0, 20)) for j, nid in enumerate(idx): spikes[15 + j // 5, nid] = 1.0 else: idx = list(range(20, 40)) for j, nid in enumerate(idx): spikes[30 + j // 7, nid] = 1.0 return spikes

这段代码的效果是,模式A中每个时间步活跃的输入数量大约5个,持续5个时间步;模式B则持续3个时间步,每个时间步活跃7个左右。强度差距不大,避免某一类模式天然更容易引发输出发放。

为了让STDP的分化更快更稳定,我会在初始化时给两个输出神经元一个很小的偏置:神经元1对前20个输入权重略高,神经元2对后20个输入权重略高。这个偏置不是标签,只是模拟了两个神经元初始响应偏好不同的情况,真正的权重演化还是靠STDP完成。

layer = SNNLayer(n_in=40, n_out=2) with torch.no_grad(): layer.weight[0, :20] += 0.2 layer.weight[1, 20:] += 0.2

5.2 训练循环:重置、前向、输出脉冲记录

训练循环比DNN简单许多,因为没有损失函数、没有反向传播,只有“喂数据、跑前向、STDP自动更新”三步。但有一个容易被忽略的关键操作:每个样本开始前必须调用reset(),把膜电位和Trace全部清零。如果忘了,上一轮的放电状态会串到下一轮,导致结果完全不可复现。

epochs = 200 for epoch in range(epochs): for pattern_id in [0, 1]: x = generate_pattern(pattern_id) layer.reset() out_spikes = layer(x) total_spikes = out_spikes.sum(dim=0).tolist() if epoch % 20 == 0: print(f"epoch {epoch:3d} | pattern {pattern_id} | output spikes: {total_spikes}")

训练过程中你会看到比较有趣的现象:一开始两个输出神经元对两类模式的响应差别不大,甚至可能出现混乱。几十个epoch之后,模式A会让神经元0发放更多脉冲,模式B会让神经元1发放更多脉冲,输出脉冲数会呈现出可复现的差异。

5.3 验证结果:神经元响应分化与权重可视化

训练结束后,分别对两种模式做一次测试,并统计每个输出神经元的脉冲总数:

def test_pattern(layer, pattern_id): x = generate_pattern(pattern_id) layer.reset() spikes = layer(x) return spikes.sum(dim=0).tolist() with torch.no_grad(): a_result = test_pattern(layer, 0) b_result = test_pattern(layer, 1) print("Pattern A ->", a_result) print("Pattern B ->", b_result)

如果训练成功,你会看到类似Pattern A -> [5, 0]、Pattern B -> [0, 4]的结果。这代表两个输出神经元已经各自对一种脉冲模式产生选择性响应。

除了输出脉冲数,权重矩阵本身也很有看头。把layer.weight打印出来,你应该能看到一种明显的“对角分布”:神经元0那行权重在前20个位置明显偏大,神经元1那行权重在后20个位置明显偏大。

如果再细致一点,你还能观察到STDP带来的“时序选择”效应:如果某个输入神经元总是在输出神经元发放之前15个时间步左右出现,它的连接就会被增强;如果它的脉冲出现在输出之后,它的连接就会被削弱。这正是STDP和普通Hebbian规则的根本差异——时间顺序,而不是简单共现。

6. 常见问题与排查技巧:我踩过的坑你直接避开

6.1 神经元完全不发放,或脉冲刷屏

最常遇到的情况是训练一跑,输出脉冲全是0。排查顺序一般是:

  • 输入太弱:检查weight @ x求和后膜电位是不是远小于阈值。我初始权重设为0.05到0.2之间,每个时间步注入5到7个输入,电流在0.25到1.4之间,阈值1.0下刚好处于“偶尔发放、偶尔不发放”的临界状态。你的场景如果神经元数更少,可以把初始权重整体抬高,或者调低阈值。
  • 时间常数太小:tau如果设成2或3,一次注入的电流会在一两个时间步内迅速漏光,根本积累不到阈值。tau取10到30之间相对合理。
  • 和“完全不发放”相反的问题是“刷屏”,即每个时间步都发放。主要原因是权重经STDP增强后涨得太快、正反馈循环。最简单的对策是给神经元加一个不应期(refractory period),在代码里维护一个refractory_cnt计数,发放后若干个时间步内脉冲被强制置0:
self.refractory_cnt = torch.zeros(self.n_out) # 在发放判断前: unrefractory = (self.refractory_cnt == 0).float() spk = ((self.V >= self.V_th).float()) * unrefractory # 更新不应期: self.refractory_cnt = torch.clamp(self.refractory_cnt - 1, min=0) self.refractory_cnt += spk * refractory_steps

6.2 权重连锁增强或全部归零

权重一直涨到1,原因是增强条件太容易满足、削弱条件太难触发。如果你发现输出神经元发放频率很高,那pre_trace和post_trace的“错配”机会就会少,STDP变成纯Hebbian增强,权重自然顶到上限。解决办法是适当调高V_th,降低输出发放频率;或者增大A_minus,让削弱机制更敏感。

权重全部归零则是另一个极端。常见原因是A_minus过大,或者输出神经元发放太少,导致“增强”事件很少发生,而每次pre脉冲都在无情地削弱。我实测下来,A_plus和A_minus的比例在1:1到1:2之间比较安全,配合clamp下界0.0,网络不至于彻底失去学习能力。

6.3 STDP训练不收敛怎么办

如果跑了200轮还是没分化,先别急着调A+、A-,建议按顺序检查三件事:

  • 固定随机种子。SNN的随机性比DNN大,不固定种子你根本分不清是算法问题还是运气问题。在代码开头加torch.manual_seed(0)。
  • 检查输出神经元是否每个样本都发放过脉冲。如果某类模式从不引发任何输出脉冲,那它与该模式相关的连接永远不会被增强,自然学不到。这种情况需要降低阈值,或者给初始权重加偏置。
  • 如果两个输出神经元总是被同一个模式触发,说明两者竞争不够激烈。可以在每个时间步做一次简单的赢者全取(Winner-Take-All),同一时间步只允许膜电位最高的那个神经元发放:
if spk.sum() > 1: winner = torch.argmax(self.V) spk = torch.zeros_like(spk) spk[winner] = 1.0

加了WTA之后,两个输出神经元会被迫错开各自的专属模式,分化速度会明显加快。

6.4 性能优化与后续扩展想法

用纯Python循环模拟50个时间步、40个输入神经元,规模很小,跑起来毫无压力。但如果要往大扩展,这个写法会有性能瓶颈,主要瓶颈是时间步的串行依赖,GPU的并行优势发挥不出来。大规模场景我建议优先考虑事件驱动方式:只在有脉冲发生时更新权重和膜电位,而不是每个时间步都全量计算。另外在软件里,STDP本身并不是效率最高的学习方法,它的真正价值在于硬件友好性——在神经形态芯片上,每个脉冲触发本地更新,几乎没有全局通信开销。

如果你对后续扩展有兴趣,我建议按这个顺序往上加东西:第一个扩展是侧抑制(lateral inhibition),能让多个输出神经元各自学到不同特征,不做WTA也能达到类似效果;第二个扩展是卷积SNN加STDP,在MNIST上做无监督初级特征提取;第三个扩展是把奖励信号引入STDP,变成R-STDP,用来做简单的强化学习任务。

我刚做完前两个扩展的时候,最大的感受是:一开始以为搞懂了STDP,实际上只有手写一遍才发现很多细节,比如Trace的更新时机、权重clamp边界、A+和A-的比例,都会直接影响最终能不能收敛。这也是为什么我强烈建议你亲手跑一遍这份代码,先跑通,再一个个改参数看变化。等你把每个参数都“玩坏”过一次,SNN的基础就算真正打牢了。

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

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

立即咨询