07

实现第一个 SNN

从手写 LIF 到可训练分类器

75 MIN
CHAPTER 7/12

先建立一个能用的直觉

训练循环的骨架与普通 PyTorch 类似:数据、前向、损失、backward、optimizer。差异集中在时间循环、状态初始化、spike 读出和 surrogate backward。

它到底解决了什么麻烦

只有亲手管理状态与时间维,才能辨认框架 API 隐藏了哪些假设。

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

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

公式 1 / 2分类读出
y^=argmaxct=1TSc[t]\hat{y}=\arg\max_c\sum_{t=1}^{T}S_c[t]
先用一句话读懂

每个类别对应一个输出神经元;把整个时间窗的 spike 数起来,发放最多的类别作为预测。

符号分别是谁

Sc[t]
类别 c 的输出神经元在第 t 步是否发放
Σt Sc[t]
类别 c 的总 spike count
argmaxc
在所有类别中选 count 最大者

按这个顺序计算

  1. 沿时间维求和,得到每个类别的 count。
  2. 保留每个样本的类别向量。
  3. 推理时对类别维做 argmax。

代入一组数

三类 count 为 [2,7,4],预测类别 1。训练时把 [2,7,4] 当 logits 算 loss,不要先 argmax。

落到代码就是

logits = spike_record.sum(dim=0) loss = F.cross_entropy(logits, labels)
公式 2 / 2时间平均损失
L=1TtCE(Uout[t],y)L=\frac{1}{T}\sum_t\mathrm{CE}(U_{out}[t],y)
先用一句话读懂

不等到序列末尾才监督,而是每一步都用输出膜电位算一次交叉熵,再取平均。这样早期时间步也直接收到学习信号。

符号分别是谁

Uout[t]
第 t 步输出层膜电位或连续读出
y
真实类别标签
CE
交叉熵损失
1/T
取平均,避免仅因增加时间步就放大 loss

按这个顺序计算

  1. 每一步保存输出膜电位。
  2. 分别与同一标签计算 CE。
  3. 把 T 个 loss 相加并除以 T。

代入一组数

四步 CE 为 [1.2,0.9,0.7,0.6],则 L=(1.2+0.9+0.7+0.6)/4=0.85。

落到代码就是

loss = torch.stack([F.cross_entropy(u, y) for u in volts]).mean()

把几个容易混的概念拆开

State reset

每个 batch 开始创建零状态,或明确处理跨片段的 stateful 训练。

Readout

可用 spike count、末步膜电位、时间平均膜电位;选择必须与损失对应。

跟着机制一步一步走

01

一个可复现训练实验至少记录什么?

除了 test accuracy,还要保存随机种子、T、β、阈值、reset、编码、输出读出、每层 spike rate、训练耗时和峰值显存。否则很难判断提升来自网络还是更高的仿真成本。

  • 先在很小数据子集上过拟合,验证梯度真的工作。
  • 再跑单 batch trace,检查膜电位与 spike 是否饱和。
  • 最后才进行完整训练与超参数扫描。
02

输出层一定要发 spike 吗?

不一定。可用 spike count,也可使用输出膜电位作为 logits。前者保持脉冲读出,后者常带来更直接的监督;需要在报告中说明。

公式在代码里对应哪一行

for images, labels in train_loader:
    images, labels = images.to(device), labels.to(device)
    mem1 = torch.zeros(images.size(0), 256, device=device)
    mem2 = torch.zeros(images.size(0), 10, device=device)
    spk_out = []
    for _ in range(num_steps):
        cur1 = fc1(images.flatten(1))
        spk1, mem1 = lif1(cur1, mem1)
        spk2, mem2 = lif2(fc2(spk1), mem2)
        spk_out.append(spk2)
    logits = torch.stack(spk_out).sum(0)
    loss = F.cross_entropy(logits, labels)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

先写下预测,再动参数

在 Code Lab 按数据→编码→网络→前向→损失→更新逐段阅读;先把 num_steps 改为 1,再解释行为变化。

打开交互实验室 →

我最常见到的误解

对 spike 做 argmax 后再算 loss。argmax 会截断梯度;应对 spike count 或膜电位 logits 计算可导损失。

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

snnTorch 与 SpikingJelly 提供神经元、surrogate、multi-step 模式和数据集工具;先掌握手写版本再使用。

进入论文库 →

别急着翻页,先验算一下

本章检测0/3 已作答

01哪个操作应该在 loss.backward() 之前?

02训练前最有价值的快速正确性检查是?

03为什么只报告 accuracy 不足以比较 SNN?