0. 环境与最小正确性检查
先只安装 PyTorch 与 torchvision。正式训练前,用 32~128 个样本验证模型能否过拟合,再打印一个样本的膜电位和 spike trace。这个步骤比直接跑完整数据集更容易发现状态未清零、梯度断裂或全层沉默。
python -m venv .venv
# Windows: .venv\Scripts\activate
pip install torch torchvision
python train_mnist_snn.py --epochs 5 --steps 12 --beta 0.9下载完整可运行脚本1. 手写 LIF 与真实 surrogate backward
前向用硬阈值;backward 只替换局部导数。reset 路径用 detach,避免把 reset 当作连续可导过程。
import torch
import torch.nn as nn
import torch.nn.functional as F
class SpikeFn(torch.autograd.Function):
@staticmethod
def forward(ctx, u):
ctx.save_for_backward(u)
return (u >= 0).to(u.dtype)
@staticmethod
def backward(ctx, grad):
(u,) = ctx.saved_tensors
return grad / (1 + 10 * u.abs()).pow(2)
spike_fn = SpikeFn.apply
class LIF(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
spk = spike_fn(mem - self.threshold)
mem = mem - spk.detach() * self.threshold
return spk, mem2. 完整的时间展开与训练
示例使用直接编码:静态图像在每一步作为电流输入。输出层 spike count 作为分类 logits。完整下载脚本进一步加入 Fashion-MNIST、Spiking CNN、测试循环、随机种子、spike rate、墙钟时间与 checkpoint。
class SimpleSNN(nn.Module):
def __init__(self, steps=20):
super().__init__()
self.steps = steps
self.fc1 = nn.Linear(28 * 28, 256)
self.fc2 = nn.Linear(256, 10)
self.lif1, self.lif2 = LIF(), LIF()
def forward(self, x):
x = x.flatten(1)
mem1 = x.new_zeros(x.size(0), 256)
mem2 = x.new_zeros(x.size(0), 10)
outputs = []
for _ in range(self.steps):
spk1, mem1 = self.lif1(self.fc1(x), mem1)
spk2, mem2 = self.lif2(self.fc2(spk1), mem2)
outputs.append(spk2)
return torch.stack(outputs).sum(0)
model = SimpleSNN().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
loss = F.cross_entropy(model(images), labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()读代码时逐项确认:状态在哪里创建、每一步读入什么、reset 发生在哪里、loss 对什么量计算、batch 结束后状态是否被清空。
3. 再使用 snnTorch
框架减少样板代码,但不会替你决定状态何时清零、时间维放在哪里、用什么读出和损失。
import snntorch as snn
from snntorch import surrogate
spike_grad = surrogate.fast_sigmoid(slope=25)
lif = snn.Leaky(beta=0.9, spike_grad=spike_grad)
# forward 返回 spike 与 membrane state
spk, mem = lif(current, mem)
# 多步训练时仍要明确初始化、时间循环与读出:
# logits = torch.stack(spk_rec).sum(0)4. SpikingJelly 的 multi-step 模式
SpikingJelly 可以让模块一次接收 [T,B,…]。效率更高之前,先确认 step_mode、detach_reset 与 reset_net 的语义。
from spikingjelly.activation_based import neuron, surrogate, functional
lif = neuron.LIFNode(
tau=2.0,
surrogate_function=surrogate.ATan(),
detach_reset=True,
step_mode="m", # multi-step: [T, B, ...]
)
spike_seq = lif(current_seq)
# 每个独立 batch 结束后显式清空所有有状态模块
functional.reset_net(model)5. 评估协议
Accuracy
固定时间步和读出规则后,在测试集统计 top-1。
Spike rate
各层 spike 数除以神经元数与时间步,分层报告。
Latency
同时报告算法时间步与真实墙钟时间,两者不是一回事。
Compute
区分理论 AC/MAC 估算、GPU 实测与神经形态芯片实测。