05

Surrogate Gradient

前向是真 spike,反向借一座可导的桥

55 MIN
CHAPTER 5/12

先建立一个能用的直觉

Heaviside 在阈值外导数为 0,在阈值处不可导。直接交给梯度下降,权重几乎收不到有用信号。

Surrogate gradient 不改变前向的 0/1 spike,只在 backward 把不存在或无用的导数替换为阈值附近的平滑斜率。

它到底解决了什么麻烦

它让误差信号能够从 spike 穿回膜电位,再穿回突触权重,同时保留真实的离散前向行为。

Math · 不背公式,顺着计算读

每张卡先讲人话,再拆符号、走计算顺序,最后代入一组数。读完后请合上说明,自己复算一次。

公式 1 / 3真实前向
S=H(VVth)S=H(V-V_{th})
先用一句话读懂

前向传播只问一个硬问题:膜电位差 u=V−Vth 是否非负。答案仍然是离散的 0/1。

符号分别是谁

u=V−Vth
离阈值还有多远;u=0 正好在阈值上
H(u)
硬阶跃函数
S
真实参与下一层计算的 spike

按这个顺序计算

  1. 计算膜电位与阈值的差 u。
  2. u<0 输出 0;u≥0 输出 1。
  3. 保存 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)
公式 2 / 3替代反向
SVσ~(VVth)\frac{\partial S}{\partial V}\approx \tilde{\sigma}'(V-V_{th})
先用一句话读懂

反向时不再使用 H 的真实导数,而是在阈值附近借一个平滑函数的斜率,让损失能告诉 V 应该往哪边移动。

符号分别是谁

∂S/∂V
spike 对膜电位的局部梯度;真实值几乎处处为 0
σ̃′
人为选择的替代导数
V−Vth
梯度窗口以阈值为中心

按这个顺序计算

  1. 上游先给出 ∂L/∂S。
  2. 用 σ̃′(V−Vth) 代替不可用的 ∂S/∂V。
  3. 链式法则相乘,得到 ∂L/∂V,再继续传给权重和旧状态。

代入一组数

上游梯度为 0.4、替代导数在当前 u 处为 0.25,则传给膜电位的梯度是 0.4×0.25=0.1。

落到代码就是

grad_mem = grad_spike * surrogate_derivative(u)
公式 3 / 3Fast sigmoid
σ~(u)=1(1+ku)2\tilde{\sigma}'(u)=\frac{1}{(1+k|u|)^2}
先用一句话读懂

离阈值越远,替代梯度越小;k 决定衰减有多快。它的目的不是拟合前向输出,而是规定 backward 的有效窗口。

符号分别是谁

u
膜电位与阈值的差
k
斜率/窗口参数;越大窗口越窄
σ̃′(u)
传给膜电位的局部梯度倍率

按这个顺序计算

  1. 取距离 |u|。
  2. 计算 1+k|u|。
  3. 平方后取倒数;在 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 路径的梯度,避免额外递归项;这是一项建模选择,不是固定真理。

跟着机制一步一步走

01

替代导数的宽度与尺度

斜率太大时,只有极少膜电位落在有效梯度区;斜率太小时,梯度分布很宽,但与硬阈值局部行为偏差更大。最佳选择常与学习率、阈值分布和网络深度耦合。

  • Sigmoid:平滑、两侧指数衰减。
  • Fast sigmoid:重尾更明显,计算简单。
  • Arctan:平滑且尾部衰减较慢。
  • Triangular:有限支持集,窗口外梯度严格为零。
02

Autograd Function 在做什么?

forward 保存膜电位差 u=V−Vth 并返回硬阈值;backward 接收上游梯度,再乘 surrogate derivative。权重梯度仍由 PyTorch 链式法则自动累积。

带走这一点Surrogate gradient 是 biased gradient estimator,不是 spike 的真实经典导数。

公式在代码里对应哪一行

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

先写下预测,再动参数

切换三种 surrogate,改变 slope,观察梯度有效区如何围绕阈值收缩或展开。

打开交互实验室 →

我最常见到的误解

说“训练时前向用 sigmoid”。标准 surrogate-gradient 训练的关键正相反:前向仍用硬 spike。

到了论文里,它会怎么出现

近年的工作研究 surrogate 形状、梯度消失、在线训练和不保存完整时间图的近似反向。

进入论文库 →

别急着翻页,先验算一下

本章检测0/3 已作答

01Surrogate gradient 主要替换哪一部分?

02Surrogate slope 极大可能导致什么?

03自定义 autograd backward 返回的是什么?