Skip to content
GitHub

NIR export

Learn

NIR export

A trained spiking network is only useful where it can run. NIR describes one in continuous time, so another simulator, or a neuromorphic chip's toolchain, can read it.

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.

  1. Inputinput
  2. Affinenode 0weight [64 × 7] -0.965 to 1.15bias [64] -0.233 to 0.283
  3. LIFnode 1tau [64] 0.0300r [64] 3.53v_leak [64] 0.00v_threshold [64] 1.00v_reset [64] 0.00
  4. Affinenode 2weight [64 × 64] -0.952 to 0.630bias [64] -0.0993 to 0.342
  5. LIFnode 3tau [64] 0.0300r [64] 3.53v_leak [64] 0.00v_threshold [64] 1.00v_reset [64] 0.00
  6. Affinenode 4weight [2 × 64] -1.14 to 0.684bias [2] -0.0308 to 7.58e-3
  7. LInode 5tau [2] 0.0500r [2] 5.52v_leak [2] 0.00
  8. 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 nn
import jax
import jax.numpy as jnp
import nir
from sparx.nir import from_nir, to_nir
from 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 seconds
nir.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.0

What 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.

All of Learn