from torch import nn import snntorch as snn from snntorch import surrogate beta = 0.9 spike_grad = surrogate.fast_sigmoid() net = nn.Sequential( nn.Flatten(), nn.Linear(20, 256), snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True), nn.Linear(256, 256), snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True), nn.Linear(256, 256), snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True), nn.Linear(256, 35), snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True, output=True), )