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 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}")