先建立一个能用的直觉
把膜电位想成一只有漏孔的杯子:输入电流向杯中加水,leak 持续排水;水位越过阈值就产生 spike,随后水位被重置。
IF 没有漏孔,过去输入会永久积累;LIF 用泄漏让较早输入的影响指数衰减,因此具有有限时间尺度。
它到底解决了什么麻烦
LIF 是复杂生物神经元的可计算近似:它保留最关键的积分、记忆与发放机制,同时容易离散化和批量训练。
Math · 不背公式,顺着计算读
每张卡先讲人话,再拆符号、走计算顺序,最后代入一组数。读完后请合上说明,自己复算一次。
膜电位变化由两股力量拉扯:leak 把 V 拉回静息电位,输入电流 RI(t) 把它推高或压低。τm 决定这种变化有多快。
符号分别是谁
- V(t)
- 随连续时间变化的膜电位
- Vrest
- 没有输入时膜电位趋向的静息值
- I(t)
- 输入电流;R 把电流尺度换成电压尺度
- τm
- 膜时间常数;越大表示遗忘越慢
按这个顺序计算
- 若 V 高于 Vrest,−(V−Vrest) 为负,leak 会把 V 往下拉。
- RI(t) 与 leak 相加,决定此刻净变化方向。
- 再除以 τm,得到 dV/dt;τm 越大,同样驱动力造成的变化越慢。
代入一组数
取 Vrest=0、V=0.2、RI=0.6、τm=20 ms,则 dV/dt=(−0.2+0.6)/20=0.02 每毫秒:输入大于泄漏,膜电位正在上升。
把连续变化切成一格一格来算:先保留上一时刻的 β 倍,再加本步输入。β 就是这一步留下多少记忆。
符号分别是谁
- V[t−1]
- 上一步 reset 之后留下的膜电位
- β
- 离散衰减系数,通常在 0 到 1 之间
- I[t]
- 已经按离散时间步缩放后的输入电流
- V[t]
- 本步用于阈值判定的膜电位
按这个顺序计算
- 用 β 乘旧膜电位,得到漏掉一部分后的记忆。
- 加上当前输入 I[t]。
- 拿更新后的 V[t] 与阈值比较;若发放,再执行 reset。
代入一组数
V[t−1]=0.7、β=0.8、I[t]=0.5 时,V[t]=0.8×0.7+0.5=1.06。阈值若为 1,本步会发放。
落到代码就是
mem = beta * mem + current第一段决定发不发,第二段处理发放后的状态。subtract reset 只减去一个阈值,所以超过阈值的余量会留下。
符号分别是谁
- S[t]
- 本步是否发放
- Vth
- 阈值,同时也是 subtract reset 的减量
- V[t]←…
- 箭头表示原地更新发放后的状态,不是新方程的等号
按这个顺序计算
- 先用发放前 V[t] 计算 S[t]。
- 若 S[t]=0,减去 0,膜电位不因 reset 改变。
- 若 S[t]=1,减去 Vth,保留超过阈值的余量。
代入一组数
V=1.06、Vth=1 时先输出 S=1,再得到 reset 后 V=0.06;hard reset 则会直接得到 0。
落到代码就是
spike = (mem >= threshold).float()
mem = mem - spike.detach() * threshold把几个容易混的概念拆开
Membrane potential
神经元的隐藏状态,不是输出概率。它积累输入,也随时间泄漏。
Reset
hard reset 设回固定值;subtract reset 减去阈值。两者在高输入下会产生不同 spike 时序。
跟着机制一步一步走
从连续方程到 β:离散化没有消失
对 τₘ dV/dt = −(V−Vrest)+RI(t) 使用前向 Euler,可得到 V[t+1]≈V[t]+Δt/τₘ·(−(V[t]−Vrest)+RI[t])。整理后,上一时刻状态的系数就是近似的 β。
- Δt 变大而 τₘ 不变:每一步遗忘更强。
- 常见精确衰减写法是 β=exp(−Δt/τₘ)。
- 代码里只写 beta,不代表物理时间尺度可以忽略。
Hard reset 与 subtract reset 会改变什么?
Hard reset 发放后把膜电位设到固定值;subtract reset 只减去阈值,保留超出的余量。持续强输入下,后者更接近积分守恒,也可能产生更密集的 spike。
- 比较 reset 前膜电位,而不是只看 spike count。
- 训练与推理必须使用同一 reset 规则。
- reset 路径是否参与梯度是另一项独立选择。
公式在代码里对应哪一行
class LIF(torch.nn.Module):
def __init__(self, beta=0.9, threshold=1.0):
super().__init__()
self.beta, self.threshold = beta, threshold
def forward(self, current, mem):
mem = self.beta * mem + current
spike = (mem >= self.threshold).to(mem.dtype)
mem = mem - spike.detach() * self.threshold
return spike, mem