FlexSN#
Authors: Yifan Huang (AllenYolk) and Wei Fang
中文版: FlexSN
FlexSN turns a pure-PyTorch single-step neuron function into a stateful
SpikingJelly neuron and can generate a Triton kernel for multi-step CUDA
execution. See Automatic neuron execution for device-based neuron execution.
Use FlexSN for custom multi-step neurons and public neuron/functional interfaces for fixed neurons. The old Auto CUDA translator and neuron code-generation templates have been removed.
备注
The torch.sigmoid examples produce continuous outputs to demonstrate
composition and gradients. Built-in LIF uses a step forward and surrogate
backward, so these outputs are not equivalent. Construction runs the core
on unit tensors. Keep it pure, without captured tensors or modules, and pass
parameters through static_inputs.
Describing neuron dynamics with a function#
Most spiking neurons can be written at one discrete time step as
Here \(X_i\) denotes an input, \(Y_i\) an output, and \(V_i\) a
state carried between time steps. FlexSN represents this equation with
core(*step_inputs, *states, *static_inputs)
-> (*outputs, *updated_states)
The final num_states return values must update the input states in order.
For example, this function describes a soft-reset LIF neuron without input
decay:
import torch
def lif_core(x: torch.Tensor, v: torch.Tensor):
h = 0.5 * v + x
spike = torch.sigmoid(h - 1.0)
v = h - spike
return spike, v
core must be pure: it must not capture a Tensor or nn.Module. Ordinary
numeric hyperparameters may live in a closure. Tensors that should train with
the model or appear in its state_dict belong in static_inputs.
Building a neuron with several states#
Consider a neuron with two inputs, two outputs, and two states. rho adapts
the threshold of the first output, while y blends hard and soft membrane
reset:
import torch
def complicated_lif_core_generator(beta: float, gamma: float):
def complicated_lif_core(
x: torch.Tensor,
y: torch.Tensor,
v: torch.Tensor,
rho: torch.Tensor,
):
h = beta * v + x
s1 = torch.sigmoid(h - (rho + 1.0))
s2 = torch.sigmoid(h - 1.0)
rho = gamma * rho + s1
v_hard = h * (1.0 - s1)
v_soft = h - s2
modulation = torch.sigmoid(y)
v = v_hard * modulation + v_soft * (1.0 - modulation)
return s1, s2, v, rho
return complicated_lif_core
The first two returns are outputs; the last two update v and rho:
Pass the state count to the constructor. FlexSN infers the input and output counts from the signature and one call with unit tensors, without example inputs:
from spikingjelly.activation_based import neuron
f = neuron.FlexSN(
core=complicated_lif_core_generator(beta=0.5, gamma=0.9),
num_states=2,
step_mode="m",
store_state_seqs=True,
).cuda()
x = torch.randn([16, 3, 32, 32], device="cuda")
y = torch.randn([16, 3, 32, 32], device="cuda")
s1, s2 = f(x, y)
v_seq, rho_seq = f.state_seqs
final_v, final_rho = f.states
print(s1.shape, s2.shape)
print(v_seq.shape, rho_seq.shape)
print(final_v.shape, final_rho.shape)
forward returns a Tensor for one output and a tuple for several outputs.
states and state_seqs are always tuples. Call reset() after each
independent sequence to clear managed state.
Managed and functional state#
forward initializes, updates, and stores state automatically.
Use functional_forward when state ownership belongs to the caller. It does
not modify the module's states:
f_torch = neuron.FlexSN(
core=complicated_lif_core_generator(beta=0.5, gamma=0.9),
num_states=2,
)
initial_states = (
torch.zeros_like(x[0]),
torch.zeros_like(x[0]),
)
(s1, s2), (final_v, final_rho) = f_torch.functional_forward(
(x, y), initial_states, static_inputs=()
)
assert f_torch.states == (None, None)
States default to zero tensors shaped like one step of the first input. Override
init_states when a model needs a different rule:
class NonzeroFlexSN(neuron.FlexSN):
@staticmethod
def init_states(num_states, step_mode, *inputs):
reference = inputs[0] if step_mode == "s" else inputs[0][0]
return tuple(torch.ones_like(reference) for _ in range(num_states))
Static inputs#
Tensors reused at every time step are passed through static_inputs.
Parameters are registered as parameters and other tensors as buffers; both are
included in state_dict. The PLIF dynamics below use a trainable
membrane-decay parameter:
def plif_core(x, v, w):
reciprocal_tau = w.sigmoid()
h = v + reciprocal_tau * (x - v)
spike = torch.sigmoid(h - 1.0)
return spike, h * (1.0 - spike)
w = torch.nn.Parameter(torch.tensor(0.0))
plif = neuron.FlexSN(
plif_core,
num_states=1,
static_inputs=(w,),
)
A functional call supplies static values explicitly, so it can use another value without replacing the module parameter:
x_seq = torch.randn(8, 4)
v0 = (torch.zeros_like(x_seq[0]),)
outputs, states = plif.functional_forward(
(x_seq,), v0, static_inputs=(torch.tensor(1.0),)
)
A static tensor must be a scalar or have the same number of elements as one input step. Arbitrary broadcasting is not supported.
Automatic execution and torch.compile#
FlexSN selects its execution automatically. CPU uses the Torch implementation. On CUDA, supported cores use the generated Triton kernels; known unsupported compositions use the Torch/HOP path. The constructor has no backend argument.
import torch.nn as nn
from spikingjelly.activation_based import neuron
flex = neuron.FlexSN(lif_core, 1).cuda()
model = nn.Sequential(nn.Linear(512, 512), flex, nn.Linear(512, 512)).cuda()
compiled = torch.compile(model, fullgraph=True)
output = compiled(torch.randn(8, 64, 512, device="cuda"))
A supported CUDA core compiles when it first receives a real CUDA input. An unsupported operation or a kernel error after selection is reported directly.
Limits and migration#
The leading dimension of a multi-step input is time
T;T == 0is rejected.Single-step mode executes the core directly; automatic CUDA fusion targets multi-step mode.
Changing step mode preserves final states and clears derived
state_seqs.The old
num_inputs,num_outputs,example_inputs,example_outputs, andrequires_gradconstructor arguments are removed.FlexSNKernelandFlexSN.kernelare removed. Usefunctional_forwardfor explicit-state execution.
Training and state management#
This training example uses trainable static inputs and stores full state traces. Reset independent batches after backward and parameter updates; retain state for continuous sequences. The final explicit-state call leaves module memory unchanged:
import torch
from spikingjelly.activation_based import neuron
def smooth_core(x, v, w):
h = v + w.sigmoid() * (x - v)
output = torch.sigmoid(h - 1.0)
return output, h * (1.0 - output)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
node = neuron.FlexSN(
smooth_core,
1,
static_inputs=(torch.nn.Parameter(torch.tensor(0.0)),),
store_state_seqs=True,
).to(device)
optimizer = torch.optim.SGD(node.parameters(), lr=0.01)
x = torch.rand(4, 2, 8, device=device, requires_grad=True)
output = node(x)
assert node.state_seqs[0].shape == x.shape
(output.sum() + node.states[0].sum()).backward()
assert node.static_inputs[0].grad is not None
optimizer.step()
node.reset()
assert node.states == (None,)
# Explicit-state calls return output/state tuples without changing managed state.
outputs, states = node.functional_forward(
(x.detach(),), (torch.zeros_like(x[0]),), static_inputs=node.static_inputs
)
assert outputs[0].shape == x.shape and states[0].shape == x.shape[1:]
assert node.states == (None,)