sparxml 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (89) hide show
  1. sparxml-0.1.0/LICENSE +21 -0
  2. sparxml-0.1.0/PKG-INFO +230 -0
  3. sparxml-0.1.0/README.md +180 -0
  4. sparxml-0.1.0/pyproject.toml +116 -0
  5. sparxml-0.1.0/setup.cfg +4 -0
  6. sparxml-0.1.0/src/sparx/__init__.py +69 -0
  7. sparxml-0.1.0/src/sparx/config.py +102 -0
  8. sparxml-0.1.0/src/sparx/datasets.py +246 -0
  9. sparxml-0.1.0/src/sparx/dynamics/__init__.py +178 -0
  10. sparxml-0.1.0/src/sparx/dynamics/core.py +400 -0
  11. sparxml-0.1.0/src/sparx/dynamics/homeostasis.py +101 -0
  12. sparxml-0.1.0/src/sparx/dynamics/ml.py +925 -0
  13. sparxml-0.1.0/src/sparx/dynamics/neurons.py +648 -0
  14. sparxml-0.1.0/src/sparx/dynamics/plasticity.py +395 -0
  15. sparxml-0.1.0/src/sparx/dynamics/synapses.py +424 -0
  16. sparxml-0.1.0/src/sparx/encode.py +167 -0
  17. sparxml-0.1.0/src/sparx/graph/__init__.py +79 -0
  18. sparxml-0.1.0/src/sparx/graph/connectivity.py +189 -0
  19. sparxml-0.1.0/src/sparx/graph/connectome.py +324 -0
  20. sparxml-0.1.0/src/sparx/graph/delivery.py +214 -0
  21. sparxml-0.1.0/src/sparx/graph/models.py +253 -0
  22. sparxml-0.1.0/src/sparx/graph/network.py +1284 -0
  23. sparxml-0.1.0/src/sparx/graph/simulate.py +269 -0
  24. sparxml-0.1.0/src/sparx/learn/__init__.py +58 -0
  25. sparxml-0.1.0/src/sparx/learn/convert.py +282 -0
  26. sparxml-0.1.0/src/sparx/learn/diffusion.py +338 -0
  27. sparxml-0.1.0/src/sparx/learn/events.py +174 -0
  28. sparxml-0.1.0/src/sparx/learn/online.py +439 -0
  29. sparxml-0.1.0/src/sparx/learn/predictive.py +240 -0
  30. sparxml-0.1.0/src/sparx/learn/reinforce.py +106 -0
  31. sparxml-0.1.0/src/sparx/losses.py +146 -0
  32. sparxml-0.1.0/src/sparx/metrics.py +40 -0
  33. sparxml-0.1.0/src/sparx/models.py +339 -0
  34. sparxml-0.1.0/src/sparx/nir.py +497 -0
  35. sparxml-0.1.0/src/sparx/nn/__init__.py +57 -0
  36. sparxml-0.1.0/src/sparx/nn/delays.py +112 -0
  37. sparxml-0.1.0/src/sparx/nn/hebbian.py +114 -0
  38. sparxml-0.1.0/src/sparx/nn/neurons.py +433 -0
  39. sparxml-0.1.0/src/sparx/nn/parallel.py +156 -0
  40. sparxml-0.1.0/src/sparx/nn/reshape.py +67 -0
  41. sparxml-0.1.0/src/sparx/objectives.py +729 -0
  42. sparxml-0.1.0/src/sparx/py.typed +0 -0
  43. sparxml-0.1.0/src/sparx/rates.py +72 -0
  44. sparxml-0.1.0/src/sparx/registry.py +52 -0
  45. sparxml-0.1.0/src/sparx/serve.py +226 -0
  46. sparxml-0.1.0/src/sparx/spiketrains.py +122 -0
  47. sparxml-0.1.0/src/sparx/surrogate.py +171 -0
  48. sparxml-0.1.0/src/sparx/tasks.py +121 -0
  49. sparxml-0.1.0/src/sparxml.egg-info/PKG-INFO +230 -0
  50. sparxml-0.1.0/src/sparxml.egg-info/SOURCES.txt +87 -0
  51. sparxml-0.1.0/src/sparxml.egg-info/dependency_links.txt +1 -0
  52. sparxml-0.1.0/src/sparxml.egg-info/requires.txt +30 -0
  53. sparxml-0.1.0/src/sparxml.egg-info/top_level.txt +1 -0
  54. sparxml-0.1.0/tests/test_connectome.py +173 -0
  55. sparxml-0.1.0/tests/test_datasets.py +96 -0
  56. sparxml-0.1.0/tests/test_delays.py +106 -0
  57. sparxml-0.1.0/tests/test_diffusion.py +199 -0
  58. sparxml-0.1.0/tests/test_distributed.py +59 -0
  59. sparxml-0.1.0/tests/test_dynamics.py +279 -0
  60. sparxml-0.1.0/tests/test_encode.py +106 -0
  61. sparxml-0.1.0/tests/test_examples.py +24 -0
  62. sparxml-0.1.0/tests/test_fast_weights.py +372 -0
  63. sparxml-0.1.0/tests/test_fixtures.py +30 -0
  64. sparxml-0.1.0/tests/test_flynn.py +110 -0
  65. sparxml-0.1.0/tests/test_graph.py +686 -0
  66. sparxml-0.1.0/tests/test_graph_distributed.py +74 -0
  67. sparxml-0.1.0/tests/test_homeostasis.py +266 -0
  68. sparxml-0.1.0/tests/test_learn.py +703 -0
  69. sparxml-0.1.0/tests/test_losses.py +135 -0
  70. sparxml-0.1.0/tests/test_microcircuit.py +144 -0
  71. sparxml-0.1.0/tests/test_ml.py +473 -0
  72. sparxml-0.1.0/tests/test_models.py +206 -0
  73. sparxml-0.1.0/tests/test_nir.py +250 -0
  74. sparxml-0.1.0/tests/test_nn.py +254 -0
  75. sparxml-0.1.0/tests/test_objectives.py +724 -0
  76. sparxml-0.1.0/tests/test_package.py +59 -0
  77. sparxml-0.1.0/tests/test_parallel.py +92 -0
  78. sparxml-0.1.0/tests/test_plasticity.py +151 -0
  79. sparxml-0.1.0/tests/test_predictive.py +128 -0
  80. sparxml-0.1.0/tests/test_readme.py +136 -0
  81. sparxml-0.1.0/tests/test_recipe.py +73 -0
  82. sparxml-0.1.0/tests/test_reference.py +349 -0
  83. sparxml-0.1.0/tests/test_research.py +44 -0
  84. sparxml-0.1.0/tests/test_serve.py +198 -0
  85. sparxml-0.1.0/tests/test_signalling.py +431 -0
  86. sparxml-0.1.0/tests/test_simulators.py +388 -0
  87. sparxml-0.1.0/tests/test_spiketrains.py +44 -0
  88. sparxml-0.1.0/tests/test_surrogate.py +98 -0
  89. sparxml-0.1.0/tests/test_tutorials.py +60 -0
