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