flybrainer 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.
- flybrainer/__init__.py +77 -0
- flybrainer/brain.py +523 -0
- flybrainer/connectome/__init__.py +59 -0
- flybrainer/connectome/compile.py +469 -0
- flybrainer/connectome/download.py +122 -0
- flybrainer/connectome/normalize.py +299 -0
- flybrainer/connectome/sources.py +61 -0
- flybrainer/connectome/verify.py +193 -0
- flybrainer/encoders.py +546 -0
- flybrainer/eyemap.py +133 -0
- flybrainer/interfaces.py +268 -0
- flybrainer/kernel.py +423 -0
- flybrainer/paths.py +53 -0
- flybrainer/plasticity.py +210 -0
- flybrainer/readout/__init__.py +5 -0
- flybrainer/readout/fixed.py +176 -0
- flybrainer/readout/reservoir.py +361 -0
- flybrainer-0.1.0.dist-info/METADATA +180 -0
- flybrainer-0.1.0.dist-info/RECORD +21 -0
- flybrainer-0.1.0.dist-info/WHEEL +4 -0
- flybrainer-0.1.0.dist-info/licenses/LICENSE +25 -0
flybrainer/__init__.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
"""flybrainer: a Numba-JIT accelerated fruit-fly connectome.
|
|
2
|
+
|
|
3
|
+
Extracted from OpenFly (github.com/marketcalls/openfly). Independent of
|
|
4
|
+
broker / market / front-end code.
|
|
5
|
+
"""
|
|
6
|
+
from flybrainer import readout
|
|
7
|
+
from flybrainer.brain import Brain, build_populations, load_graph, r8_ame12_edges
|
|
8
|
+
from flybrainer.encoders import (
|
|
9
|
+
ENCODER_NAMES,
|
|
10
|
+
FEATURE_NAMES,
|
|
11
|
+
BarsEncoder,
|
|
12
|
+
ChartEncoder,
|
|
13
|
+
FeatureEncoder,
|
|
14
|
+
make_encoder,
|
|
15
|
+
)
|
|
16
|
+
from flybrainer.eyemap import EyeMap, as_eye_map, default_eye_map, resolve_eye_map
|
|
17
|
+
from flybrainer.interfaces import (
|
|
18
|
+
REQUIRED_POPULATIONS,
|
|
19
|
+
BrainProtocol,
|
|
20
|
+
Decision,
|
|
21
|
+
EncoderProtocol,
|
|
22
|
+
ObservationResult,
|
|
23
|
+
Prediction,
|
|
24
|
+
ReadoutProtocol,
|
|
25
|
+
SensorFrame,
|
|
26
|
+
Stimulus,
|
|
27
|
+
)
|
|
28
|
+
from flybrainer.kernel import (
|
|
29
|
+
KERNEL_PARAMETERS,
|
|
30
|
+
KERNEL_VERSION,
|
|
31
|
+
Kernel,
|
|
32
|
+
KernelState,
|
|
33
|
+
)
|
|
34
|
+
from flybrainer.plasticity import MV_PER_CONTACT, KCMBONPlasticity, PlasticityConfig
|
|
35
|
+
|
|
36
|
+
__version__ = "0.1.0"
|
|
37
|
+
|
|
38
|
+
__all__ = [
|
|
39
|
+
# contracts
|
|
40
|
+
"Stimulus",
|
|
41
|
+
"ObservationResult",
|
|
42
|
+
"BrainProtocol",
|
|
43
|
+
"REQUIRED_POPULATIONS",
|
|
44
|
+
"SensorFrame",
|
|
45
|
+
"Decision",
|
|
46
|
+
"Prediction",
|
|
47
|
+
"EncoderProtocol",
|
|
48
|
+
"ReadoutProtocol",
|
|
49
|
+
# eyemap
|
|
50
|
+
"EyeMap",
|
|
51
|
+
"as_eye_map",
|
|
52
|
+
"default_eye_map",
|
|
53
|
+
"resolve_eye_map",
|
|
54
|
+
# encoders
|
|
55
|
+
"ChartEncoder",
|
|
56
|
+
"BarsEncoder",
|
|
57
|
+
"FeatureEncoder",
|
|
58
|
+
"make_encoder",
|
|
59
|
+
"FEATURE_NAMES",
|
|
60
|
+
"ENCODER_NAMES",
|
|
61
|
+
# kernel
|
|
62
|
+
"Kernel",
|
|
63
|
+
"KernelState",
|
|
64
|
+
"KERNEL_PARAMETERS",
|
|
65
|
+
"KERNEL_VERSION",
|
|
66
|
+
# plasticity
|
|
67
|
+
"KCMBONPlasticity",
|
|
68
|
+
"PlasticityConfig",
|
|
69
|
+
"MV_PER_CONTACT",
|
|
70
|
+
# brain
|
|
71
|
+
"Brain",
|
|
72
|
+
"load_graph",
|
|
73
|
+
"build_populations",
|
|
74
|
+
"r8_ame12_edges",
|
|
75
|
+
# subpackage
|
|
76
|
+
"readout",
|
|
77
|
+
]
|
flybrainer/brain.py
ADDED
|
@@ -0,0 +1,523 @@
|
|
|
1
|
+
"""Brain: the BrainProtocol implementation over the compiled graph.
|
|
2
|
+
|
|
3
|
+
Populations (name -> int32 index array in graph order):
|
|
4
|
+
|
|
5
|
+
R1-R6 R1-R6 photoreceptors mapped to an ommatidial column (stimulus order)
|
|
6
|
+
R8 all mapped R8p and R8y cells (stimulus order for `Stimulus.r8`)
|
|
7
|
+
R8p, R8y subsets of R8 by channel (2 = R8p blue, 1 = R8y green)
|
|
8
|
+
lamina L1, L2, L3 and L5 cells
|
|
9
|
+
KC type starts with KC
|
|
10
|
+
PAM11, PPL101 dopamine neurons of the two compartments used by the plastic arm
|
|
11
|
+
MBON07, MBON11 the corresponding output neurons; MBON = every type starting with MBON
|
|
12
|
+
DNp20_L, DNp20_R type DNp20 by somaSide; DNpe017 both sides
|
|
13
|
+
DN superclass == descending_neuron (the tbc, sensory_descending and
|
|
14
|
+
efferent_descending variants are excluded)
|
|
15
|
+
central_complex class == CX (the MaleCNS `class` column)
|
|
16
|
+
random2000 2,000 neurons sampled without replacement (numpy default_rng, seed
|
|
17
|
+
20260912) from superclass cb_* excluding photoreceptors and lamina
|
|
18
|
+
|
|
19
|
+
Stimulus: `r16` and `r8` are luminance in [0, 1] per mapped R1-R6 and R8
|
|
20
|
+
cell, low-passed with tau 10 ms (updated per 10 ms bin). Drive: lamina 12 mV
|
|
21
|
+
constant, photoreceptors 30 x L / (h + L) mV with half-saturation h
|
|
22
|
+
(constructor parameter, settings key neural.half_saturation). Pulses add a
|
|
23
|
+
constant to a population's drive for the first `duration_ms` of the
|
|
24
|
+
observation. The simulation runs in 10 ms bins so pulses can end mid
|
|
25
|
+
observation; the clock continues across observations.
|
|
26
|
+
|
|
27
|
+
Declared assumption (Brain(r8_ame12_excitatory=True), settings key
|
|
28
|
+
neural.r8_ame12_excitatory, default on): the 390 edges from R8
|
|
29
|
+
photoreceptors onto the six aMe12 cells are made excitatory on the brain's
|
|
30
|
+
own copy of the weights, because R8 drives aMe12 in vivo
|
|
31
|
+
(https://doi.org/10.1038/s41586-023-06681-6) while the histamine sign rule
|
|
32
|
+
would silence it. Without it no visual signal reaches the Kenyon cells.
|
|
33
|
+
Recorded in parameters() and provenance() with the edge count and hashes.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
from __future__ import annotations
|
|
37
|
+
|
|
38
|
+
import json
|
|
39
|
+
import math
|
|
40
|
+
import time
|
|
41
|
+
from dataclasses import dataclass
|
|
42
|
+
from pathlib import Path
|
|
43
|
+
from typing import Any
|
|
44
|
+
|
|
45
|
+
import numpy as np
|
|
46
|
+
|
|
47
|
+
from flybrainer import kernel as K
|
|
48
|
+
from flybrainer.connectome.compile import GRAPH_FORMAT_VERSION
|
|
49
|
+
from flybrainer.connectome.verify import sha256_array
|
|
50
|
+
from flybrainer.interfaces import REQUIRED_POPULATIONS, ObservationResult, Stimulus
|
|
51
|
+
from flybrainer.kernel import KERNEL_PARAMETERS, KERNEL_VERSION, Kernel, KernelState
|
|
52
|
+
from flybrainer.paths import PATHS
|
|
53
|
+
from flybrainer.plasticity import KCMBONPlasticity, PlasticityConfig
|
|
54
|
+
|
|
55
|
+
BIN_STEPS = 100
|
|
56
|
+
BIN_MS = BIN_STEPS * K.DT_MS
|
|
57
|
+
LUMINANCE_TAU_MS = 10.0
|
|
58
|
+
LAMINA_DRIVE_MV = 12.0
|
|
59
|
+
PHOTORECEPTOR_MAX_MV = 30.0
|
|
60
|
+
DEFAULT_HALF_SATURATION = 0.5
|
|
61
|
+
RANDOM_SAMPLE_SEED = 20260912
|
|
62
|
+
RANDOM_SAMPLE_SIZE = 2000
|
|
63
|
+
DESCENDING_SUPERCLASS = "descending_neuron"
|
|
64
|
+
CENTRAL_COMPLEX_CLASS = "CX"
|
|
65
|
+
CENTRAL_BRAIN_PREFIX = "cb_"
|
|
66
|
+
LAMINA_ALL_TYPES = ("L1", "L2", "L3", "L4", "L5")
|
|
67
|
+
PHOTORECEPTOR_PREFIXES = ("R1-R6", "R7", "R8")
|
|
68
|
+
CHECKPOINT_VERSION = 1
|
|
69
|
+
|
|
70
|
+
# Declared modeling assumption: R8 photoreceptors drive the aMe12 accessory medulla
|
|
71
|
+
# neurons (Nature 2023, https://doi.org/10.1038/s41586-023-06681-6). The transmitter
|
|
72
|
+
# sign rule makes every R8 output inhibitory (histamine), which silences aMe12 and,
|
|
73
|
+
# through its 191 synapses onto Kenyon cells, the whole downstream brain. With the
|
|
74
|
+
# flag on, the R8 to aMe12 edges are made excitatory on the brain's own copy of the
|
|
75
|
+
# weights; the compiled graph and its hashes stay pristine.
|
|
76
|
+
R8_AME12_SOURCE_PREFIX = "R8"
|
|
77
|
+
R8_AME12_TARGET_TYPE = "aMe12"
|
|
78
|
+
R8_AME12_CITATION = "https://doi.org/10.1038/s41586-023-06681-6"
|
|
79
|
+
|
|
80
|
+
EXPECTED_POPULATION_SIZES = {
|
|
81
|
+
"PAM11": 15,
|
|
82
|
+
"PPL101": 2,
|
|
83
|
+
"MBON07": 4,
|
|
84
|
+
"MBON11": 2,
|
|
85
|
+
"R1-R6": 3335,
|
|
86
|
+
"R8": 811,
|
|
87
|
+
}
|
|
88
|
+
|
|
89
|
+
GRAPH_ARRAYS = (
|
|
90
|
+
"ptr",
|
|
91
|
+
"post",
|
|
92
|
+
"weight",
|
|
93
|
+
"ids",
|
|
94
|
+
"superclass",
|
|
95
|
+
"type",
|
|
96
|
+
"cell_class",
|
|
97
|
+
"soma_side",
|
|
98
|
+
"root_side",
|
|
99
|
+
"nt_sign",
|
|
100
|
+
"nt_uncertain",
|
|
101
|
+
"modulatory",
|
|
102
|
+
"r16",
|
|
103
|
+
"r16_uv",
|
|
104
|
+
"r16_eye",
|
|
105
|
+
"r8",
|
|
106
|
+
"r8_uv",
|
|
107
|
+
"r8_eye",
|
|
108
|
+
"r8_channel",
|
|
109
|
+
"lamina",
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
@dataclass(frozen=True)
|
|
114
|
+
class EyeMap:
|
|
115
|
+
"""Photoreceptor geometry for the sensory encoders.
|
|
116
|
+
|
|
117
|
+
`uv_r16[i]` is the (u, v) position in [0, 1] x [0, 1] of the cell
|
|
118
|
+
`r16[i]`, which is also `populations["R1-R6"][i]` and receives
|
|
119
|
+
`Stimulus.r16[i]`. Likewise for R8 with `Stimulus.r8`. `eye` is 0 for
|
|
120
|
+
the left eye and 1 for the right; both eyes cover the full field.
|
|
121
|
+
`r8_channel` is 1 for R8y (green) and 2 for R8p (blue).
|
|
122
|
+
"""
|
|
123
|
+
|
|
124
|
+
uv_r16: np.ndarray
|
|
125
|
+
eye_r16: np.ndarray
|
|
126
|
+
uv_r8: np.ndarray
|
|
127
|
+
eye_r8: np.ndarray
|
|
128
|
+
r8_channel: np.ndarray
|
|
129
|
+
r16: np.ndarray
|
|
130
|
+
r8: np.ndarray
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def load_graph(path: Path | str | None = None) -> dict[str, np.ndarray]:
|
|
134
|
+
path = Path(path) if path else PATHS.graph
|
|
135
|
+
if not path.exists():
|
|
136
|
+
raise FileNotFoundError(f"compiled graph not found at {path} (run: openfly prepare)")
|
|
137
|
+
with np.load(path, allow_pickle=False) as z:
|
|
138
|
+
missing = [k for k in GRAPH_ARRAYS if k not in z.files]
|
|
139
|
+
if missing:
|
|
140
|
+
raise ValueError(f"{path} is missing arrays: {', '.join(missing)}")
|
|
141
|
+
return {k: z[k] for k in z.files}
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
def _startswith(values: np.ndarray, prefix: str) -> np.ndarray:
|
|
145
|
+
return np.char.startswith(values.astype(str), prefix)
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def build_populations(g: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
|
149
|
+
types = g["type"].astype(str)
|
|
150
|
+
superclass = g["superclass"].astype(str)
|
|
151
|
+
soma = g["soma_side"].astype(str)
|
|
152
|
+
cell_class = g["cell_class"].astype(str)
|
|
153
|
+
|
|
154
|
+
def idx(mask: np.ndarray) -> np.ndarray:
|
|
155
|
+
return np.flatnonzero(mask).astype(np.int32)
|
|
156
|
+
|
|
157
|
+
r8 = np.asarray(g["r8"], dtype=np.int32)
|
|
158
|
+
channel = np.asarray(g["r8_channel"], dtype=np.int8)
|
|
159
|
+
pops: dict[str, np.ndarray] = {
|
|
160
|
+
"R1-R6": np.asarray(g["r16"], dtype=np.int32),
|
|
161
|
+
"R8": r8,
|
|
162
|
+
"R8p": r8[channel == 2],
|
|
163
|
+
"R8y": r8[channel == 1],
|
|
164
|
+
"lamina": np.asarray(g["lamina"], dtype=np.int32),
|
|
165
|
+
"KC": idx(_startswith(types, "KC")),
|
|
166
|
+
"PAM11": idx(types == "PAM11"),
|
|
167
|
+
"PPL101": idx(types == "PPL101"),
|
|
168
|
+
"MBON07": idx(types == "MBON07"),
|
|
169
|
+
"MBON11": idx(types == "MBON11"),
|
|
170
|
+
"MBON": idx(_startswith(types, "MBON")),
|
|
171
|
+
"DNp20_L": idx((types == "DNp20") & (soma == "L")),
|
|
172
|
+
"DNp20_R": idx((types == "DNp20") & (soma == "R")),
|
|
173
|
+
"DNpe017": idx(types == "DNpe017"),
|
|
174
|
+
"DN": idx(superclass == DESCENDING_SUPERCLASS),
|
|
175
|
+
"central_complex": idx(cell_class == CENTRAL_COMPLEX_CLASS),
|
|
176
|
+
}
|
|
177
|
+
photoreceptor = np.zeros(len(types), dtype=bool)
|
|
178
|
+
for prefix in PHOTORECEPTOR_PREFIXES:
|
|
179
|
+
photoreceptor |= _startswith(types, prefix)
|
|
180
|
+
lamina_any = np.isin(types, LAMINA_ALL_TYPES)
|
|
181
|
+
pool = idx(_startswith(superclass, CENTRAL_BRAIN_PREFIX) & ~photoreceptor & ~lamina_any)
|
|
182
|
+
rng = np.random.default_rng(RANDOM_SAMPLE_SEED)
|
|
183
|
+
size = min(RANDOM_SAMPLE_SIZE, len(pool))
|
|
184
|
+
sample = rng.choice(pool, size=size, replace=False) if size else np.zeros(0, np.int32)
|
|
185
|
+
pops["random2000"] = np.sort(sample).astype(np.int32)
|
|
186
|
+
return pops
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def r8_ame12_edges(g: dict[str, np.ndarray]) -> np.ndarray:
|
|
190
|
+
"""Edge indices from any R8 photoreceptor type onto aMe12 cells (sorted)."""
|
|
191
|
+
types = g["type"].astype(str)
|
|
192
|
+
ptr = g["ptr"]
|
|
193
|
+
post = g["post"]
|
|
194
|
+
sources = np.flatnonzero(np.char.startswith(types, R8_AME12_SOURCE_PREFIX))
|
|
195
|
+
is_target = types == R8_AME12_TARGET_TYPE
|
|
196
|
+
out = []
|
|
197
|
+
for i in sources:
|
|
198
|
+
e = np.arange(ptr[i], ptr[i + 1], dtype=np.int64)
|
|
199
|
+
if len(e):
|
|
200
|
+
sel = is_target[post[e]]
|
|
201
|
+
if sel.any():
|
|
202
|
+
out.append(e[sel])
|
|
203
|
+
return np.concatenate(out) if out else np.zeros(0, dtype=np.int64)
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def check_population_sizes(
|
|
207
|
+
pops: dict[str, np.ndarray], expected: dict[str, int] | None = None
|
|
208
|
+
) -> None:
|
|
209
|
+
expected = expected or EXPECTED_POPULATION_SIZES
|
|
210
|
+
problems = [
|
|
211
|
+
f"{name} has {len(pops[name])} cells, expected {size}"
|
|
212
|
+
for name, size in expected.items()
|
|
213
|
+
if len(pops.get(name, ())) != size
|
|
214
|
+
]
|
|
215
|
+
missing = [name for name in REQUIRED_POPULATIONS if name not in pops]
|
|
216
|
+
if missing:
|
|
217
|
+
problems.append("missing populations: " + ", ".join(missing))
|
|
218
|
+
if problems:
|
|
219
|
+
raise ValueError("population check failed: " + "; ".join(problems))
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
class Brain:
|
|
223
|
+
"""BrainProtocol implementation. See the module docstring for conventions."""
|
|
224
|
+
|
|
225
|
+
def __init__(
|
|
226
|
+
self,
|
|
227
|
+
graph_path: Path | str | None = None,
|
|
228
|
+
half_saturation: float = DEFAULT_HALF_SATURATION,
|
|
229
|
+
plastic: bool = False,
|
|
230
|
+
plasticity_config: PlasticityConfig | None = None,
|
|
231
|
+
check_counts: bool = True,
|
|
232
|
+
graph: dict[str, np.ndarray] | None = None,
|
|
233
|
+
r8_ame12_excitatory: bool = True,
|
|
234
|
+
):
|
|
235
|
+
if half_saturation <= 0:
|
|
236
|
+
raise ValueError("half_saturation must be positive")
|
|
237
|
+
self.r8_ame12_excitatory = bool(r8_ame12_excitatory)
|
|
238
|
+
self.graph_path = (
|
|
239
|
+
str(Path(graph_path) if graph_path else PATHS.graph) if graph is None else "<in-memory>"
|
|
240
|
+
)
|
|
241
|
+
g = graph if graph is not None else load_graph(graph_path)
|
|
242
|
+
self._graph = g
|
|
243
|
+
self.n = int(len(g["ptr"]) - 1)
|
|
244
|
+
self.edges = int(len(g["post"]))
|
|
245
|
+
self.half_saturation = float(half_saturation)
|
|
246
|
+
self.plastic = bool(plastic)
|
|
247
|
+
self.types = g["type"].astype(str)
|
|
248
|
+
self.populations = build_populations(g)
|
|
249
|
+
if check_counts:
|
|
250
|
+
check_population_sizes(self.populations)
|
|
251
|
+
is_kc = np.zeros(self.n, dtype=np.uint8)
|
|
252
|
+
is_kc[self.populations["KC"]] = 1
|
|
253
|
+
# A plastic or sign-corrected brain rewrites weights, so it works on its own
|
|
254
|
+
# copy and the compiled graph (and its hashes) stays pristine.
|
|
255
|
+
weight = g["weight"].copy() if (self.plastic or self.r8_ame12_excitatory) else g["weight"]
|
|
256
|
+
self._r8_ame12 = self._apply_r8_ame12(g, weight)
|
|
257
|
+
self.kernel = Kernel(g["ptr"], g["post"], weight, g["modulatory"], is_kc)
|
|
258
|
+
self.r16 = self.populations["R1-R6"]
|
|
259
|
+
self.r8 = self.populations["R8"]
|
|
260
|
+
self.lamina = self.populations["lamina"]
|
|
261
|
+
self._drive = np.zeros(self.n, dtype=np.float64)
|
|
262
|
+
self._bin_counts = np.zeros(self.n, dtype=np.int32)
|
|
263
|
+
self.lp16 = np.zeros(len(self.r16), dtype=np.float64)
|
|
264
|
+
self.lp8 = np.zeros(len(self.r8), dtype=np.float64)
|
|
265
|
+
self.observations = 0
|
|
266
|
+
self.plasticity: KCMBONPlasticity | None = None
|
|
267
|
+
if self.plastic:
|
|
268
|
+
self.plasticity = KCMBONPlasticity(
|
|
269
|
+
self.kernel.ptr,
|
|
270
|
+
self.kernel.post,
|
|
271
|
+
self.kernel.weight,
|
|
272
|
+
self.populations,
|
|
273
|
+
plasticity_config,
|
|
274
|
+
)
|
|
275
|
+
self._hashes: dict[str, str] | None = None
|
|
276
|
+
|
|
277
|
+
def _apply_r8_ame12(self, g: dict[str, np.ndarray], weight: np.ndarray) -> dict[str, Any]:
|
|
278
|
+
"""Make the R8 to aMe12 edges excitatory on `weight` if the flag is on."""
|
|
279
|
+
edges = r8_ame12_edges(g)
|
|
280
|
+
before = float(g["weight"][edges].sum()) if len(edges) else 0.0
|
|
281
|
+
if self.r8_ame12_excitatory and len(edges):
|
|
282
|
+
weight[edges] = np.abs(weight[edges])
|
|
283
|
+
after = float(weight[edges].sum()) if len(edges) else 0.0
|
|
284
|
+
return {
|
|
285
|
+
"enabled": self.r8_ame12_excitatory,
|
|
286
|
+
"source": f"type starts with {R8_AME12_SOURCE_PREFIX}",
|
|
287
|
+
"target": f"type == {R8_AME12_TARGET_TYPE}",
|
|
288
|
+
"edges": int(len(edges)),
|
|
289
|
+
"edge_index_sha256": sha256_array(edges),
|
|
290
|
+
"weight_sha256": sha256_array(np.ascontiguousarray(weight[edges])),
|
|
291
|
+
"weight_sum_before_mv": before,
|
|
292
|
+
"weight_sum_after_mv": after,
|
|
293
|
+
"citation": R8_AME12_CITATION,
|
|
294
|
+
}
|
|
295
|
+
|
|
296
|
+
# ------------------------------------------------------------------
|
|
297
|
+
# Introspection
|
|
298
|
+
# ------------------------------------------------------------------
|
|
299
|
+
|
|
300
|
+
@property
|
|
301
|
+
def clock(self) -> int:
|
|
302
|
+
return self.kernel.clock
|
|
303
|
+
|
|
304
|
+
@property
|
|
305
|
+
def sim_ms(self) -> float:
|
|
306
|
+
return self.kernel.sim_ms
|
|
307
|
+
|
|
308
|
+
@property
|
|
309
|
+
def active_count(self) -> int:
|
|
310
|
+
return self.kernel.state.active_count
|
|
311
|
+
|
|
312
|
+
def population_sizes(self) -> dict[str, int]:
|
|
313
|
+
return {name: int(len(idx)) for name, idx in self.populations.items()}
|
|
314
|
+
|
|
315
|
+
def eye_map(self) -> EyeMap:
|
|
316
|
+
g = self._graph
|
|
317
|
+
return EyeMap(
|
|
318
|
+
uv_r16=np.asarray(g["r16_uv"], dtype=np.float32),
|
|
319
|
+
eye_r16=np.asarray(g["r16_eye"], dtype=np.int8),
|
|
320
|
+
uv_r8=np.asarray(g["r8_uv"], dtype=np.float32),
|
|
321
|
+
eye_r8=np.asarray(g["r8_eye"], dtype=np.int8),
|
|
322
|
+
r8_channel=np.asarray(g["r8_channel"], dtype=np.int8),
|
|
323
|
+
r16=self.r16,
|
|
324
|
+
r8=self.r8,
|
|
325
|
+
)
|
|
326
|
+
|
|
327
|
+
def graph_hashes(self) -> dict[str, str]:
|
|
328
|
+
"""SHA-256 per graph array (same digest as data/graph.lock.json)."""
|
|
329
|
+
if self._hashes is None:
|
|
330
|
+
self._hashes = {k: sha256_array(v) for k, v in sorted(self._graph.items())}
|
|
331
|
+
return dict(self._hashes)
|
|
332
|
+
|
|
333
|
+
def parameters(self) -> dict[str, Any]:
|
|
334
|
+
return {
|
|
335
|
+
**KERNEL_PARAMETERS,
|
|
336
|
+
"bin_ms": BIN_MS,
|
|
337
|
+
"luminance_tau_ms": LUMINANCE_TAU_MS,
|
|
338
|
+
"lamina_drive_mv": LAMINA_DRIVE_MV,
|
|
339
|
+
"photoreceptor_max_mv": PHOTORECEPTOR_MAX_MV,
|
|
340
|
+
"half_saturation": self.half_saturation,
|
|
341
|
+
"plastic": self.plastic,
|
|
342
|
+
"plasticity": self.plasticity.config.as_dict() if self.plasticity else None,
|
|
343
|
+
"random_sample_seed": RANDOM_SAMPLE_SEED,
|
|
344
|
+
"r8_ame12_excitatory": self.r8_ame12_excitatory,
|
|
345
|
+
"r8_ame12_edges": self._r8_ame12["edges"],
|
|
346
|
+
"r8_ame12_edges_sha256": self._r8_ame12["edge_index_sha256"],
|
|
347
|
+
"r8_ame12_weights_sha256": self._r8_ame12["weight_sha256"],
|
|
348
|
+
}
|
|
349
|
+
|
|
350
|
+
def provenance(self) -> dict:
|
|
351
|
+
return {
|
|
352
|
+
"kernel_version": KERNEL_VERSION,
|
|
353
|
+
"graph_format_version": GRAPH_FORMAT_VERSION,
|
|
354
|
+
"graph_path": self.graph_path,
|
|
355
|
+
"graph_hashes": self.graph_hashes(),
|
|
356
|
+
"n": self.n,
|
|
357
|
+
"edges": self.edges,
|
|
358
|
+
"parameters": self.parameters(),
|
|
359
|
+
"population_sizes": self.population_sizes(),
|
|
360
|
+
"population_definitions": {
|
|
361
|
+
"DN": f"superclass == {DESCENDING_SUPERCLASS}",
|
|
362
|
+
"central_complex": f"class == {CENTRAL_COMPLEX_CLASS}",
|
|
363
|
+
"random2000": f"seed {RANDOM_SAMPLE_SEED}, superclass {CENTRAL_BRAIN_PREFIX}* minus photoreceptors and lamina",
|
|
364
|
+
"modulatory": "no postsynaptic effect in the base model",
|
|
365
|
+
},
|
|
366
|
+
"assumptions": {"r8_ame12_excitatory": dict(self._r8_ame12)},
|
|
367
|
+
"plasticity_state": self.plasticity.summary() if self.plasticity else None,
|
|
368
|
+
"clock": self.clock,
|
|
369
|
+
"sim_ms": self.sim_ms,
|
|
370
|
+
"observations": self.observations,
|
|
371
|
+
}
|
|
372
|
+
|
|
373
|
+
def population_rates(self, counts: np.ndarray, neural_ms: float) -> dict[str, float]:
|
|
374
|
+
"""Mean firing rate in Hz per population for one observation's counts."""
|
|
375
|
+
seconds = neural_ms / 1000.0
|
|
376
|
+
out = {}
|
|
377
|
+
for name, idx in self.populations.items():
|
|
378
|
+
out[name] = float(counts[idx].mean() / seconds) if len(idx) and seconds > 0 else 0.0
|
|
379
|
+
return out
|
|
380
|
+
|
|
381
|
+
# ------------------------------------------------------------------
|
|
382
|
+
# Simulation
|
|
383
|
+
# ------------------------------------------------------------------
|
|
384
|
+
|
|
385
|
+
def reset(self) -> None:
|
|
386
|
+
"""Fresh mutable state: membranes at rest, clock 0, plastic factors 1."""
|
|
387
|
+
self.kernel.reset()
|
|
388
|
+
self.lp16[:] = 0.0
|
|
389
|
+
self.lp8[:] = 0.0
|
|
390
|
+
self.observations = 0
|
|
391
|
+
if self.plasticity is not None:
|
|
392
|
+
self.plasticity.reset()
|
|
393
|
+
|
|
394
|
+
def _luminance(self, values: Any, expected: int, name: str) -> np.ndarray:
|
|
395
|
+
arr = np.asarray(values, dtype=np.float64).reshape(-1)
|
|
396
|
+
if len(arr) != expected:
|
|
397
|
+
raise ValueError(f"stimulus.{name} has {len(arr)} values, expected {expected}")
|
|
398
|
+
if not np.all(np.isfinite(arr)):
|
|
399
|
+
raise ValueError(f"stimulus.{name} contains non-finite values")
|
|
400
|
+
return np.clip(arr, 0.0, 1.0)
|
|
401
|
+
|
|
402
|
+
def _photoreceptor_drive(self, lp: np.ndarray) -> np.ndarray:
|
|
403
|
+
return PHOTORECEPTOR_MAX_MV * lp / (self.half_saturation + lp)
|
|
404
|
+
|
|
405
|
+
def observe(self, stimulus: Stimulus, neural_ms: float) -> ObservationResult:
|
|
406
|
+
t0 = time.perf_counter()
|
|
407
|
+
n_steps = int(round(float(neural_ms) / K.DT_MS))
|
|
408
|
+
if n_steps <= 0:
|
|
409
|
+
raise ValueError("neural_ms must be at least one step (0.1 ms)")
|
|
410
|
+
r16 = self._luminance(stimulus.r16, len(self.r16), "r16")
|
|
411
|
+
r8 = self._luminance(stimulus.r8, len(self.r8), "r8")
|
|
412
|
+
pulses = []
|
|
413
|
+
for pulse in stimulus.pulses or ():
|
|
414
|
+
name, mv, duration = pulse
|
|
415
|
+
if name not in self.populations:
|
|
416
|
+
raise KeyError(f"pulse population {name!r} is not a known population")
|
|
417
|
+
if not (math.isfinite(mv) and math.isfinite(duration)) or duration < 0:
|
|
418
|
+
raise ValueError(
|
|
419
|
+
f"pulse {name!r}: current and duration must be finite and duration non-negative"
|
|
420
|
+
)
|
|
421
|
+
pulses.append((self.populations[name], float(mv), float(duration)))
|
|
422
|
+
|
|
423
|
+
counts = np.zeros(self.n, dtype=np.int32)
|
|
424
|
+
drive = self._drive
|
|
425
|
+
bin_counts = self._bin_counts
|
|
426
|
+
done = 0
|
|
427
|
+
while done < n_steps:
|
|
428
|
+
steps = min(BIN_STEPS, n_steps - done)
|
|
429
|
+
bin_ms = steps * K.DT_MS
|
|
430
|
+
alpha = 1.0 - math.exp(-bin_ms / LUMINANCE_TAU_MS)
|
|
431
|
+
self.lp16 += (r16 - self.lp16) * alpha
|
|
432
|
+
self.lp8 += (r8 - self.lp8) * alpha
|
|
433
|
+
drive.fill(0.0)
|
|
434
|
+
drive[self.lamina] = LAMINA_DRIVE_MV
|
|
435
|
+
drive[self.r16] = self._photoreceptor_drive(self.lp16)
|
|
436
|
+
drive[self.r8] = self._photoreceptor_drive(self.lp8)
|
|
437
|
+
elapsed_ms = done * K.DT_MS
|
|
438
|
+
for idx, mv, duration in pulses:
|
|
439
|
+
if elapsed_ms < duration:
|
|
440
|
+
drive[idx] += mv
|
|
441
|
+
bin_counts.fill(0)
|
|
442
|
+
self.kernel.set_drive(drive)
|
|
443
|
+
self.kernel.run(steps, bin_counts)
|
|
444
|
+
counts += bin_counts
|
|
445
|
+
if self.plasticity is not None:
|
|
446
|
+
self.plasticity.update(bin_counts, bin_ms)
|
|
447
|
+
done += steps
|
|
448
|
+
self.observations += 1
|
|
449
|
+
return ObservationResult(
|
|
450
|
+
counts=counts,
|
|
451
|
+
neural_ms=n_steps * K.DT_MS,
|
|
452
|
+
sim_ms=self.sim_ms,
|
|
453
|
+
compute_seconds=time.perf_counter() - t0,
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
# ------------------------------------------------------------------
|
|
457
|
+
# Checkpoints
|
|
458
|
+
# ------------------------------------------------------------------
|
|
459
|
+
|
|
460
|
+
def _metadata(self) -> dict[str, Any]:
|
|
461
|
+
return {
|
|
462
|
+
"checkpoint_version": CHECKPOINT_VERSION,
|
|
463
|
+
"kernel_version": KERNEL_VERSION,
|
|
464
|
+
"graph_format_version": GRAPH_FORMAT_VERSION,
|
|
465
|
+
"graph_hashes": self.graph_hashes(),
|
|
466
|
+
"n": self.n,
|
|
467
|
+
"edges": self.edges,
|
|
468
|
+
"parameters": self.parameters(),
|
|
469
|
+
"population_sizes": self.population_sizes(),
|
|
470
|
+
"observations": self.observations,
|
|
471
|
+
}
|
|
472
|
+
|
|
473
|
+
def checkpoint(self, path: str) -> None:
|
|
474
|
+
"""Save only mutable state plus metadata to a compressed npz."""
|
|
475
|
+
arrays: dict[str, np.ndarray] = {}
|
|
476
|
+
for k, v in self.kernel.state.to_arrays().items():
|
|
477
|
+
arrays["k_" + k] = v
|
|
478
|
+
arrays["lp16"] = self.lp16.copy()
|
|
479
|
+
arrays["lp8"] = self.lp8.copy()
|
|
480
|
+
if self.plasticity is not None:
|
|
481
|
+
for k, v in self.plasticity.state().items():
|
|
482
|
+
arrays["p_" + k] = v
|
|
483
|
+
arrays["meta"] = np.array(json.dumps(self._metadata(), sort_keys=True))
|
|
484
|
+
path = Path(path)
|
|
485
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
486
|
+
tmp = path.with_name(path.name + ".partial")
|
|
487
|
+
with open(tmp, "wb") as fh:
|
|
488
|
+
np.savez_compressed(fh, **arrays)
|
|
489
|
+
tmp.replace(path)
|
|
490
|
+
|
|
491
|
+
def restore(self, path: str) -> None:
|
|
492
|
+
"""Load a checkpoint written by `checkpoint`. Validates the metadata."""
|
|
493
|
+
with np.load(Path(path), allow_pickle=False) as z:
|
|
494
|
+
arrays = {k: z[k] for k in z.files}
|
|
495
|
+
meta = json.loads(str(arrays.pop("meta")))
|
|
496
|
+
mine = self._metadata()
|
|
497
|
+
problems = []
|
|
498
|
+
for key in ("checkpoint_version", "kernel_version", "graph_format_version", "n", "edges"):
|
|
499
|
+
if meta.get(key) != mine[key]:
|
|
500
|
+
problems.append(f"{key}: checkpoint {meta.get(key)!r}, brain {mine[key]!r}")
|
|
501
|
+
if meta.get("graph_hashes") != mine["graph_hashes"]:
|
|
502
|
+
changed = sorted(
|
|
503
|
+
k
|
|
504
|
+
for k in set(meta.get("graph_hashes", {})) | set(mine["graph_hashes"])
|
|
505
|
+
if meta.get("graph_hashes", {}).get(k) != mine["graph_hashes"].get(k)
|
|
506
|
+
)
|
|
507
|
+
problems.append("graph arrays differ: " + ", ".join(changed))
|
|
508
|
+
for key in ("half_saturation", "plastic", "plasticity", "r8_ame12_excitatory"):
|
|
509
|
+
if meta.get("parameters", {}).get(key) != mine["parameters"].get(key):
|
|
510
|
+
problems.append(
|
|
511
|
+
f"parameter {key}: checkpoint {meta.get('parameters', {}).get(key)!r}, brain {mine['parameters'].get(key)!r}"
|
|
512
|
+
)
|
|
513
|
+
if meta.get("population_sizes") != mine["population_sizes"]:
|
|
514
|
+
problems.append("population sizes differ")
|
|
515
|
+
if problems:
|
|
516
|
+
raise ValueError("checkpoint does not match this brain: " + "; ".join(problems))
|
|
517
|
+
kernel_arrays = {k[2:]: v for k, v in arrays.items() if k.startswith("k_")}
|
|
518
|
+
self.kernel.state = KernelState.from_arrays(kernel_arrays, self.kernel.rest)
|
|
519
|
+
self.lp16 = np.array(arrays["lp16"], dtype=np.float64)
|
|
520
|
+
self.lp8 = np.array(arrays["lp8"], dtype=np.float64)
|
|
521
|
+
if self.plasticity is not None:
|
|
522
|
+
self.plasticity.load_state({k[2:]: v for k, v in arrays.items() if k.startswith("p_")})
|
|
523
|
+
self.observations = int(meta.get("observations", 0))
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""MaleCNS v1.0 connectome: download, normalize, compile.
|
|
2
|
+
|
|
3
|
+
Requires `[feather]` for normalize.py; download.py uses urllib only.
|
|
4
|
+
"""
|
|
5
|
+
from flybrainer.connectome.sources import (
|
|
6
|
+
DATASET,
|
|
7
|
+
LICENSE,
|
|
8
|
+
SOURCES,
|
|
9
|
+
URL_PREFIX,
|
|
10
|
+
Source,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"SOURCES",
|
|
15
|
+
"Source",
|
|
16
|
+
"URL_PREFIX",
|
|
17
|
+
"LICENSE",
|
|
18
|
+
"DATASET",
|
|
19
|
+
]
|
|
20
|
+
|
|
21
|
+
# Lazy imports for sub-modules (require optional extras)
|
|
22
|
+
def __getattr__(name):
|
|
23
|
+
if name == "download_source":
|
|
24
|
+
from flybrainer.connectome.download import download_source
|
|
25
|
+
return download_source
|
|
26
|
+
if name == "is_verified":
|
|
27
|
+
from flybrainer.connectome.download import is_verified
|
|
28
|
+
return is_verified
|
|
29
|
+
if name == "DownloadError":
|
|
30
|
+
from flybrainer.connectome.download import DownloadError
|
|
31
|
+
return DownloadError
|
|
32
|
+
if name == "compile_graph":
|
|
33
|
+
from flybrainer.connectome.compile import compile_graph
|
|
34
|
+
return compile_graph
|
|
35
|
+
if name == "read_manifest":
|
|
36
|
+
from flybrainer.connectome.compile import read_manifest
|
|
37
|
+
return read_manifest
|
|
38
|
+
if name == "normalize_uv":
|
|
39
|
+
from flybrainer.connectome.compile import normalize_uv
|
|
40
|
+
return normalize_uv
|
|
41
|
+
if name == "photoreceptor_geometry":
|
|
42
|
+
from flybrainer.connectome.compile import photoreceptor_geometry
|
|
43
|
+
return photoreceptor_geometry
|
|
44
|
+
if name == "GRAPH_FORMAT_VERSION":
|
|
45
|
+
from flybrainer.connectome.compile import GRAPH_FORMAT_VERSION
|
|
46
|
+
return GRAPH_FORMAT_VERSION
|
|
47
|
+
if name == "sign_from_nt":
|
|
48
|
+
from flybrainer.connectome.normalize import sign_from_nt
|
|
49
|
+
return sign_from_nt
|
|
50
|
+
if name == "tokenize_nt":
|
|
51
|
+
from flybrainer.connectome.normalize import tokenize_nt
|
|
52
|
+
return tokenize_nt
|
|
53
|
+
if name == "load_annotations":
|
|
54
|
+
from flybrainer.connectome.normalize import load_annotations
|
|
55
|
+
return load_annotations
|
|
56
|
+
if name == "node_arrays":
|
|
57
|
+
from flybrainer.connectome.normalize import node_arrays
|
|
58
|
+
return node_arrays
|
|
59
|
+
raise AttributeError(f"module 'flybrainer.connectome' has no attribute {name!r}")
|