sparxml-0.1.0/LICENSE ADDED
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2024-2026 Ashish Kumar Singh
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
sparxml-0.1.0/PKG-INFO ADDED
@@ -0,0 +1,230 @@
1
+ Metadata-Version: 2.4
2
+ Name: sparxml
3
+ Version: 0.1.0
4
+ Summary: Spiking neural networks in JAX and Flax
5
+ Author-email: Ashish Kumar Singh <ashishkmr472@gmail.com>
6
+ License-Expression: MIT
7
+ Project-URL: Homepage, https://sparxml.dev
8
+ Project-URL: Repository, https://github.com/AshishKumar4/sparx
9
+ Project-URL: Documentation, https://sparxml.dev/docs/
10
+ Project-URL: Issues, https://github.com/AshishKumar4/sparx/issues
11
+ Project-URL: Changelog, https://github.com/AshishKumar4/sparx/blob/main/CHANGELOG.md
12
+ Keywords: jax,flax,spiking-neural-networks,neuromorphic,computational-neuroscience,connectomics,dew
13
+ Classifier: Development Status :: 3 - Alpha
14
+ Classifier: Intended Audience :: Science/Research
15
+ Classifier: Intended Audience :: Developers
16
+ Classifier: Programming Language :: Python :: 3
17
+ Classifier: Programming Language :: Python :: 3.12
18
+ Classifier: Programming Language :: Python :: 3.13
19
+ Classifier: Programming Language :: Python :: 3.14
20
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
21
+ Classifier: Topic :: Scientific/Engineering :: Bio-Informatics
22
+ Classifier: Typing :: Typed
23
+ Requires-Python: >=3.12
24
+ Description-Content-Type: text/markdown
25
+ License-File: LICENSE
26
+ Requires-Dist: dewml<0.2,>=0.1.0
27
+ Requires-Dist: jax
28
+ Requires-Dist: flax
29
+ Requires-Dist: optax
30
+ Requires-Dist: numpy
31
+ Provides-Extra: cuda12
32
+ Requires-Dist: dewml[cuda12]<0.2,>=0.1.0; extra == "cuda12"
33
+ Provides-Extra: cuda13
34
+ Requires-Dist: dewml[cuda13]<0.2,>=0.1.0; extra == "cuda13"
35
+ Provides-Extra: tpu
36
+ Requires-Dist: dewml[tpu]<0.2,>=0.1.0; extra == "tpu"
37
+ Provides-Extra: datasets
38
+ Requires-Dist: h5py; extra == "datasets"
39
+ Provides-Extra: connectome
40
+ Requires-Dist: pyarrow; extra == "connectome"
41
+ Provides-Extra: nir
42
+ Requires-Dist: nir; extra == "nir"
43
+ Provides-Extra: test
44
+ Requires-Dist: pytest; extra == "test"
45
+ Requires-Dist: ruff==0.14.3; extra == "test"
46
+ Requires-Dist: h5py; extra == "test"
47
+ Requires-Dist: nir; extra == "test"
48
+ Requires-Dist: pyarrow; extra == "test"
49
+ Dynamic: license-file
50
+
51
+ <picture>
52
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/banner-dark.svg">
53
+ <img alt="sparx: spiking neural networks in JAX" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/banner-light.svg" width="100%">
54
+ </picture>
55
+
56
+ sparx trains spiking neural networks and simulates circuits of biological neurons, in JAX. Its spiking layers are Flax modules, so they train with optax or [dew](https://github.com/AshishKumar4/dew) and work with `jit`, `grad`, `vmap` and sharding. The same neuron models also run in millivolts and milliseconds, wired into circuits and whole connectomes, and there they match NEST and Brian2.
57
+
58
+ [sparxml.dev](https://sparxml.dev): a course from one neuron to a spiking network that drives from events, the docs and the API reference.
59
+
60
+ [Guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md) · [Train, serve and export](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/train-and-deploy.md) · [From NEST and Brian2](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/nest-and-brian2.md) · [Fit a circuit](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/fit-a-circuit.md) · [Units](https://github.com/AshishKumar4/sparx/blob/main/docs/units.md) · [Status](https://github.com/AshishKumar4/sparx/blob/main/docs/status.md) · [Fidelity ledger](https://github.com/AshishKumar4/sparx/blob/main/docs/fidelity.md) · [Design](https://github.com/AshishKumar4/sparx/blob/main/docs/design.md) · [Performance](https://github.com/AshishKumar4/sparx/blob/main/docs/performance.md)
61
+
62
+ ## Install
63
+
64
+ sparx is published as `sparxml` and imported as `sparx`, as [dew](https://github.com/AshishKumar4/dew) is published as `dewml` and imported as `dew`:
65
+
66
+ | Where | Install |
67
+ | --- | --- |
68
+ | CPU | `pip install sparxml` |
69
+ | NVIDIA GPU (CUDA 12 or 13) | `pip install "sparxml[cuda12]"` or `"sparxml[cuda13]"` |
70
+ | TPU | `pip install "sparxml[tpu]"` |
71
+ | SHD and MNIST readers, connectome tables, NIR | add the extras `datasets`, `connectome` and `nir` |
72
+
73
+ From a clone, on the dew commit and the jax build the tests run on:
74
+
75
+ ```bash
76
+ git clone https://github.com/AshishKumar4/sparx.git && cd sparx
77
+ uv venv --python 3.12 && source .venv/bin/activate
78
+ uv pip install -e ".[datasets,test]" -c constraints.txt && pytest -q
79
+ ```
80
+
81
+ sparx needs Python 3.12 or later, and CI tests it on 3.12, 3.13 and 3.14 on Linux and on 3.12 on macOS, with JAX 0.11.2 and Flax 0.12 on CPU. The API may change before 1.0.
82
+
83
+ ## A first network
84
+
85
+ ```python
86
+ import flax.linen as nn
87
+ import jax
88
+ import jax.numpy as jnp
89
+ import optax
90
+
91
+ import sparx
92
+
93
+
94
+ class Net(nn.Module):
95
+ @nn.compact
96
+ def __call__(self, spikes): # [T, B, 784]
97
+ x = sparx.nn.LIF(tau=2.0)(nn.Dense(256)(spikes))
98
+ return sparx.nn.LI(tau=2.0)(nn.Dense(10)(x)) # membrane [T, B, 10]
99
+
100
+
101
+ net = Net()
102
+ images = jax.random.uniform(jax.random.key(0), (32, 784)) # intensities in [0, 1]
103
+ labels = jnp.zeros(32, jnp.int32)
104
+ spikes = sparx.encode.RateEncoder(steps=8)(jax.random.key(1), images) # [8, 32, 784]
105
+ params = net.init(jax.random.key(2), spikes)
106
+
107
+
108
+ def loss(params):
109
+ logits = jnp.mean(net.apply(params, spikes), axis=0) # mean membrane over time
110
+ return optax.softmax_cross_entropy_with_integer_labels(logits, labels).mean()
111
+
112
+
113
+ grads = jax.grad(loss)(params)
114
+ ```
115
+
116
+ `LIF` turns input currents into spikes, exactly 0 or 1, and `LI` integrates them into a membrane, which is the readout. Everything else is Flax and optax. [`examples/train_mnist.py`](https://github.com/AshishKumar4/sparx/blob/main/examples/train_mnist.py) trains a network like this on MNIST under dew's `Trainer`.
117
+
118
+ <picture>
119
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-dark.webp">
120
+ <img alt="A spiking classifier learning MNIST: a test digit, the hidden layer's spikes for it, the output spike counts and the test accuracy rising over 400 training steps" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-light.webp" width="100%">
121
+ </picture>
122
+
123
+ A 784-200-10 network of LIF neurons learning MNIST by surrogate gradients. It follows one test digit through 400 training steps, showing the hidden layer's spikes, the ten output neurons' spike counts and the test accuracy, which reaches 92.8%.
124
+
125
+ ## How sparx fits together
126
+
127
+ <picture>
128
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/architecture-dark.svg">
129
+ <img alt="Training tools and simulation tools both run neuron models through one protocol" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/architecture-light.svg" width="100%">
130
+ </picture>
131
+
132
+ Every neuron model implements one protocol, `init_state` and `step`, and `run` scans a model over time. Dimensionless cells serve deep learning, and physical models in mV and ms serve neuroscience. The two halves mix: a layer can hold a physical model (`nn.Dynamics(AdEx())`), and a simulated population can hold a dimensionless cell.
133
+
134
+ ## A spiking layer over time
135
+
136
+ <picture>
137
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/over_time-dark.svg">
138
+ <img alt="A layer scans one step function over time; a LIF membrane rises to threshold and resets at each spike; the spike's gradient is a smooth surrogate" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/over_time-light.svg" width="100%">
139
+ </picture>
140
+
141
+ Arrays are time-major, `[T, B, ...]`. Synaptic layers run over all steps in one matrix product, and only the neurons step through time. A spike is a step function with zero derivative almost everywhere, so the backward pass uses a surrogate's slope instead. The layers are LIF, IF, LI, current-based synaptic LIF, adaptive LIF (ALIF), rate units, parallel spiking neurons and dense layers with learned delays ([guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#neuron-layers)).
142
+
143
+ ## Recurrence and fast weights
144
+
145
+ <picture>
146
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/recurrence-dark.svg">
147
+ <img alt="A recurrent cell sends its output back through a dense, sparse or delayed wiring, with optional fast weights from a Hebbian trace" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/recurrence-light.svg" width="100%">
148
+ </picture>
149
+
150
+ `RecurrentCell` feeds any model's output back through a wiring: dense, an edge list such as a connectome's, or edges with their own delays. Fast weights add a Hebbian trace that each sequence writes as it runs (Miconi et al. 2018, 2019). Against Miconi et al.'s four networks in PyTorch, activity, traces and gradients agree within 5e-14. On their pattern completion task, the plastic network gets 0.3% of the zeroed bits wrong, and the same network without fast weights 50.1%.
151
+
152
+ ## Learning rules
153
+
154
+ <picture>
155
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/learning-dark.svg">
156
+ <img alt="Nine ways to train a spiking network in sparx, each with a schematic of its learning signal" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/learning-light.svg" width="100%">
157
+ </picture>
158
+
159
+ Each rule is checked against what defines it. e-prop meets the two identities its authors verify their code with, OTTT matches their PyTorch modules, PC-ALM matches their JAX reference to 5e-14, and conversion matches their toolbox. REINFORCE is checked on enumerated trajectories, and exact spike times against finite differences. The [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#learning-rules) describes each rule.
160
+
161
+ ## Simulating circuits
162
+
163
+ <picture>
164
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/circuits-dark.svg">
165
+ <img alt="Two populations connected by excitatory and inhibitory projections with delays, driven by Poisson input and simulated in chunks" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/circuits-light.svg" width="100%">
166
+ </picture>
167
+
168
+ ```python
169
+ import jax
170
+ from sparx.graph import PopulationRate, SpikeRaster, simulate
171
+ from sparx.graph.models import brunel
172
+
173
+ network = brunel(250, g=5.0, eta=2.0) # 1,250 LIF neurons; brunel(2500) is the paper's 12,500
174
+ result = simulate(network, network.init(jax.random.key(0)), duration=200.0, key=jax.random.key(1),
175
+ monitors={"spikes": SpikeRaster("e"), "rate": PopulationRate("e")})
176
+ spikes = result.records["spikes"] # [2000, 1000]: one row of booleans per 0.1 ms step
177
+ ```
178
+
179
+ <picture>
180
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/network-dark.webp">
181
+ <img alt="A raster of 200 neurons of Brunel's network firing irregularly over 300 ms, with the population rate below" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/network-light.webp" width="100%">
182
+ </picture>
183
+
184
+ Brunel's balanced network at the paper's size, 10,000 excitatory and 2,500 inhibitory LIF neurons, in its asynchronous irregular regime at 37 Hz.
185
+
186
+ The physical models match NEST 3.10 and Brian2 2.10 spike for spike where the dynamics are deterministic, for the integration scheme, dtype and step each check states ([status](https://github.com/AshishKumar4/sparx/blob/main/docs/status.md#capabilities-and-limits)), and in rate, irregularity and synchrony where they are chaotic. Potjans and Diesmann's cortical microcircuit, built as its reference implementation builds it, fires spike for spike with NEST on the same network ([from NEST and Brian2](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/nest-and-brian2.md)). On a 4-core CPU, sparx simulates a second of Brunel's network in 9.6 s, NEST in 7.5 s and Brian2 in 11.8 s ([performance](https://github.com/AshishKumar4/sparx/blob/main/docs/performance.md#against-nest-and-brian2)). Populations can hold graded neurons and connect through stochastic release, gap junctions and neuromodulators. Projections can carry STDP, triplet STDP, dopamine-modulated STDP and short-term plasticity ([guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#simulating-circuits)).
187
+
188
+ ## Connectomes
189
+
190
+ `sparx.graph.connectome` builds Shiu et al.'s (2024) model of the whole fly brain from FlyWire. It reproduces their published runs, with a rate correlation of 0.999 and the motor neuron MN9 at 67.1 Hz against their 67.0 ± 6.6, at about 30 s per simulated second on 4 CPU cores. `FLYNN` (Wang and Chen 2026) trains a connectome as a recurrent rate network with one learned weight per synapse; against their PyTorch cell its activity and gradients agree within 1e-15.
191
+
192
+ ## RNeuralNet
193
+
194
+ <picture>
195
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/messages-dark.webp">
196
+ <img alt="Messages travelling along the connections of a small RNeuralNet, each connection with its own delay" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/messages-light.webp" width="100%">
197
+ </picture>
198
+
199
+ `sparx.learn.RNeuralNet` rebuilds RNeuralNet-Research (2018), an early project of the author's, deterministically. Graded neurons sit on a random graph, each connection delivers its messages after its own delay, and a reward spreads backward by a softmax of activity. Compiled and run in a fixed order, the original C++ and sparx agree within 7.2e-7. On a delayed cue-order task, REINFORCE through the same network learns the task on four of five seeds, and the reward-diffusion rule never changes the network's choice. AGREL's update, a signed error sent back from the chosen output through the weights, learns it on the same four seeds; the same error spread by the original's shares does not.
200
+
201
+ ## Training on dew
202
+
203
+ <picture>
204
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-dark.svg">
205
+ <img alt="A Flax model and a sparx objective go to dew's trainer, which writes a run record that reloads, serves and exports" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-light.svg" width="100%">
206
+ </picture>
207
+
208
+ The `Trainer` from [dew](https://github.com/AshishKumar4/dew) runs sparx's objectives for classification, activity fitting, e-prop, predictive coding and RNeuralNet's rewards. A run's record names every class by import path, so `dew.pipeline("runs/shd", trust=("sparx",))` loads a trained network in a new process, and `sparx.serve.StreamServer` serves it to many streams at once. The [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#training-on-dew) has a full SHD script.
209
+
210
+ ## Results
211
+
212
+ | Task | Network | Test accuracy |
213
+ | --- | --- | --- |
214
+ | MNIST, rate-coded, 8 steps | 784-512-512 LIF | 97.5% after 2 epochs |
215
+ | SHD, Hammouamri et al.'s recipe, 150 epochs, three seeds | 140-256-256 LIF with learned delays | 93.99 ± 0.29% at the last epoch (their code on the same GPU: 93.89 ± 0.26%) |
216
+ | SHD, 140 channels | 140-128 ALIF, with and without learned delays | 74.6% and 64.5% |
217
+ | Fashion-MNIST, Seely and Gould's headline cell | ReLU residual MLP, depth 32 | PC-ALM 75.1%, PC 62.2%, backpropagation 77.8% |
218
+ | Pattern completion, Miconi et al.'s task | plastic recurrent network | 0.3% of bits wrong; 50.1% without fast weights |
219
+
220
+ The SHD row is the full recipe beside the authors' code, both on an A100, three seeds each ([research/shd](https://github.com/AshishKumar4/sparx/blob/main/research/shd/README.md)). Both train on every training recording and score the test set after each epoch. The paper's 95.07 ± 0.24% (a 95% confidence interval over ten runs) is the best epoch on the test set, which chooses with the test set; here, as mean ± standard deviation over three seeds, their code's best epoch is 95.17 ± 0.61% and sparx's 94.96 ± 0.89%. With a tenth of the training set held out to choose the epoch, sparx scores 94.14 ± 0.98% on test. The other rows are short, untuned runs on a 4-core CPU; the [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#results-in-detail) gives their commands, times and comparisons.
221
+
222
+ ## Correctness
223
+
224
+ Every model is checked against a reference: a float64 loop of its equations, the original authors' code, or NEST and Brian2. [docs/fidelity.md](https://github.com/AshishKumar4/sparx/blob/main/docs/fidelity.md) lists each model's reference, the check, the observed error and every known difference. `pytest -q` runs all of it on CPU in about 16 minutes.
225
+
226
+ [`tools/make_figures.py`](https://github.com/AshishKumar4/sparx/blob/main/tools/make_figures.py) draws the banner and diagrams, and [`tools/make_clips.py`](https://github.com/AshishKumar4/sparx/blob/main/tools/make_clips.py) renders the clips. The spikes in the banner and the clips come from sparx runs.
227
+
228
+ ## License
229
+
230
+ MIT
@@ -0,0 +1,180 @@
1
+ <picture>
2
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/banner-dark.svg">
3
+ <img alt="sparx: spiking neural networks in JAX" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/banner-light.svg" width="100%">
4
+ </picture>
5
+
6
+ sparx trains spiking neural networks and simulates circuits of biological neurons, in JAX. Its spiking layers are Flax modules, so they train with optax or [dew](https://github.com/AshishKumar4/dew) and work with `jit`, `grad`, `vmap` and sharding. The same neuron models also run in millivolts and milliseconds, wired into circuits and whole connectomes, and there they match NEST and Brian2.
7
+
8
+ [sparxml.dev](https://sparxml.dev): a course from one neuron to a spiking network that drives from events, the docs and the API reference.
9
+
10
+ [Guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md) · [Train, serve and export](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/train-and-deploy.md) · [From NEST and Brian2](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/nest-and-brian2.md) · [Fit a circuit](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/fit-a-circuit.md) · [Units](https://github.com/AshishKumar4/sparx/blob/main/docs/units.md) · [Status](https://github.com/AshishKumar4/sparx/blob/main/docs/status.md) · [Fidelity ledger](https://github.com/AshishKumar4/sparx/blob/main/docs/fidelity.md) · [Design](https://github.com/AshishKumar4/sparx/blob/main/docs/design.md) · [Performance](https://github.com/AshishKumar4/sparx/blob/main/docs/performance.md)
11
+
12
+ ## Install
13
+
14
+ sparx is published as `sparxml` and imported as `sparx`, as [dew](https://github.com/AshishKumar4/dew) is published as `dewml` and imported as `dew`:
15
+
16
+ | Where | Install |
17
+ | --- | --- |
18
+ | CPU | `pip install sparxml` |
19
+ | NVIDIA GPU (CUDA 12 or 13) | `pip install "sparxml[cuda12]"` or `"sparxml[cuda13]"` |
20
+ | TPU | `pip install "sparxml[tpu]"` |
21
+ | SHD and MNIST readers, connectome tables, NIR | add the extras `datasets`, `connectome` and `nir` |
22
+
23
+ From a clone, on the dew commit and the jax build the tests run on:
24
+
25
+ ```bash
26
+ git clone https://github.com/AshishKumar4/sparx.git && cd sparx
27
+ uv venv --python 3.12 && source .venv/bin/activate
28
+ uv pip install -e ".[datasets,test]" -c constraints.txt && pytest -q
29
+ ```
30
+
31
+ sparx needs Python 3.12 or later, and CI tests it on 3.12, 3.13 and 3.14 on Linux and on 3.12 on macOS, with JAX 0.11.2 and Flax 0.12 on CPU. The API may change before 1.0.
32
+
33
+ ## A first network
34
+
35
+ ```python
36
+ import flax.linen as nn
37
+ import jax
38
+ import jax.numpy as jnp
39
+ import optax
40
+
41
+ import sparx
42
+
43
+
44
+ class Net(nn.Module):
45
+ @nn.compact
46
+ def __call__(self, spikes): # [T, B, 784]
47
+ x = sparx.nn.LIF(tau=2.0)(nn.Dense(256)(spikes))
48
+ return sparx.nn.LI(tau=2.0)(nn.Dense(10)(x)) # membrane [T, B, 10]
49
+
50
+
51
+ net = Net()
52
+ images = jax.random.uniform(jax.random.key(0), (32, 784)) # intensities in [0, 1]
53
+ labels = jnp.zeros(32, jnp.int32)
54
+ spikes = sparx.encode.RateEncoder(steps=8)(jax.random.key(1), images) # [8, 32, 784]
55
+ params = net.init(jax.random.key(2), spikes)
56
+
57
+
58
+ def loss(params):
59
+ logits = jnp.mean(net.apply(params, spikes), axis=0) # mean membrane over time
60
+ return optax.softmax_cross_entropy_with_integer_labels(logits, labels).mean()
61
+
62
+
63
+ grads = jax.grad(loss)(params)
64
+ ```
65
+
66
+ `LIF` turns input currents into spikes, exactly 0 or 1, and `LI` integrates them into a membrane, which is the readout. Everything else is Flax and optax. [`examples/train_mnist.py`](https://github.com/AshishKumar4/sparx/blob/main/examples/train_mnist.py) trains a network like this on MNIST under dew's `Trainer`.
67
+
68
+ <picture>
69
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-dark.webp">
70
+ <img alt="A spiking classifier learning MNIST: a test digit, the hidden layer's spikes for it, the output spike counts and the test accuracy rising over 400 training steps" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-light.webp" width="100%">
71
+ </picture>
72
+
73
+ A 784-200-10 network of LIF neurons learning MNIST by surrogate gradients. It follows one test digit through 400 training steps, showing the hidden layer's spikes, the ten output neurons' spike counts and the test accuracy, which reaches 92.8%.
74
+
75
+ ## How sparx fits together
76
+
77
+ <picture>
78
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/architecture-dark.svg">
79
+ <img alt="Training tools and simulation tools both run neuron models through one protocol" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/architecture-light.svg" width="100%">
80
+ </picture>
81
+
82
+ Every neuron model implements one protocol, `init_state` and `step`, and `run` scans a model over time. Dimensionless cells serve deep learning, and physical models in mV and ms serve neuroscience. The two halves mix: a layer can hold a physical model (`nn.Dynamics(AdEx())`), and a simulated population can hold a dimensionless cell.
83
+
84
+ ## A spiking layer over time
85
+
86
+ <picture>
87
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/over_time-dark.svg">
88
+ <img alt="A layer scans one step function over time; a LIF membrane rises to threshold and resets at each spike; the spike's gradient is a smooth surrogate" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/over_time-light.svg" width="100%">
89
+ </picture>
90
+
91
+ Arrays are time-major, `[T, B, ...]`. Synaptic layers run over all steps in one matrix product, and only the neurons step through time. A spike is a step function with zero derivative almost everywhere, so the backward pass uses a surrogate's slope instead. The layers are LIF, IF, LI, current-based synaptic LIF, adaptive LIF (ALIF), rate units, parallel spiking neurons and dense layers with learned delays ([guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#neuron-layers)).
92
+
93
+ ## Recurrence and fast weights
94
+
95
+ <picture>
96
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/recurrence-dark.svg">
97
+ <img alt="A recurrent cell sends its output back through a dense, sparse or delayed wiring, with optional fast weights from a Hebbian trace" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/recurrence-light.svg" width="100%">
98
+ </picture>
99
+
100
+ `RecurrentCell` feeds any model's output back through a wiring: dense, an edge list such as a connectome's, or edges with their own delays. Fast weights add a Hebbian trace that each sequence writes as it runs (Miconi et al. 2018, 2019). Against Miconi et al.'s four networks in PyTorch, activity, traces and gradients agree within 5e-14. On their pattern completion task, the plastic network gets 0.3% of the zeroed bits wrong, and the same network without fast weights 50.1%.
101
+
102
+ ## Learning rules
103
+
104
+ <picture>
105
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/learning-dark.svg">
106
+ <img alt="Nine ways to train a spiking network in sparx, each with a schematic of its learning signal" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/learning-light.svg" width="100%">
107
+ </picture>
108
+
109
+ Each rule is checked against what defines it. e-prop meets the two identities its authors verify their code with, OTTT matches their PyTorch modules, PC-ALM matches their JAX reference to 5e-14, and conversion matches their toolbox. REINFORCE is checked on enumerated trajectories, and exact spike times against finite differences. The [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#learning-rules) describes each rule.
110
+
111
+ ## Simulating circuits
112
+
113
+ <picture>
114
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/circuits-dark.svg">
115
+ <img alt="Two populations connected by excitatory and inhibitory projections with delays, driven by Poisson input and simulated in chunks" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/circuits-light.svg" width="100%">
116
+ </picture>
117
+
118
+ ```python
119
+ import jax
120
+ from sparx.graph import PopulationRate, SpikeRaster, simulate
121
+ from sparx.graph.models import brunel
122
+
123
+ network = brunel(250, g=5.0, eta=2.0) # 1,250 LIF neurons; brunel(2500) is the paper's 12,500
124
+ result = simulate(network, network.init(jax.random.key(0)), duration=200.0, key=jax.random.key(1),
125
+ monitors={"spikes": SpikeRaster("e"), "rate": PopulationRate("e")})
126
+ spikes = result.records["spikes"] # [2000, 1000]: one row of booleans per 0.1 ms step
127
+ ```
128
+
129
+ <picture>
130
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/network-dark.webp">
131
+ <img alt="A raster of 200 neurons of Brunel's network firing irregularly over 300 ms, with the population rate below" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/network-light.webp" width="100%">
132
+ </picture>
133
+
134
+ Brunel's balanced network at the paper's size, 10,000 excitatory and 2,500 inhibitory LIF neurons, in its asynchronous irregular regime at 37 Hz.
135
+
136
+ The physical models match NEST 3.10 and Brian2 2.10 spike for spike where the dynamics are deterministic, for the integration scheme, dtype and step each check states ([status](https://github.com/AshishKumar4/sparx/blob/main/docs/status.md#capabilities-and-limits)), and in rate, irregularity and synchrony where they are chaotic. Potjans and Diesmann's cortical microcircuit, built as its reference implementation builds it, fires spike for spike with NEST on the same network ([from NEST and Brian2](https://github.com/AshishKumar4/sparx/blob/main/docs/tutorials/nest-and-brian2.md)). On a 4-core CPU, sparx simulates a second of Brunel's network in 9.6 s, NEST in 7.5 s and Brian2 in 11.8 s ([performance](https://github.com/AshishKumar4/sparx/blob/main/docs/performance.md#against-nest-and-brian2)). Populations can hold graded neurons and connect through stochastic release, gap junctions and neuromodulators. Projections can carry STDP, triplet STDP, dopamine-modulated STDP and short-term plasticity ([guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#simulating-circuits)).
137
+
138
+ ## Connectomes
139
+
140
+ `sparx.graph.connectome` builds Shiu et al.'s (2024) model of the whole fly brain from FlyWire. It reproduces their published runs, with a rate correlation of 0.999 and the motor neuron MN9 at 67.1 Hz against their 67.0 ± 6.6, at about 30 s per simulated second on 4 CPU cores. `FLYNN` (Wang and Chen 2026) trains a connectome as a recurrent rate network with one learned weight per synapse; against their PyTorch cell its activity and gradients agree within 1e-15.
141
+
142
+ ## RNeuralNet
143
+
144
+ <picture>
145
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/messages-dark.webp">
146
+ <img alt="Messages travelling along the connections of a small RNeuralNet, each connection with its own delay" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/messages-light.webp" width="100%">
147
+ </picture>
148
+
149
+ `sparx.learn.RNeuralNet` rebuilds RNeuralNet-Research (2018), an early project of the author's, deterministically. Graded neurons sit on a random graph, each connection delivers its messages after its own delay, and a reward spreads backward by a softmax of activity. Compiled and run in a fixed order, the original C++ and sparx agree within 7.2e-7. On a delayed cue-order task, REINFORCE through the same network learns the task on four of five seeds, and the reward-diffusion rule never changes the network's choice. AGREL's update, a signed error sent back from the chosen output through the weights, learns it on the same four seeds; the same error spread by the original's shares does not.
150
+
151
+ ## Training on dew
152
+
153
+ <picture>
154
+ <source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-dark.svg">
155
+ <img alt="A Flax model and a sparx objective go to dew's trainer, which writes a run record that reloads, serves and exports" src="https://raw.githubusercontent.com/AshishKumar4/sparx/main/docs/assets/training-light.svg" width="100%">
156
+ </picture>
157
+
158
+ The `Trainer` from [dew](https://github.com/AshishKumar4/dew) runs sparx's objectives for classification, activity fitting, e-prop, predictive coding and RNeuralNet's rewards. A run's record names every class by import path, so `dew.pipeline("runs/shd", trust=("sparx",))` loads a trained network in a new process, and `sparx.serve.StreamServer` serves it to many streams at once. The [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#training-on-dew) has a full SHD script.
159
+
160
+ ## Results
161
+
162
+ | Task | Network | Test accuracy |
163
+ | --- | --- | --- |
164
+ | MNIST, rate-coded, 8 steps | 784-512-512 LIF | 97.5% after 2 epochs |
165
+ | SHD, Hammouamri et al.'s recipe, 150 epochs, three seeds | 140-256-256 LIF with learned delays | 93.99 ± 0.29% at the last epoch (their code on the same GPU: 93.89 ± 0.26%) |
166
+ | SHD, 140 channels | 140-128 ALIF, with and without learned delays | 74.6% and 64.5% |
167
+ | Fashion-MNIST, Seely and Gould's headline cell | ReLU residual MLP, depth 32 | PC-ALM 75.1%, PC 62.2%, backpropagation 77.8% |
168
+ | Pattern completion, Miconi et al.'s task | plastic recurrent network | 0.3% of bits wrong; 50.1% without fast weights |
169
+
170
+ The SHD row is the full recipe beside the authors' code, both on an A100, three seeds each ([research/shd](https://github.com/AshishKumar4/sparx/blob/main/research/shd/README.md)). Both train on every training recording and score the test set after each epoch. The paper's 95.07 ± 0.24% (a 95% confidence interval over ten runs) is the best epoch on the test set, which chooses with the test set; here, as mean ± standard deviation over three seeds, their code's best epoch is 95.17 ± 0.61% and sparx's 94.96 ± 0.89%. With a tenth of the training set held out to choose the epoch, sparx scores 94.14 ± 0.98% on test. The other rows are short, untuned runs on a 4-core CPU; the [guide](https://github.com/AshishKumar4/sparx/blob/main/docs/guide.md#results-in-detail) gives their commands, times and comparisons.
171
+
172
+ ## Correctness
173
+
174
+ Every model is checked against a reference: a float64 loop of its equations, the original authors' code, or NEST and Brian2. [docs/fidelity.md](https://github.com/AshishKumar4/sparx/blob/main/docs/fidelity.md) lists each model's reference, the check, the observed error and every known difference. `pytest -q` runs all of it on CPU in about 16 minutes.
175
+
176
+ [`tools/make_figures.py`](https://github.com/AshishKumar4/sparx/blob/main/tools/make_figures.py) draws the banner and diagrams, and [`tools/make_clips.py`](https://github.com/AshishKumar4/sparx/blob/main/tools/make_clips.py) renders the clips. The spikes in the banner and the clips come from sparx runs.
177
+
178
+ ## License
179
+
180
+ MIT
@@ -0,0 +1,116 @@
1
+ [build-system]
2
+ requires = ["setuptools>=77", "wheel"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "sparxml"
7
+ dynamic = ["version"]
8
+ description = "Spiking neural networks in JAX and Flax"
9
+ readme = "README.md"
10
+ requires-python = ">=3.12"
11
+ authors = [
12
+ { name="Ashish Kumar Singh", email="ashishkmr472@gmail.com" }
13
+ ]
14
+ license = "MIT"
15
+ license-files = ["LICENSE"]
16
+ keywords = ["jax", "flax", "spiking-neural-networks", "neuromorphic", "computational-neuroscience",
17
+ "connectomics", "dew"]
18
+ classifiers = [
19
+ "Development Status :: 3 - Alpha",
20
+ "Intended Audience :: Science/Research",
21
+ "Intended Audience :: Developers",
22
+ "Programming Language :: Python :: 3",
23
+ "Programming Language :: Python :: 3.12",
24
+ "Programming Language :: Python :: 3.13",
25
+ "Programming Language :: Python :: 3.14",
26
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
27
+ "Topic :: Scientific/Engineering :: Bio-Informatics",
28
+ "Typing :: Typed",
29
+ ]
30
+ dependencies = [
31
+ # Dew is the platform sparx trains, distributes, checkpoints and serves
32
+ # through (docs/design.md), and bounds jax, flax and optax to the versions
33
+ # it supports. Development and CI install the dew commit constraints.txt
34
+ # names.
35
+ "dewml>=0.1.0,<0.2",
36
+ "jax",
37
+ "flax",
38
+ "optax",
39
+ "numpy",
40
+ ]
41
+
42
+ [project.urls]
43
+ Homepage = "https://sparxml.dev"
44
+ Repository = "https://github.com/AshishKumar4/sparx"
45
+ Documentation = "https://sparxml.dev/docs/"
46
+ Issues = "https://github.com/AshishKumar4/sparx/issues"
47
+ Changelog = "https://github.com/AshishKumar4/sparx/blob/main/CHANGELOG.md"
48
+
49
+ [project.optional-dependencies]
50
+ # The accelerator builds of jax, as dew installs them.
51
+ cuda12 = ["dewml[cuda12]>=0.1.0,<0.2"]
52
+ cuda13 = ["dewml[cuda13]>=0.1.0,<0.2"]
53
+ tpu = ["dewml[tpu]>=0.1.0,<0.2"]
54
+ # sparx.datasets reads the published HDF5 files.
55
+ datasets = ["h5py"]
56
+ # sparx.graph.connectome reads connectome tables (CSV and parquet).
57
+ connectome = ["pyarrow"]
58
+ # sparx.nir exchanges networks through NIR.
59
+ nir = ["nir"]
60
+ test = ["pytest", "ruff==0.14.3", "h5py", "nir", "pyarrow"]
61
+
62
+ [tool.setuptools.dynamic]
63
+ version = { attr = "sparx.__version__" }
64
+
65
+ [tool.setuptools.packages.find]
66
+ where = ["src"]
67
+
68
+ [tool.setuptools.package-data]
69
+ sparx = ["py.typed"]
70
+
71
+ [tool.ruff]
72
+ target-version = "py312"
73
+ # dew's width (dew's pyproject.toml says why).
74
+ line-length = 110
75
+
76
+ [tool.ruff.lint]
77
+ # dew's rules. Scripts use stdout as their interface; src/sparx/ruff.toml adds
78
+ # T20 for the library, as dew's src/dew/ruff.toml does.
79
+ select = [
80
+ "E", "F", "W", "I", "UP", "B", "C4", "PERF", "PIE", "RET", "SIM", "PGH",
81
+ "ERA", "ANN401", "FBT003", "PLW0120", "RUF", "C901", "PLR0912", "PLR0915",
82
+ ]
83
+ allowed-confusables = ["×", "γ", "σ", "τ", "θ", "β", "α", "ρ"]
84
+
85
+ [tool.ruff.lint.per-file-ignores]
86
+ # dew's checker, verbatim at the commit sparx pins (CI checks it); its rule
87
+ # walkers are as complex as dew's own limit of 25 allows.
88
+ "tools/lint_slop.py" = ["C901"]
89
+
90
+ [tool.ruff.lint.pylint]
91
+ # dew's limits. Sparx keeps the stricter mccabe default of 10.
92
+ max-branches = 25
93
+ max-statements = 75
94
+
95
+ [tool.ruff.lint.isort]
96
+ known-first-party = ["sparx"]
97
+ combine-as-imports = true
98
+ split-on-trailing-comma = false
99
+
100
+ [tool.pyright]
101
+ # As dew's: sparx's own source and its checker, against the project venv. A
102
+ # worktree has no .venv of its own, so the gate passes --pythonpath instead.
103
+ include = ["src/sparx", "tools/lint_slop.py"]
104
+ pythonVersion = "3.12"
105
+ typeCheckingMode = "standard"
106
+ venvPath = "."
107
+ venv = ".venv"
108
+
109
+ [tool.pytest.ini_options]
110
+ testpaths = ["tests"]
111
+ pythonpath = ["src", "tests"]
112
+ filterwarnings = [
113
+ 'error::DeprecationWarning:sparx($|\.)',
114
+ 'error::PendingDeprecationWarning:sparx($|\.)',
115
+ 'error::FutureWarning:sparx($|\.)',
116
+ ]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1,69 @@
1
+ """Sparx: spiking neural networks in JAX and Flax, trained, simulated and served through dew.
2
+
3
+ Networks are Flax linen modules over time-major spike trains `[T, ...]`.
4
+
5
+ Deep spiking networks:
6
+
7
+ - `sparx.nn`: Flax layers over the neuron models, parallel spiking neurons and delayed synapses.
8
+ - `sparx.models`: architectures built from them (`SEWResNet`, `SpikingMLP`).
9
+ - `sparx.surrogate`: the spike and its surrogate gradients.
10
+ - `sparx.encode`: the encoders that turn data into spike trains.
11
+ - `sparx.losses` and `sparx.rates`: losses over time, firing-rate readouts and penalties.
12
+ - `sparx.learn`: rules beyond backpropagation through time (e-prop, OTTT, EventProp, conversion).
13
+
14
+ Circuits in physical units:
15
+
16
+ - `sparx.dynamics`: neuron, synapse and plasticity models as pure JAX, the dimensionless family deep
17
+ networks train with and the physical one, and `run`, which scans any of them over time.
18
+ - `sparx.graph`: populations and projections wired into a `Network`, `simulate`, and connectomes.
19
+ - `sparx.spiketrains`: statistics of and distances between recorded spike trains.
20
+
21
+ Around them:
22
+
23
+ - `sparx.objectives`: the objectives that train spiking networks under dew's `Trainer`, with
24
+ `sparx.metrics` (their accuracy), `sparx.tasks` (the trained classifier `dew.pipeline` loads) and
25
+ `sparx.config` (`SNNRunConfig`, the run a recipe trains and `run.json` records).
26
+ - `sparx.datasets`: spiking datasets as dew datasets (SHD).
27
+ - `sparx.serve`: `StreamServer`, many streaming sessions in one batch.
28
+ - `sparx.nir`: exchange through the Neuromorphic Intermediate Representation.
29
+
30
+ `sparx.graph`, `sparx.learn`, `sparx.objectives`, `sparx.metrics`,
31
+ `sparx.tasks`, `sparx.config`, `sparx.datasets`, `sparx.serve` and
32
+ `sparx.nir` load on first access (`sparx.graph.Network` after `import
33
+ sparx`). The graph, the objectives and the datasets import dew's trainer
34
+ and data stack, about 0.9 s on a 4-core CPU, which a script that only
35
+ trains a network in its own loop does not need.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import importlib
41
+ from types import ModuleType
42
+ from typing import TYPE_CHECKING
43
+
44
+ from sparx import dynamics, encode, losses, models, nn, rates, spiketrains, surrogate
45
+ from sparx.dynamics import run
46
+ from sparx.rates import firing_rates, rate_penalty
47
+ from sparx.surrogate import spike
48
+
49
+ if TYPE_CHECKING:
50
+ from sparx import config, datasets, graph, learn, metrics, nir, objectives, serve, tasks
51
+
52
+ __version__ = "0.1.0"
53
+
54
+ _LAZY = ("config", "datasets", "graph", "learn", "metrics", "nir", "objectives", "serve", "tasks")
55
+
56
+ __all__ = ["__version__", "config", "datasets", "dynamics", "encode", "firing_rates", "graph", "learn",
57
+ "losses", "metrics", "models", "nir", "nn", "objectives", "rate_penalty", "rates", "run", "serve",
58
+ "spike", "spiketrains", "surrogate", "tasks"]
59
+
60
+
61
+ def __getattr__(name: str) -> ModuleType:
62
+ # PEP 562: importing the submodule binds it on the package, so this runs once per name.
63
+ if name in _LAZY:
64
+ return importlib.import_module(f"sparx.{name}")
65
+ raise AttributeError(f"module 'sparx' has no attribute {name!r}")
66
+
67
+
68
+ def __dir__() -> list[str]:
69
+ return sorted(set(globals()) | set(_LAZY))