A network as a graph of equations
The Neuromorphic Intermediate Representation (Pedersen et al., Nature Communications 2024) stores a spiking network as a graph of nodes, each a primitive with its parameters: affine maps, convolutions, and neurons such as LIF written as differential equations in seconds. It leaves the step to the reader. A sparx LIF is the discrete v = decay * v + x, so exporting one names how time was discretized: by default the leak's exact solution over a step, which gives tau = -dt / ln(decay) and r = 1 / (1 - decay).
The pilot, exported
This is the network flying the drone on the front page, as sparx.nir.to_nir writes it with a step of 10 ms. Its LIF neurons' time constant of 3 steps becomes 30 ms.
- Inputinput
- Affinenode 0
weight[64 × 7] -0.965 to 1.15bias[64] -0.233 to 0.283 - LIFnode 1
tau[64] 0.0300r[64] 3.53v_leak[64] 0.00v_threshold[64] 1.00v_reset[64] 0.00 - Affinenode 2
weight[64 × 64] -0.952 to 0.630bias[64] -0.0993 to 0.342 - LIFnode 3
tau[64] 0.0300r[64] 3.53v_leak[64] 0.00v_threshold[64] 1.00v_reset[64] 0.00 - Affinenode 4
weight[2 × 64] -1.14 to 0.684bias[2] -0.0308 to 7.58e-3 - LInode 5
tau[2] 0.0500r[2] 5.52v_leak[2] 0.00 - Outputoutput
Download pilot.nir (89 KB, HDF5). sparx.nir.from_nir reads it back into an nn.Sequential of the same layers: run on 200 steps of random input, the network read back gives exactly the original's output, a largest difference of 0.
The code
import flax.linen as nnimport jaximport jax.numpy as jnpimport nir
from sparx.nir import from_nir, to_nirfrom sparx.nn import LI, LIF
pilot = nn.Sequential([nn.Dense(64), LIF(tau=3.0, reset="zero"), nn.Dense(64), LIF(tau=3.0, reset="zero"), nn.Dense(2), LI(tau=5.0)])variables = pilot.init(jax.random.key(0), jnp.zeros((1, 1, 7)))graph = to_nir(pilot, variables, dt=0.01) # a step is 10 ms; NIR counts secondsnir.write("pilot.nir", graph)
model, read = from_nir(nir.read("pilot.nir"), dt=0.01)x = jax.random.normal(jax.random.key(1), (200, 8, 7))print(jnp.abs(pilot.apply(variables, x) - model.apply(read, x)).max()) # 0.0What exports
nn.Sequential stacks of dense, convolution and flatten layers, hard-reset LIF and IF with one threshold and time constant, LI, and dense recurrent LIF. Soft reset, adaptive and physical neurons, sparse or plastic recurrence and learned delays have no NIR form here, and SpikingMLP and SEWResNet do not export as models (status). The pilot was built from exportable layers on purpose: its LIF neurons reset to zero.
Checked
sparx's export and import are checked against snnTorch (tests/test_nir.py): the node types sparx writes run spike for spike in snnTorch, and stacks round-trip bit for bit through both discretizations. The pilot's own round trip above ran in site/lab/pilot.py nir, with NIR 1.0.8.