sparx.nir
Exchanging networks through NIR, the Neuromorphic Intermediate Representation (Pedersen et al. 2024).
graph = to_nir(model, variables, dt=1e-3) # a flax nn.Sequential, see belowmodel, variables = from_nir(graph, dt=1e-4, discretization="euler")NIR describes neurons in continuous time, tau dv/dt = (v_leak - v) + r I
with a hard reset to v_reset, and leaves the step to each library. A
sparx LIF is the discrete v <- decay v + x, so a conversion names how
time was discretized:
"exact"(sparx’s own): the leak’s exact solution over a step, the input held over it, sotau = -dt / ln(decay)andr = 1 / (1 - decay)."euler": forward Euler,decay = 1 - dt / tauandr = tau / dt, which snnTorch’s NIR import and export assume atdt = 1e-4s.
r is folded into the preceding weight layer on import. NIR times are in
seconds. Supported: nn.Sequential stacks of these layers, with anything
else refused:
flax.linen.Dense, as NIRAffine(orLinearon import);flax.linen.Convwith a 2-d kernel, asConv2d;sparx.nn.Flattenover each example’s whole shape, asFlatten;sparx.nn.LIFwithreset="zero"(NIR’s reset) and one threshold, asLIF;sparx.nn.IFwithreset="zero", asIF, whoser = 1 / dtmakes a step add its input to the membrane, under either discretization;sparx.nn.LI, a classifier’s readout, asLI, the same leak without a threshold;sparx.nn.Recurrent(LIF(...))on a flat input, as aLIFnode with aLinearedge from its output back to its input.
Images. NIR follows PyTorch: one example of a 2-d convolution’s input is
[C, H, W], and its weight is [C_out, C_in / groups, kH, kW]. Flax is
channels last: the input is [..., H, W, C] and the kernel
[kH, kW, C_in / groups, C_out]. The weights are transposed both ways, and
a sparx tensor is a NIR tensor with its channel axis moved last; per-neuron
parameters of a LIF node after a convolution (tau, r, v_threshold)
are [C, H, W] in NIR. A NIR Flatten orders the features C, H, W, as
PyTorch does, and so does sparx.nn.Flatten, which moves the channel axis
first before it reshapes; a plain flax reshape would order them H, W, C
and scramble the dense layer after it. NIR’s convolution pads both sides
equally, so a flax padding that pads one side more ("SAME" with an even
kernel or a stride) is refused. NIR 1.0.8 types a Conv2d from its
kernel’s height alone and as if it were ungrouped, so its graph check
rejects a non-square or grouped kernel. from_nir reads both from a
graph built without the check, and to_nir raises NIR’s type error.
to_nir needs input_shape, one example’s shape in sparx’s layout
((H, W, C)), when the stack does not start with a Dense layer.
Recurrence. sparx’s Recurrent adds s[t-1] @ W to the input of step t.
The feedback is the previous step’s spikes, delayed by one step, since a
spike cannot reach its own neuron within the step that fired it. snnTorch’s
RLeaky does the same, and exports as a LIF node "k.lif" and an
Affine node "k.w_rec" with edges both ways between them, inside the
chain; this module reads that cycle as one Recurrent layer and writes it
the same way, with node names to match. The NIR weight is [out, in], so
W is its transpose. sparx’s recurrent matrix has no bias: an Affine
feedback’s bias is a constant input every step, so it is added to the bias
of the layer before the LIF node on import, and the feedback exports as a
Linear node. A recurrent subgraph nested as a NIRGraph node, with its
own Input and Output, imports too.
snnTorch 1.0’s import_nir cannot read a recurrent graph (its subgraph
step fails on the graph it writes), so to_nir’s recurrent graphs are
checked against sparx’s own import and snnTorch’s export only.
Contents
Section titled “Contents”| Name | |
|---|---|
from_nir | A sequential stack and its variables from a NIR graph of a chain of the nodes above, read with discretization at dt seconds. Image layers take inputs [T, B, H, W, C]. |
to_nir | The NIR graph of a sequential stack of the layers above, at step dt seconds. |
from_nir
Section titled “from_nir”def from_nir(graph: nir.NIRGraph, dt: float, discretization: Discretization = 'exact') -> tuple[nn.Sequential, dict[str, dict[str, dict[str, np.ndarray]]]]A sequential stack and its variables from a NIR graph of a chain of the nodes above, read with
discretization at dt seconds. Image layers take inputs [T, B, H, W, C].
to_nir
Section titled “to_nir”def to_nir(model: nn.Sequential, variables: StackVariables, dt: float, discretization: Discretization = 'exact', input_shape: Sequence[int] | None = None) -> nir.NIRGraphThe NIR graph of a sequential stack of the layers above, at step dt seconds.
input_shape is one example’s input shape in sparx’s layout, (H, W, C)
for images; a stack that starts with a Dense layer has it from the
kernel.