sparxml 0.1.0__py3-none-any.whl
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.
- sparx/__init__.py +69 -0
- sparx/config.py +102 -0
- sparx/datasets.py +246 -0
- sparx/dynamics/__init__.py +178 -0
- sparx/dynamics/core.py +400 -0
- sparx/dynamics/homeostasis.py +101 -0
- sparx/dynamics/ml.py +925 -0
- sparx/dynamics/neurons.py +648 -0
- sparx/dynamics/plasticity.py +395 -0
- sparx/dynamics/synapses.py +424 -0
- sparx/encode.py +167 -0
- sparx/graph/__init__.py +79 -0
- sparx/graph/connectivity.py +189 -0
- sparx/graph/connectome.py +324 -0
- sparx/graph/delivery.py +214 -0
- sparx/graph/models.py +253 -0
- sparx/graph/network.py +1284 -0
- sparx/graph/simulate.py +269 -0
- sparx/learn/__init__.py +58 -0
- sparx/learn/convert.py +282 -0
- sparx/learn/diffusion.py +338 -0
- sparx/learn/events.py +174 -0
- sparx/learn/online.py +439 -0
- sparx/learn/predictive.py +240 -0
- sparx/learn/reinforce.py +106 -0
- sparx/losses.py +146 -0
- sparx/metrics.py +40 -0
- sparx/models.py +339 -0
- sparx/nir.py +497 -0
- sparx/nn/__init__.py +57 -0
- sparx/nn/delays.py +112 -0
- sparx/nn/hebbian.py +114 -0
- sparx/nn/neurons.py +433 -0
- sparx/nn/parallel.py +156 -0
- sparx/nn/reshape.py +67 -0
- sparx/objectives.py +729 -0
- sparx/py.typed +0 -0
- sparx/rates.py +72 -0
- sparx/registry.py +52 -0
- sparx/serve.py +226 -0
- sparx/spiketrains.py +122 -0
- sparx/surrogate.py +171 -0
- sparx/tasks.py +121 -0
- sparxml-0.1.0.dist-info/METADATA +230 -0
- sparxml-0.1.0.dist-info/RECORD +48 -0
- sparxml-0.1.0.dist-info/WHEEL +5 -0
- sparxml-0.1.0.dist-info/licenses/LICENSE +21 -0
- sparxml-0.1.0.dist-info/top_level.txt +1 -0
sparx/__init__.py
ADDED
|
@@ -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))
|
sparx/config.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
"""The run a spiking classifier trains as, one typed record that `run.json` holds.
|
|
2
|
+
|
|
3
|
+
python recipes/snn/train.py --data.channels 140 --trainer.batch-size 64 --trainer.steps 3000 \\
|
|
4
|
+
--model.hidden 128 --model.classes 20 --model.delays 15 \\
|
|
5
|
+
--objective.schedules '{"sigma": {"class": "linear", "fields": {"peak": 7.5, "end": 0.5}}}'
|
|
6
|
+
dew train runs/<name>/run.json --trust sparx --set trainer.steps=6000 # the same run, trained on
|
|
7
|
+
|
|
8
|
+
`SNNRunConfig` is dew's `RunConfig` with the data, encoder and sample field
|
|
9
|
+
a spiking classifier reads. Its model is any spiking model over `[T, B,
|
|
10
|
+
channels]` by import path (`--model my_package.models:Net`), each of its
|
|
11
|
+
fields a flag and the neuron template a record (`--model.neuron '{"class":
|
|
12
|
+
"sparx.nn.neurons:LIF", "fields": {"tau": 3.0}}'`). The objective is
|
|
13
|
+
`SpikingClassifierObjective`, each of its keyword arguments a flag
|
|
14
|
+
(`--objective.readout max`), and `prepare` builds it around the model, the
|
|
15
|
+
sample field and the encoder. Parameter groups with optimizers of their own
|
|
16
|
+
are dew's (`--optim.param-groups`).
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import dataclasses
|
|
22
|
+
from pathlib import Path
|
|
23
|
+
from typing import TYPE_CHECKING
|
|
24
|
+
|
|
25
|
+
from dew.config import ModelConfig, ObjectiveConfig, OptimConfig, Prepared, RunConfig
|
|
26
|
+
from dew.data import Loading
|
|
27
|
+
from dew.inputs import Field
|
|
28
|
+
from dew.registry import import_path
|
|
29
|
+
|
|
30
|
+
from sparx.datasets import SHD, write_synthetic_shd
|
|
31
|
+
from sparx.encode import EventsEncoder, SpikeEncoder
|
|
32
|
+
from sparx.metrics import Accuracy
|
|
33
|
+
from sparx.models import SpikingMLP
|
|
34
|
+
from sparx.nn.neurons import ALIF
|
|
35
|
+
from sparx.objectives import SpikingClassifierObjective
|
|
36
|
+
from sparx.registry import spike_encoders
|
|
37
|
+
|
|
38
|
+
if TYPE_CHECKING:
|
|
39
|
+
# tyro reads the runtime annotation, a union of the encoders, and a type
|
|
40
|
+
# checker cannot read a variable in a type expression, so both get what
|
|
41
|
+
# they need, as dew.config's DataSpec does.
|
|
42
|
+
type EncoderSpec = SpikeEncoder
|
|
43
|
+
else:
|
|
44
|
+
EncoderSpec = spike_encoders.union
|
|
45
|
+
|
|
46
|
+
__all__ = ["SNNRunConfig"]
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclasses.dataclass(frozen=True)
|
|
50
|
+
class SNNRunConfig(RunConfig):
|
|
51
|
+
"""A run of `SpikingClassifierObjective` on SHD, with the encoder and the field it reads."""
|
|
52
|
+
|
|
53
|
+
objective: ObjectiveConfig = dataclasses.field(default_factory=lambda: ObjectiveConfig(
|
|
54
|
+
import_path(SpikingClassifierObjective), {"readout": "max", "rates": {"lower": 0.01, "upper": 0.3}}))
|
|
55
|
+
"""The objective and its keyword arguments (`--objective.readout mean`)."""
|
|
56
|
+
model: ModelConfig = dataclasses.field(default_factory=lambda: ModelConfig(import_path(SpikingMLP), {
|
|
57
|
+
"hidden": [128], "classes": 20, "dtype": "float32",
|
|
58
|
+
"neuron": {"class": import_path(ALIF), "fields": {
|
|
59
|
+
"tau": 5.0, "tau_adapt": 20.0, "beta": 0.2, "learn_tau": True, "detach_reset": True}}}))
|
|
60
|
+
data: SHD = dataclasses.field(default_factory=SHD)
|
|
61
|
+
optim: OptimConfig = dataclasses.field(default_factory=lambda: OptimConfig(learning_rate=2e-3,
|
|
62
|
+
clip_grads=1.0))
|
|
63
|
+
encoder: EncoderSpec = dataclasses.field(default_factory=EventsEncoder)
|
|
64
|
+
"""How a batch field becomes spikes; `encoder:rate --encoder.steps 8` for static data."""
|
|
65
|
+
sample: str = "spikes"
|
|
66
|
+
"""The batch field the encoder reads."""
|
|
67
|
+
smoke: bool = False
|
|
68
|
+
"""Train a small network for a few steps on synthetic SHD-layout recordings; nothing is downloaded."""
|
|
69
|
+
|
|
70
|
+
def smoked(self) -> SNNRunConfig:
|
|
71
|
+
"""This run shrunk to a few seconds on CPU, reading synthetic recordings.
|
|
72
|
+
|
|
73
|
+
The recordings go under the checkpoint directory, so the run reads
|
|
74
|
+
them there and leaves nothing elsewhere. The run it returns is no
|
|
75
|
+
longer a smoke run: it records the small run as it trains, so its
|
|
76
|
+
`run.json` trains that run again.
|
|
77
|
+
"""
|
|
78
|
+
cache = write_synthetic_shd(Path(self.trainer.checkpoint_dir) / "synthetic-shd")
|
|
79
|
+
synthetic = SHD(steps=20, channels=70, cache=str(cache),
|
|
80
|
+
loading=Loading(workers=0, threads=1, read_buffer=1))
|
|
81
|
+
model = dataclasses.replace(self.model, fields={**self.model.fields, "hidden": [16], "classes": 2})
|
|
82
|
+
trainer = dataclasses.replace(self.trainer, batch_size=16, steps=8, log_every=4, eval_every=8,
|
|
83
|
+
checkpoint_every=8)
|
|
84
|
+
return dataclasses.replace(self, data=synthetic, model=model, trainer=trainer, smoke=False)
|
|
85
|
+
|
|
86
|
+
def sample_field(self) -> Field:
|
|
87
|
+
"""The batch field the encoder reads, at the per-record shape the dataset writes."""
|
|
88
|
+
return Field(self.sample, (self.data.steps, self.data.channels))
|
|
89
|
+
|
|
90
|
+
def prepare(self) -> Prepared:
|
|
91
|
+
"""The objective the run names, around its model, the sample field and the encoder, on SHD.
|
|
92
|
+
|
|
93
|
+
Its schedules run over the run's steps unless the run states
|
|
94
|
+
`schedule_steps`.
|
|
95
|
+
"""
|
|
96
|
+
run = self.smoked() if self.smoke else self
|
|
97
|
+
dataset = run.data.load(batch=run.trainer.batch_size)
|
|
98
|
+
derived = ({} if "schedule_steps" in run.objective.fields
|
|
99
|
+
else {"schedule_steps": run.trainer.total_steps(dataset)})
|
|
100
|
+
objective = run.objective.build(model=run.model.build(), sample=run.sample_field(),
|
|
101
|
+
encoder=run.encoder, **derived)
|
|
102
|
+
return Prepared(run, lambda name: run.train(objective, dataset, name=name, metrics=[Accuracy()]))
|
sparx/datasets.py
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
"""Neuromorphic datasets as dense, binned spike counts.
|
|
2
|
+
|
|
3
|
+
`shd` reads the Spiking Heidelberg Digits (Cramer et al., "The Heidelberg
|
|
4
|
+
Spiking Data Sets for the Systematic Evaluation of Spiking Neural Networks",
|
|
5
|
+
IEEE TNNLS 2020): spoken digits 0-9 in English and German, 20 classes,
|
|
6
|
+
rendered as spikes on 700 cochlear channels. Each record becomes a
|
|
7
|
+
`[steps, channels]` array of spike counts, batch-major as dew's loaders
|
|
8
|
+
expect; `sparx.encode.EventsEncoder()` moves the time axis to the front. `SHD` is the
|
|
9
|
+
same data as a dew dataset spec, which a run's `data` holds (`--data.channels
|
|
10
|
+
140` on the recipe's command line). `write_synthetic_shd` writes small files
|
|
11
|
+
in SHD's layout, which smoke runs and tests read in its place.
|
|
12
|
+
|
|
13
|
+
`mnist` reads MNIST (LeCun et al. 1998) or Fashion-MNIST (Xiao et al. 2017)
|
|
14
|
+
as uint8 images and labels.
|
|
15
|
+
|
|
16
|
+
`holdout` splits records into a part to train on and a part to validate on.
|
|
17
|
+
|
|
18
|
+
Reading the files needs h5py (`pip install "sparxml[datasets]"`).
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
from __future__ import annotations
|
|
22
|
+
|
|
23
|
+
import gzip
|
|
24
|
+
import shutil
|
|
25
|
+
import urllib.request
|
|
26
|
+
from collections.abc import Mapping
|
|
27
|
+
from dataclasses import dataclass
|
|
28
|
+
from pathlib import Path
|
|
29
|
+
from typing import Literal
|
|
30
|
+
|
|
31
|
+
import numpy as np
|
|
32
|
+
from dew.data import Dataset, DatasetSpec
|
|
33
|
+
from dew.data.dataset import Tokenize
|
|
34
|
+
from dew.files import replacing
|
|
35
|
+
from numpy.typing import ArrayLike
|
|
36
|
+
|
|
37
|
+
__all__ = ["MNIST_URLS", "SHD", "SHD_URL", "Binning", "bin_events", "holdout", "mnist", "shd",
|
|
38
|
+
"write_synthetic_shd"]
|
|
39
|
+
|
|
40
|
+
SHD_URL = "https://zenkelab.org/datasets/shd/{split}.h5.gz"
|
|
41
|
+
MNIST_URLS = {
|
|
42
|
+
"mnist": "https://storage.googleapis.com/cvdf-datasets/mnist/{name}.gz",
|
|
43
|
+
"fashion": "https://raw.githubusercontent.com/zalandoresearch/fashion-mnist/master/data/fashion/{name}.gz",
|
|
44
|
+
}
|
|
45
|
+
"""Where `mnist` fetches each IDX file: the MNIST mirror and Zalando's Fashion-MNIST repository."""
|
|
46
|
+
_CHANNELS = 700
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
type Binning = Literal["grid", "events"]
|
|
50
|
+
"""How `bin_events` cuts time into steps.
|
|
51
|
+
|
|
52
|
+
- `grid`: equal bins from time 0, the usual binning.
|
|
53
|
+
- `events`: SpikingJelly's SHD frames by duration, in the releases SNN-delays
|
|
54
|
+
(Hammouamri et al., ICLR 2024) trained on: each step opens at the first
|
|
55
|
+
event not yet counted and holds every event within one step's duration of
|
|
56
|
+
it. Silences longer than a step are dropped, so a recording is shorter
|
|
57
|
+
than on the grid and its timing is compressed; times are scaled to
|
|
58
|
+
milliseconds in the file's own float16, as theirs are.
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def bin_events(times: ArrayLike, units: ArrayLike, steps: int, max_time: float, channels: int,
|
|
63
|
+
source_channels: int = _CHANNELS, binning: Binning = "grid") -> np.ndarray:
|
|
64
|
+
"""Count one record's spikes into `[steps, channels]` uint8 bins.
|
|
65
|
+
|
|
66
|
+
A step lasts `max_time / steps`. On the `grid`, time `[0, max_time)` is
|
|
67
|
+
cut into `steps` equal bins and events at or after `max_time` are
|
|
68
|
+
dropped; with `binning="events"` steps open at events (`Binning`) and
|
|
69
|
+
steps past the `steps`-th are dropped. `channels` must divide
|
|
70
|
+
`source_channels`; adjacent source channels are pooled into each output
|
|
71
|
+
channel. Counts saturate at 255.
|
|
72
|
+
"""
|
|
73
|
+
if source_channels % channels:
|
|
74
|
+
raise ValueError(f"channels must divide {source_channels}, not {channels}")
|
|
75
|
+
if binning == "grid":
|
|
76
|
+
step = np.floor(np.asarray(times, np.float64) / max_time * steps).astype(np.int64)
|
|
77
|
+
elif binning == "events":
|
|
78
|
+
step = _event_frames(np.asarray(times), round(1000 * max_time / steps, 9))
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError(f"binning must be grid or events, not {binning!r}")
|
|
81
|
+
kept = step < steps
|
|
82
|
+
counts = np.zeros((steps, channels), np.int64)
|
|
83
|
+
np.add.at(counts, (step[kept], np.asarray(units, np.int64)[kept] // (source_channels // channels)), 1)
|
|
84
|
+
return np.minimum(counts, 255).astype(np.uint8)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def _event_frames(times: np.ndarray, duration: float) -> np.ndarray:
|
|
88
|
+
"""Each event's frame under `Binning`'s `events`, `duration` in ms.
|
|
89
|
+
|
|
90
|
+
SpikingJelly's `integrate_events_by_fixed_duration_shd` scans the events
|
|
91
|
+
once: a frame starts at event `l` and takes every following event `r`
|
|
92
|
+
while `t[r] - t[l] <= duration`. Its arithmetic is kept, `1000 * t` in
|
|
93
|
+
the times' dtype, so the frames are theirs exactly.
|
|
94
|
+
"""
|
|
95
|
+
t = 1000 * times
|
|
96
|
+
if t.size == 0:
|
|
97
|
+
return np.zeros(0, np.int64)
|
|
98
|
+
starts = [0]
|
|
99
|
+
while True:
|
|
100
|
+
start = starts[-1]
|
|
101
|
+
beyond = t[start:] - t[start] > duration
|
|
102
|
+
if not beyond.any():
|
|
103
|
+
break
|
|
104
|
+
starts.append(start + int(np.argmax(beyond)))
|
|
105
|
+
return np.searchsorted(np.asarray(starts), np.arange(t.size), side="right") - 1
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _fetched(url: str, path: Path) -> Path:
|
|
109
|
+
"""`path`, downloaded from `url` the first time and published whole (`dew.files.replacing`), so an
|
|
110
|
+
interrupted download leaves nothing a later call would read."""
|
|
111
|
+
if not path.exists():
|
|
112
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
113
|
+
with replacing(path) as partial:
|
|
114
|
+
urllib.request.urlretrieve(url, str(partial))
|
|
115
|
+
return path
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def _download(split: str, cache: Path) -> Path:
|
|
119
|
+
"""The decompressed `{split}.h5` in `cache`, fetched and decompressed once."""
|
|
120
|
+
path = cache / f"{split}.h5"
|
|
121
|
+
if not path.exists():
|
|
122
|
+
archive = _fetched(SHD_URL.format(split=split), cache / f"{split}.h5.gz")
|
|
123
|
+
with gzip.open(archive) as source, replacing(path) as partial, open(partial, "wb") as target:
|
|
124
|
+
shutil.copyfileobj(source, target)
|
|
125
|
+
return path
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def mnist(split: Literal["train", "test"], *, fashion: bool = False,
|
|
129
|
+
cache: str | Path | None = None) -> dict[str, np.ndarray]:
|
|
130
|
+
"""MNIST `split`, or Fashion-MNIST's with `fashion=True`: uint8 images `[N, 28, 28]` under `"image"`
|
|
131
|
+
and int32 labels `[N]` under `"label"`.
|
|
132
|
+
|
|
133
|
+
The gzipped IDX files download once to `cache`, `~/.cache/sparx` by
|
|
134
|
+
default, Fashion-MNIST's under `fashion/`.
|
|
135
|
+
"""
|
|
136
|
+
root = Path(cache) if cache is not None else Path.home() / ".cache" / "sparx"
|
|
137
|
+
root = root / "fashion" if fashion else root
|
|
138
|
+
prefix = "train" if split == "train" else "t10k"
|
|
139
|
+
arrays = []
|
|
140
|
+
for name, header in ((f"{prefix}-images-idx3-ubyte", 16), (f"{prefix}-labels-idx1-ubyte", 8)):
|
|
141
|
+
path = _fetched(MNIST_URLS["fashion" if fashion else "mnist"].format(name=name), root / f"{name}.gz")
|
|
142
|
+
with gzip.open(path) as file:
|
|
143
|
+
arrays.append(np.frombuffer(file.read(), np.uint8, offset=header))
|
|
144
|
+
images, labels = arrays
|
|
145
|
+
return {"image": images.reshape(-1, 28, 28), "label": labels.astype(np.int32)}
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def shd(split: Literal["train", "test"], steps: int = 100, max_time: float = 1.4, channels: int = 700,
|
|
149
|
+
cache: str | Path | None = None, path: str | Path | None = None,
|
|
150
|
+
binning: Binning = "grid") -> dict[str, np.ndarray]:
|
|
151
|
+
"""The SHD `split` as `{"spikes": uint8 [N, steps, channels], "label": int32 [N]}`.
|
|
152
|
+
|
|
153
|
+
`steps` bins over the first `max_time` seconds. Every spike of the train
|
|
154
|
+
split falls before 1.37 s, and 100 steps of 14 ms is the binning of
|
|
155
|
+
Zenke's SpyTorch SHD tutorial. `binning="events"` is SNN-delays' binning
|
|
156
|
+
(`Binning`); at 10 ms its recordings last at most 124 steps (train) and
|
|
157
|
+
105 (test), so `steps=124, max_time=1.24` keeps every event.
|
|
158
|
+
The file is read from `path` when given, otherwise downloaded once into
|
|
159
|
+
`cache` (default `~/.cache/sparx`; 131 MB for train, 38 MB for test) and
|
|
160
|
+
decompressed beside it.
|
|
161
|
+
"""
|
|
162
|
+
import h5py
|
|
163
|
+
|
|
164
|
+
if path is None:
|
|
165
|
+
path = _download(f"shd_{split}", Path.home() / ".cache" / "sparx" if cache is None else Path(cache))
|
|
166
|
+
with h5py.File(path, "r") as file:
|
|
167
|
+
nodes = [file.get(name) for name in ("spikes/times", "spikes/units", "labels")]
|
|
168
|
+
times, units, labels = nodes
|
|
169
|
+
if not (isinstance(times, h5py.Dataset) and isinstance(units, h5py.Dataset)
|
|
170
|
+
and isinstance(labels, h5py.Dataset)):
|
|
171
|
+
raise ValueError(f"{path} lacks SHD's spikes/times, spikes/units and labels datasets")
|
|
172
|
+
spikes = np.stack([bin_events(times[i], units[i], steps, max_time, channels, binning=binning)
|
|
173
|
+
for i in range(len(times))])
|
|
174
|
+
return {"spikes": spikes, "label": np.asarray(labels, np.int32)}
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
@dataclass(frozen=True)
|
|
178
|
+
class SHD(DatasetSpec):
|
|
179
|
+
"""The Spiking Heidelberg Digits as a dew dataset: the train split to train on, the test
|
|
180
|
+
split to validate on (SHD has no separate validation split), binned by `shd`.
|
|
181
|
+
|
|
182
|
+
Records are `{"spikes": uint8 [steps, channels], "label": int32}`; `sparx.encode.EventsEncoder()`
|
|
183
|
+
turns a batch into the network's time-major input. Both splits are held in memory
|
|
184
|
+
(about 700 MB at 100 steps over 700 channels), shuffled from `seed` every epoch, and
|
|
185
|
+
each process reads its share of every batch.
|
|
186
|
+
"""
|
|
187
|
+
|
|
188
|
+
steps: int = 100
|
|
189
|
+
max_time: float = 1.4
|
|
190
|
+
channels: int = 700
|
|
191
|
+
cache: str | None = None
|
|
192
|
+
binning: Binning = "grid"
|
|
193
|
+
|
|
194
|
+
def load(self, *, batch: int, tokenize: Tokenize | None = None) -> Dataset:
|
|
195
|
+
self.uncaptioned(tokenize)
|
|
196
|
+
train = shd("train", self.steps, self.max_time, self.channels, self.cache, binning=self.binning)
|
|
197
|
+
test = shd("test", self.steps, self.max_time, self.channels, self.cache, binning=self.binning)
|
|
198
|
+
return Dataset.from_records(train, batch=batch, seed=self.seed, validation=test, loading=self.loading)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def write_synthetic_shd(directory: str | Path, records: int = 64, seed: int = 0) -> Path:
|
|
202
|
+
"""Write `shd_train.h5` and `shd_test.h5` in SHD's layout into `directory`, and return it.
|
|
203
|
+
|
|
204
|
+
Each split holds `records` recordings of 60 spikes over 1.4 s, labelled 0
|
|
205
|
+
or 1 by which half of the 700 channels fires, so a network learns them in
|
|
206
|
+
a few steps. `shd(split, cache=directory)` and `SHD(cache=directory)` read
|
|
207
|
+
them as they read SHD's own files, which lets a smoke run or a test
|
|
208
|
+
exercise the whole path without the 169 MB download.
|
|
209
|
+
"""
|
|
210
|
+
import h5py
|
|
211
|
+
|
|
212
|
+
directory = Path(directory)
|
|
213
|
+
directory.mkdir(parents=True, exist_ok=True)
|
|
214
|
+
rng = np.random.default_rng(seed)
|
|
215
|
+
half = _CHANNELS // 2
|
|
216
|
+
for split in ("train", "test"):
|
|
217
|
+
labels = rng.integers(0, 2, records).astype(np.uint16)
|
|
218
|
+
with h5py.File(directory / f"shd_{split}.h5", "w") as file:
|
|
219
|
+
times = file.create_dataset("spikes/times", (records,), dtype=h5py.vlen_dtype(np.float32))
|
|
220
|
+
units = file.create_dataset("spikes/units", (records,), dtype=h5py.vlen_dtype(np.uint16))
|
|
221
|
+
for i, label in enumerate(labels):
|
|
222
|
+
times[i] = np.sort(rng.uniform(0, 1.4, 60)).astype(np.float32)
|
|
223
|
+
units[i] = (rng.integers(0, half, 60) + half * int(label)).astype(np.uint16)
|
|
224
|
+
file.create_dataset("labels", data=labels)
|
|
225
|
+
return directory
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def holdout(records: Mapping[str, np.ndarray], fraction: float,
|
|
229
|
+
seed: int = 0) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]:
|
|
230
|
+
"""Split `records` (columns of equal length) into a random `1 - fraction` and `fraction`.
|
|
231
|
+
|
|
232
|
+
SHD has no validation split, and SNN-delays selects its epochs on the
|
|
233
|
+
test set; holding out part of the training set gives a validation set
|
|
234
|
+
to select on, so the test accuracy stays an estimate.
|
|
235
|
+
"""
|
|
236
|
+
sizes = {len(column) for column in records.values()}
|
|
237
|
+
if len(sizes) != 1:
|
|
238
|
+
raise ValueError(f"columns of one length split together, not lengths {sorted(sizes)}")
|
|
239
|
+
total = sizes.pop()
|
|
240
|
+
held = round(fraction * total)
|
|
241
|
+
if not 0 < held < total:
|
|
242
|
+
raise ValueError(f"holding out {fraction} of {total} records leaves one side empty")
|
|
243
|
+
order = np.random.default_rng(seed).permutation(total)
|
|
244
|
+
kept, out = np.sort(order[held:]), np.sort(order[:held])
|
|
245
|
+
return ({name: column[kept] for name, column in records.items()},
|
|
246
|
+
{name: column[out] for name, column in records.items()})
|
|
@@ -0,0 +1,178 @@
|
|
|
1
|
+
"""Neuron, synapse and plasticity models, and the runner that scans them over time.
|
|
2
|
+
|
|
3
|
+
Every neuron model meets one contract (`sparx.dynamics.core`):
|
|
4
|
+
`init_state(shape, dtype)` and `step(state, SynapticInput, dt) -> (state,
|
|
5
|
+
Output)`, and `run(model, inputs)` scans one over time. A model's output is
|
|
6
|
+
its spikes, or a graded value each step when the model is `graded`. Two families meet
|
|
7
|
+
it. `sparx.dynamics.ml` holds the dimensionless, one-step-is-one-unit
|
|
8
|
+
models deep spiking networks train with: soft resets, detached resets,
|
|
9
|
+
learnable decays, input as a jump of the membrane. `sparx.dynamics.neurons`
|
|
10
|
+
holds models of neurons as biology measures them: membrane equations in mV
|
|
11
|
+
and ms with conductances, refractoriness and reversal potentials. Synapses
|
|
12
|
+
with receptor kinetics and plasticity rules complete them; `sparx.nn` builds
|
|
13
|
+
layers from these models and `sparx.graph` builds circuits and connectomes
|
|
14
|
+
(design.md sections 4 and 5).
|
|
15
|
+
|
|
16
|
+
This package exports the models and the contract. The arithmetic the
|
|
17
|
+
models share (`fire`, `exact_linear`, `rk4`, `substeps` and the rest), which
|
|
18
|
+
a new model is written with, stays in `sparx.dynamics.core`.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
from sparx.dynamics.core import Gap, Model, NeuronModel, Output, Reset, SynapticInput, Term, decay, run
|
|
22
|
+
from sparx.dynamics.homeostasis import IntrinsicPlasticity, IntrinsicState
|
|
23
|
+
from sparx.dynamics.ml import (
|
|
24
|
+
ACTIVATIONS,
|
|
25
|
+
ALIFCell,
|
|
26
|
+
ALIFState,
|
|
27
|
+
BernoulliCell,
|
|
28
|
+
BernoulliState,
|
|
29
|
+
DecayingHebb,
|
|
30
|
+
Dense,
|
|
31
|
+
EligibleHebb,
|
|
32
|
+
FastWeights,
|
|
33
|
+
HebbianRule,
|
|
34
|
+
LICell,
|
|
35
|
+
LIFCell,
|
|
36
|
+
MembraneState,
|
|
37
|
+
ModulatedHebb,
|
|
38
|
+
OjaHebb,
|
|
39
|
+
PulseCell,
|
|
40
|
+
RateCell,
|
|
41
|
+
RateState,
|
|
42
|
+
RecurrentCell,
|
|
43
|
+
RecurrentState,
|
|
44
|
+
RetroactiveHebb,
|
|
45
|
+
Serial,
|
|
46
|
+
Sparse,
|
|
47
|
+
Wiring,
|
|
48
|
+
)
|
|
49
|
+
from sparx.dynamics.neurons import (
|
|
50
|
+
IZHIKEVICH_2003,
|
|
51
|
+
IZHIKEVICH_2004,
|
|
52
|
+
RECEPTORS,
|
|
53
|
+
AdEx,
|
|
54
|
+
AdExState,
|
|
55
|
+
GradedPotential,
|
|
56
|
+
GradedPotentialState,
|
|
57
|
+
HodgkinHuxley,
|
|
58
|
+
HodgkinHuxleyState,
|
|
59
|
+
Izhikevich,
|
|
60
|
+
IzhikevichState,
|
|
61
|
+
LeakyIntegrateAndFire,
|
|
62
|
+
LeakyIntegrateAndFireState,
|
|
63
|
+
MgBlock,
|
|
64
|
+
izhikevich_2003,
|
|
65
|
+
izhikevich_2004,
|
|
66
|
+
)
|
|
67
|
+
from sparx.dynamics.plasticity import (
|
|
68
|
+
DopamineSTDP,
|
|
69
|
+
DopamineTraces,
|
|
70
|
+
PairSTDP,
|
|
71
|
+
Plasticity,
|
|
72
|
+
Rules,
|
|
73
|
+
ScalingTraces,
|
|
74
|
+
STDPTraces,
|
|
75
|
+
SynapticScaling,
|
|
76
|
+
TripletSTDP,
|
|
77
|
+
TripletTraces,
|
|
78
|
+
TsodyksMarkram,
|
|
79
|
+
TsodyksMarkramState,
|
|
80
|
+
)
|
|
81
|
+
from sparx.dynamics.synapses import (
|
|
82
|
+
Alpha,
|
|
83
|
+
AlphaState,
|
|
84
|
+
Arrivals,
|
|
85
|
+
BiExponential,
|
|
86
|
+
BiExponentialState,
|
|
87
|
+
Delta,
|
|
88
|
+
Exponential,
|
|
89
|
+
Graded,
|
|
90
|
+
GradedState,
|
|
91
|
+
Landing,
|
|
92
|
+
PointNeuron,
|
|
93
|
+
PointNeuronState,
|
|
94
|
+
Receptor,
|
|
95
|
+
StochasticRelease,
|
|
96
|
+
SynapseModel,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
__all__ = [
|
|
100
|
+
"ACTIVATIONS",
|
|
101
|
+
"IZHIKEVICH_2003",
|
|
102
|
+
"IZHIKEVICH_2004",
|
|
103
|
+
"RECEPTORS",
|
|
104
|
+
"ALIFCell",
|
|
105
|
+
"ALIFState",
|
|
106
|
+
"AdEx",
|
|
107
|
+
"AdExState",
|
|
108
|
+
"Alpha",
|
|
109
|
+
"AlphaState",
|
|
110
|
+
"Arrivals",
|
|
111
|
+
"BernoulliCell",
|
|
112
|
+
"BernoulliState",
|
|
113
|
+
"BiExponential",
|
|
114
|
+
"BiExponentialState",
|
|
115
|
+
"DecayingHebb",
|
|
116
|
+
"Delta",
|
|
117
|
+
"Dense",
|
|
118
|
+
"DopamineSTDP",
|
|
119
|
+
"DopamineTraces",
|
|
120
|
+
"EligibleHebb",
|
|
121
|
+
"Exponential",
|
|
122
|
+
"FastWeights",
|
|
123
|
+
"Gap",
|
|
124
|
+
"Graded",
|
|
125
|
+
"GradedPotential",
|
|
126
|
+
"GradedPotentialState",
|
|
127
|
+
"GradedState",
|
|
128
|
+
"HebbianRule",
|
|
129
|
+
"HodgkinHuxley",
|
|
130
|
+
"HodgkinHuxleyState",
|
|
131
|
+
"IntrinsicPlasticity",
|
|
132
|
+
"IntrinsicState",
|
|
133
|
+
"Izhikevich",
|
|
134
|
+
"IzhikevichState",
|
|
135
|
+
"LICell",
|
|
136
|
+
"LIFCell",
|
|
137
|
+
"Landing",
|
|
138
|
+
"LeakyIntegrateAndFire",
|
|
139
|
+
"LeakyIntegrateAndFireState",
|
|
140
|
+
"MembraneState",
|
|
141
|
+
"MgBlock",
|
|
142
|
+
"Model",
|
|
143
|
+
"ModulatedHebb",
|
|
144
|
+
"NeuronModel",
|
|
145
|
+
"OjaHebb",
|
|
146
|
+
"Output",
|
|
147
|
+
"PairSTDP",
|
|
148
|
+
"Plasticity",
|
|
149
|
+
"PointNeuron",
|
|
150
|
+
"PointNeuronState",
|
|
151
|
+
"PulseCell",
|
|
152
|
+
"RateCell",
|
|
153
|
+
"RateState",
|
|
154
|
+
"Receptor",
|
|
155
|
+
"RecurrentCell",
|
|
156
|
+
"RecurrentState",
|
|
157
|
+
"Reset",
|
|
158
|
+
"RetroactiveHebb",
|
|
159
|
+
"Rules",
|
|
160
|
+
"STDPTraces",
|
|
161
|
+
"ScalingTraces",
|
|
162
|
+
"Serial",
|
|
163
|
+
"Sparse",
|
|
164
|
+
"StochasticRelease",
|
|
165
|
+
"SynapseModel",
|
|
166
|
+
"SynapticInput",
|
|
167
|
+
"SynapticScaling",
|
|
168
|
+
"Term",
|
|
169
|
+
"TripletSTDP",
|
|
170
|
+
"TripletTraces",
|
|
171
|
+
"TsodyksMarkram",
|
|
172
|
+
"TsodyksMarkramState",
|
|
173
|
+
"Wiring",
|
|
174
|
+
"decay",
|
|
175
|
+
"izhikevich_2003",
|
|
176
|
+
"izhikevich_2004",
|
|
177
|
+
"run",
|
|
178
|
+
]
|