先建立一个能用的直觉
Heaviside 在阈值外导数为 0,在阈值处不可导。直接交给梯度下降,权重几乎收不到有用信号。
Surrogate gradient 不改变前向的 0/1 spike,只在 backward 把不存在或无用的导数替换为阈值附近的平滑斜率。
它到底解决了什么麻烦
它让误差信号能够从 spike 穿回膜电位,再穿回突触权重,同时保留真实的离散前向行为。
Math · 不背公式,顺着计算读
每张卡先讲人话,再拆符号、走计算顺序,最后代入一组数。读完后请合上说明,自己复算一次。
前向传播只问一个硬问题:膜电位差 u=V−Vth 是否非负。答案仍然是离散的 0/1。
符号分别是谁
- u=V−Vth
- 离阈值还有多远;u=0 正好在阈值上
- H(u)
- 硬阶跃函数
- S
- 真实参与下一层计算的 spike
按这个顺序计算
- 计算膜电位与阈值的差 u。
- u<0 输出 0;u≥0 输出 1。
- 保存 u,供 backward 估计局部梯度。
代入一组数
Vth=1 时,V=0.97 得到 u=−0.03、S=0;V=1.02 得到 u=0.02、S=1。
落到代码就是
u = mem - threshold
spike = (u >= 0).to(u.dtype)反向时不再使用 H 的真实导数,而是在阈值附近借一个平滑函数的斜率,让损失能告诉 V 应该往哪边移动。
符号分别是谁
- ∂S/∂V
- spike 对膜电位的局部梯度;真实值几乎处处为 0
- σ̃′
- 人为选择的替代导数
- V−Vth
- 梯度窗口以阈值为中心
按这个顺序计算
- 上游先给出 ∂L/∂S。
- 用 σ̃′(V−Vth) 代替不可用的 ∂S/∂V。
- 链式法则相乘,得到 ∂L/∂V,再继续传给权重和旧状态。
代入一组数
上游梯度为 0.4、替代导数在当前 u 处为 0.25,则传给膜电位的梯度是 0.4×0.25=0.1。
落到代码就是
grad_mem = grad_spike * surrogate_derivative(u)离阈值越远,替代梯度越小;k 决定衰减有多快。它的目的不是拟合前向输出,而是规定 backward 的有效窗口。
符号分别是谁
- u
- 膜电位与阈值的差
- k
- 斜率/窗口参数;越大窗口越窄
- σ̃′(u)
- 传给膜电位的局部梯度倍率
按这个顺序计算
- 取距离 |u|。
- 计算 1+k|u|。
- 平方后取倒数;在 u=0 时梯度最大为 1。
代入一组数
k=10 时,u=0 的值为 1;u=0.1 时为 1/(1+1)²=0.25;u=0.5 时只剩 1/36≈0.028。
落到代码就是
surrogate = 1 / (1 + k * u.abs()).pow(2)把几个容易混的概念拆开
Shape choice
Sigmoid、fast sigmoid、arctan、triangular 都可用;尺度与宽度会改变梯度流。
Detach reset
常见实现阻断 reset 路径的梯度,避免额外递归项;这是一项建模选择,不是固定真理。
跟着机制一步一步走
替代导数的宽度与尺度
斜率太大时,只有极少膜电位落在有效梯度区;斜率太小时,梯度分布很宽,但与硬阈值局部行为偏差更大。最佳选择常与学习率、阈值分布和网络深度耦合。
- Sigmoid:平滑、两侧指数衰减。
- Fast sigmoid:重尾更明显,计算简单。
- Arctan:平滑且尾部衰减较慢。
- Triangular:有限支持集,窗口外梯度严格为零。
Autograd Function 在做什么?
forward 保存膜电位差 u=V−Vth 并返回硬阈值;backward 接收上游梯度,再乘 surrogate derivative。权重梯度仍由 PyTorch 链式法则自动累积。
公式在代码里对应哪一行
class Spike(torch.autograd.Function):
@staticmethod
def forward(ctx, u):
ctx.save_for_backward(u)
return (u >= 0).to(u.dtype)
@staticmethod
def backward(ctx, grad_output):
(u,) = ctx.saved_tensors
surrogate = 1 / (1 + 10 * u.abs()).pow(2)
return grad_output * surrogate