Skip to content
GitHub

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 below
model, 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, so tau = -dt / ln(decay) and r = 1 / (1 - decay).
  • "euler": forward Euler, decay = 1 - dt / tau and r = tau / dt, which snnTorch’s NIR import and export assume at dt = 1e-4 s.

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 NIR Affine (or Linear on import);
  • flax.linen.Conv with a 2-d kernel, as Conv2d;
  • sparx.nn.Flatten over each example’s whole shape, as Flatten;
  • sparx.nn.LIF with reset="zero" (NIR’s reset) and one threshold, as LIF;
  • sparx.nn.IF with reset="zero", as IF, whose r = 1 / dt makes a step add its input to the membrane, under either discretization;
  • sparx.nn.LI, a classifier’s readout, as LI, the same leak without a threshold;
  • sparx.nn.Recurrent(LIF(...)) on a flat input, as a LIF node with a Linear edge 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.

Name
from_nirA 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_nirThe NIR graph of a sequential stack of the layers above, at step dt seconds.
def from_nir(graph: nir.NIRGraph, dt: float, discretization: Discretization = 'exact') -> tuple[nn.Sequential, dict[str, dict[str, dict[str, np.ndarray]]]]

sparx.nir on GitHub

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

def to_nir(model: nn.Sequential, variables: StackVariables, dt: float, discretization: Discretization = 'exact', input_shape: Sequence[int] | None = None) -> nir.NIRGraph

sparx.nir on GitHub

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