BitBrain: generic clean-room ADE + SBC library with MNIST acceptance
Implement the BitBrain (Address Decoder Element + Sparse Binary Coincidence) classifier as a generic, deterministic Nim library under common_libs/bitbrain/, written from the published algorithm (Front. Neuroinform. 17:1125844), not from the GPL-3.0 reference C. - ade.nim: signed thresholded random projection (scale 64 / centre 127 defaults reproduce the reference), multi-width ADs, optional deterministic homeostatic threshold adaptation. Hebbian longevity and Metropolis-Hastings sampling are described but not implemented. - sbc.nim: packed class-bit coincidence memory; idempotent learn, counting inference. - bitbrain.nim: container over several ADs and SBCs, online learn/infer, argmax readout, memory accounting. - tests: 32 unit checks (idempotence, planted rule + monotone online curve, shuffled-label chance control, unseen input, homeostasis, memory). - tests/test_bitbrain_mnist.nim: loads the reference pretrained ADs/thresholds and MNIST from /tmp, reproduces the reference exactly - 97.210% corrected and 96.540% bug-compatible - confirming the port. No gun/wiring integration yet; inputs and outputs to be agreed separately.
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
# BitBrain (ADE + SBC) — generic Nim library
|
||||
|
||||
A clean-room Nim implementation of the classifier in
|
||||
|
||||
> *BitBrain and Sparse Binary Coincidence (SBC) memories*,
|
||||
> Frontiers in Neuroinformatics 17:1125844, 2023.
|
||||
|
||||
**This is a library, not a gun.** There is no I/O design, no battle wiring and no
|
||||
environment knobs here yet — input and output formats are to be agreed separately.
|
||||
|
||||
## Clean-room note
|
||||
|
||||
The implementation was written **from the published algorithm description only**.
|
||||
No code was copied from the reference C program (`full_mnist_2048.c`) shipped with
|
||||
the paper; that file is GPL-3.0-or-later, © 2022 The University of Manchester, and
|
||||
copying it would impose that licence on this repository. The published defaults
|
||||
(`scale = 64`, `centre = 127`) make the scoring rule numerically identical to the
|
||||
reference, which the MNIST acceptance test confirms to the digit.
|
||||
|
||||
## Mechanism
|
||||
|
||||
1. **ADE (Address Decoder Element)** — a sparse, signed, thresholded random
|
||||
projection. Each ADE has `width` synapses `(inputIndex, polarity)` and scores an
|
||||
input as
|
||||
|
||||
```
|
||||
raw = Σ_j polarity_j * (input[inputIndex_j] - center)
|
||||
score = scale * raw
|
||||
```
|
||||
|
||||
firing iff `score >= threshold`. Defaults `scale = 64`, `center = 127`
|
||||
reproduce the reference exactly. Multi-width ADs (the paper's best setup uses
|
||||
widths `{6, 8, 10, 12}`) detect features at different scales.
|
||||
2. **Homeostatic threshold learning** (optional, unsupervised) — accumulate each
|
||||
ADE's firing count over inputs, then nudge its threshold toward a target firing
|
||||
rate (~1% in the paper). `accumulateFiring` + `adaptThresholds` implement the
|
||||
paper's deterministic controller. The paper's additional Hebbian *longevity*
|
||||
step (retire the weakest synapse, resample a new input index) and its
|
||||
Metropolis–Hastings input-position sampling are **described in `ade.nim` but
|
||||
not implemented** — `initRandomAddressDecoder` uses the paper's uniform random
|
||||
initialisation.
|
||||
3. **SBC memory** (supervised) — two ADs on the axes; a pair of simultaneously
|
||||
firing ADEs `(i, j)` is a coincidence indexing one cell holding a class
|
||||
bitmask. `learn` **sets** the class bit and is idempotent: setting it again is a
|
||||
no-op. No clearing, no learning rate, no decay, no epochs. `infer` accesses the
|
||||
same locations but **counts** set bits per class.
|
||||
4. **BitBrain container** — several ADs (possibly different widths) plus several
|
||||
SBCs built from pairs of them. Inference aggregates class counts across SBCs;
|
||||
argmax wins (ties to the lowest class index).
|
||||
|
||||
## API
|
||||
|
||||
```nim
|
||||
import bitbrain/bitbrain # re-exports ade + sbc
|
||||
|
||||
# --- construction ---
|
||||
var bb = buildRandomBitBrain(
|
||||
widths = @[6, 8, 10, 12], # one AD per width
|
||||
nAde = 512, # ADEs per AD
|
||||
inputWidth = 256, # input vector length
|
||||
nClasses = 8,
|
||||
seed = 1234'i64) # deterministic
|
||||
|
||||
# or build ADs yourself and assemble:
|
||||
# var ad = initAddressDecoder(nAde = 2048, width = 6)
|
||||
# ad.codes = ... # signed 1-based input codes (pretrained)
|
||||
# ad.thresholds = ...
|
||||
# var bb = initBitBrain(@[ad, ...], crossPairs(4) & withinPairs(4), nClasses = 10)
|
||||
|
||||
# --- online by construction (order-free, repeatable, interleaved) ---
|
||||
bb.learn(input, class) # sets class bits, idempotent
|
||||
let (label, counts) = bb.infer(input) # argmax + per-class counts
|
||||
bb.resetLearning() # wipe SBCs (ADs unchanged)
|
||||
|
||||
# --- optional unsupervised homeostasis ---
|
||||
for input in trainingStream:
|
||||
ad.accumulateFiring(input)
|
||||
ad.adaptThresholds(interval = 2000, targetRate = 0.01, step = 1)
|
||||
|
||||
# --- accounting ---
|
||||
bb.memoryBytes # ADs + SBC tensors
|
||||
bb.sbcMemoryBytes # SBC tensors only (the dominant term)
|
||||
```
|
||||
|
||||
Generic over the input element type (`openArray[SomeInteger]`) and over the input
|
||||
width, ADE count and class count — nothing is hardcoded to 784/10. Deterministic
|
||||
given the seed. Dependencies: `std/` only.
|
||||
|
||||
Reference SBC wiring: `crossPairs(4)` gives the 6 cross-AD SBCs used by the
|
||||
reference C. `withinPairs(n)` adds the paper's 4 within-AD ("half-size") SBCs;
|
||||
this implementation stores them full-size (half-size packing is a separate memory
|
||||
optimisation).
|
||||
|
||||
## Measured results
|
||||
|
||||
All figures measured on this machine with `-d:release`, single-threaded, on the
|
||||
reference fixtures (pretrained ADs + MNIST). Nothing is committed: fixtures live
|
||||
under `/tmp/bitbrain/BitBrain_C_code` (override with `$BITBRAIN_FIXTURES`).
|
||||
|
||||
### MNIST acceptance — exact reproduction of the reference
|
||||
|
||||
4 ADs × 2048 ADEs, widths `{6, 8, 10, 12}`, 6 cross-AD SBCs, 10 classes, one
|
||||
online pass over 60,000 training samples, evaluated on all 10,000 test images:
|
||||
|
||||
| Reader | This library | Reference C |
|
||||
|---|---:|---:|
|
||||
| Corrected (clean-room, all row ADEs counted) | **97.210%** | 97.210% |
|
||||
| Bug-compatible (`uint8_t bit_test`: only `i % 32 < 8`) | **96.540%** | 96.540% |
|
||||
|
||||
The bug-compatible mode reproduces the shipped reference's truncation bug exactly,
|
||||
which proves the ADE scoring, the coincidence indexing and the idempotent SBC rule
|
||||
are all correct. The default corrected reader is the one to use.
|
||||
|
||||
### Per-sample cost (reference MNIST configuration)
|
||||
|
||||
| Operation | Measured |
|
||||
|---|---:|
|
||||
| `learn` (4 ADs + 6 SBCs, dense active lists) | **0.384 ms/sample** |
|
||||
| `infer` (corrected reader) | **0.630 ms/sample** |
|
||||
| `infer` (bug-compatible reader) | 0.454 ms/sample |
|
||||
|
||||
All well inside this project's live budget of ~13.16 ms/tick. The reference C
|
||||
measured ~0.2–0.56 ms/sample for the same geometry.
|
||||
|
||||
### Memory (measured via `sbcMemoryBytes` / `memoryBytes`)
|
||||
|
||||
The SBC tensor dominates: `nAde × nAde × nClasses` bits per SBC, rounded up to
|
||||
`uint32` slots.
|
||||
|
||||
| Configuration | SBC tensors | ADs | Total |
|
||||
|---|---:|---:|---:|
|
||||
| Reference: 6 SBCs × 2048² × 10 | 31,457,280 B (30.0 MiB) | 360,448 B | 31,817,728 B (30.3 MiB) |
|
||||
| Gun-sized: 6 SBCs × 512² × 8 | 1,572,864 B (1.5 MiB) | 90,112 B | 1,662,976 B (1.59 MiB) |
|
||||
| 10 SBCs × 2048² × 10 (paper variant) | 52,428,800 B (50.0 MiB) | 360,448 B | 52,789,248 B (50.3 MiB) |
|
||||
|
||||
The reference configuration is L2/L3-hostile; a battle-sized configuration should
|
||||
size `nAde` to the feature count, not copy 2048².
|
||||
|
||||
## Tests
|
||||
|
||||
```bash
|
||||
# Unit / sanity suite (no fixtures required, ~3 s)
|
||||
nim c -r --nimcache:/tmp/nc_j90 -d:release \
|
||||
--path:common_libs common_libs/tests/test_bitbrain.nim
|
||||
|
||||
# MNIST acceptance harness (needs fixtures in /tmp; skips cleanly if absent)
|
||||
nim c -r --nimcache:/tmp/nc_j90 -d:release --path:common_libs \
|
||||
-o:/tmp/acceptance_bitbrain_mnist common_libs/tests/test_bitbrain_mnist.nim
|
||||
```
|
||||
|
||||
The unit suite covers: hand-computed ADE scoring and the `>=` threshold, idempotent
|
||||
SBC learning (learn one sample 1000× → bit-identical memory), a planted rule
|
||||
learned near-perfectly with a monotone online accuracy curve, a shuffled-label
|
||||
control that degrades to chance, unseen-input behaviour, homeostatic
|
||||
threshold adaptation, and memory accounting.
|
||||
@@ -0,0 +1,203 @@
|
||||
## Address Decoder Element (ADE) layer — clean-room implementation.
|
||||
##
|
||||
## This module implements the ADE thresholded random projection as described in
|
||||
##
|
||||
## "BitBrain and Sparse Binary Coincidence (SBC) memories",
|
||||
## Frontiers in Neuroinformatics 17:1125844, 2023.
|
||||
##
|
||||
## It was written **from the published algorithm description only**. No code was
|
||||
## copied from the reference C program (`full_mnist_2048.c`) that ships alongside
|
||||
## the paper, which is GPL-3.0-or-later, (c) 2022 The University of Manchester.
|
||||
## Keeping this a clean-room implementation avoids imposing that licence on this
|
||||
## repository. The published defaults (`scale = 64`, `center = 127`) make this
|
||||
## module numerically identical to the reference's scoring rule.
|
||||
##
|
||||
## Mechanism
|
||||
## ---------
|
||||
## A *synapse* is `(inputIndex, polarity)`. An ADE with `width` synapses scores an
|
||||
## input vector as
|
||||
##
|
||||
## raw = Σ_j polarity_j * (input[inputIndex_j] - center)
|
||||
## score = scale * raw
|
||||
##
|
||||
## and **fires** iff `score >= threshold`. The reference C uses `scale = 64`
|
||||
## and `center = 127` on raw `uint8` MNIST pixels, i.e. exactly the default
|
||||
## parameters here. Multi-width ADs (the paper's best setup uses widths in
|
||||
## {6, 8, 10, 12}) detect features at different scales.
|
||||
##
|
||||
## Storage
|
||||
## -------
|
||||
## Synapses are stored as a flat `seq[int32]` of signed 1-based codes: the sign
|
||||
## is the polarity, `abs(code) - 1` is the 0-based input index. This is byte-for-
|
||||
## byte the layout of the reference `AD*_2048` weight files, so a pretrained AD
|
||||
## can be loaded with a single `memcpy`-equivalent read.
|
||||
##
|
||||
## Unsupervised phase (threshold homeostasis)
|
||||
## ------------------------------------------
|
||||
## The paper learns each ADE's homeostatic threshold so that it fires at a target
|
||||
## rate (~1%). `accumulateFiring` + `adaptThresholds` below implement the
|
||||
## paper's simple deterministic controller: over an interval of `interval` inputs
|
||||
## the firing counts are compared against `targetRate * interval`, and each ADE's
|
||||
## threshold is nudged up or down by `step`.
|
||||
##
|
||||
## The paper's *additional* unsupervised mechanisms are deliberately NOT
|
||||
## implemented here (the task only requires them to be described):
|
||||
##
|
||||
## * **Hebbian "longevity" learning.** Each synapse carries a longevity
|
||||
## counter. When an ADE fires, the smallest contributor to the threshold
|
||||
## crossing has its longevity decremented and the largest contributor has it
|
||||
## incremented. After an interval, any synapse whose longevity drops below a
|
||||
## critical value is *retired* and replaced: a new input index is drawn by the
|
||||
## same sampling mechanism used at initialisation, with longevity reset to
|
||||
## the default. This lets each ADE "home in" on a feature. It is purely
|
||||
## unsupervised.
|
||||
## * **Metropolis-Hastings input-position sampling.** Synapse input positions
|
||||
## are drawn from a distribution proportional to `sqrt` of the global input
|
||||
## activity histogram (the `sqrt` both stabilises the Poisson uncertainty and
|
||||
## flattens the distribution so that synapses can land slightly outside the
|
||||
## training support). `initRandomAddressDecoder` instead samples uniformly at
|
||||
## random, which is the paper's stated initialisation.
|
||||
## * **Spatial clustering.** For image/volumetric data the sampled positions
|
||||
## may be constrained to a local region (rejection sampling until the
|
||||
## constraint holds) so each ADE sees a coherent receptive field.
|
||||
|
||||
import std/random
|
||||
|
||||
const
|
||||
DefaultScale* = 64
|
||||
## Score multiplier. 64 reproduces the reference C's `yang = ±64`.
|
||||
DefaultCenter* = 127
|
||||
## Input centre. 127 reproduces the reference C's `(pixel - 127)`.
|
||||
|
||||
type
|
||||
AddressDecoder* = object
|
||||
## An AD: `nAde` Address Decoder Elements, all of `width` synapses.
|
||||
nAde*: int ## the paper's `w`
|
||||
width*: int ## the paper's `n` (synapses per ADE)
|
||||
codes*: seq[int32] ## nAde*width; sign = polarity, |code|-1 = input index
|
||||
thresholds*: seq[int32] ## per-ADE firing thresholds
|
||||
scale*: int ## score multiplier (default 64)
|
||||
center*: int ## input centre (default 127)
|
||||
fireCounts*: seq[int32] ## firing accumulator for threshold homeostasis
|
||||
|
||||
proc initAddressDecoder*(nAde, width: int, scale = DefaultScale,
|
||||
center = DefaultCenter): AddressDecoder =
|
||||
## Allocate an AD with zeroed synapses. Callers either fill `codes` directly
|
||||
## (e.g. from a pretrained weight file) or use `initRandomAddressDecoder`.
|
||||
doAssert nAde > 0, "nAde must be positive"
|
||||
doAssert width > 0, "width must be positive"
|
||||
result.nAde = nAde
|
||||
result.width = width
|
||||
result.scale = scale
|
||||
result.center = center
|
||||
result.codes = newSeq[int32](nAde * width)
|
||||
result.thresholds = newSeq[int32](nAde)
|
||||
result.fireCounts = newSeq[int32](nAde)
|
||||
|
||||
proc synCode(inputIdx: int, polarity: int8): int32 {.inline.} =
|
||||
## Encode `(index, polarity)` as the reference's signed 1-based code.
|
||||
doAssert inputIdx >= 0
|
||||
doAssert polarity == 1'i8 or polarity == -1'i8
|
||||
result = int32(polarity) * int32(inputIdx + 1)
|
||||
|
||||
proc initRandomAddressDecoder*(nAde, width, inputWidth: int, rng: var Rand,
|
||||
scale = DefaultScale, center = DefaultCenter,
|
||||
threshold: int32 = 0'i32): AddressDecoder =
|
||||
## Randomly initialise an AD. For each ADE `width` **distinct** input indices
|
||||
## are drawn uniformly from `0 ..< inputWidth` (multapses are disallowed, as in
|
||||
## the paper) with a random ± polarity. Every ADE starts at `threshold`.
|
||||
doAssert inputWidth >= width,
|
||||
"cannot draw " & $width & " distinct indices from " & $inputWidth
|
||||
result = initAddressDecoder(nAde, width, scale, center)
|
||||
var chosen = newSeq[int32](width)
|
||||
for i in 0 ..< nAde:
|
||||
result.thresholds[i] = threshold
|
||||
var k = 0
|
||||
while k < width:
|
||||
let idx = int32(rng.rand(inputWidth - 1))
|
||||
var dup = false
|
||||
for p in 0 ..< k:
|
||||
if chosen[p] == idx:
|
||||
dup = true
|
||||
break
|
||||
if not dup:
|
||||
chosen[k] = idx
|
||||
inc k
|
||||
for j in 0 ..< width:
|
||||
let pol = if rng.rand(1) == 0: 1'i8 else: -1'i8
|
||||
result.codes[i * width + j] = synCode(int(chosen[j]), pol)
|
||||
|
||||
proc score*[T: SomeInteger](ad: AddressDecoder, input: openArray[T],
|
||||
ade: int): int {.inline.} =
|
||||
## Signed, centred, scaled score of ADE `ade` on `input`.
|
||||
## `score = scale * Σ_j polarity_j * (input[idx_j] - center)`.
|
||||
doAssert ade >= 0 and ade < ad.nAde
|
||||
var raw = 0
|
||||
let base = ade * ad.width
|
||||
for j in 0 ..< ad.width:
|
||||
let c = ad.codes[base + j]
|
||||
let idx =
|
||||
if c > 0'i32: int(c) - 1
|
||||
else: int(-c) - 1
|
||||
let pol = if c > 0'i32: 1 else: -1
|
||||
raw += pol * (int(input[idx]) - ad.center)
|
||||
result = ad.scale * raw
|
||||
|
||||
proc fires*[T: SomeInteger](ad: AddressDecoder, input: openArray[T],
|
||||
ade: int): bool {.inline.} =
|
||||
## An ADE fires iff its score reaches its threshold (`score >= threshold`).
|
||||
score(ad, input, ade) >= int(ad.thresholds[ade])
|
||||
|
||||
proc activeList*[T: SomeInteger](ad: AddressDecoder, input: openArray[T],
|
||||
active: var seq[int32]) =
|
||||
## Fill `active` with the indices of the ADEs that fire on `input`.
|
||||
## This is the sparse representation consumed by the SBC memory.
|
||||
active.setLen(0)
|
||||
for i in 0 ..< ad.nAde:
|
||||
if fires(ad, input, i):
|
||||
active.add int32(i)
|
||||
|
||||
proc fireCount*[T: SomeInteger](ad: AddressDecoder, input: openArray[T]): int =
|
||||
## Number of ADEs in `ad` that fire on `input`.
|
||||
for i in 0 ..< ad.nAde:
|
||||
if fires(ad, input, i):
|
||||
inc result
|
||||
|
||||
proc accumulateFiring*[T: SomeInteger](ad: var AddressDecoder,
|
||||
input: openArray[T]) =
|
||||
## Add this input's firing events to the homeostasis accumulator.
|
||||
for i in 0 ..< ad.nAde:
|
||||
if fires(ad, input, i):
|
||||
inc ad.fireCounts[i]
|
||||
|
||||
proc resetFiringCounts*(ad: var AddressDecoder) =
|
||||
## Zero the homeostasis accumulator (does not touch thresholds).
|
||||
for i in 0 ..< ad.nAde:
|
||||
ad.fireCounts[i] = 0
|
||||
|
||||
proc adaptThresholds*(ad: var AddressDecoder, interval: int, targetRate: float,
|
||||
step: int = 1) =
|
||||
## One step of the paper's homeostatic threshold controller.
|
||||
##
|
||||
## `fireCounts` is expected to hold the firing counts accumulated over the last
|
||||
## `interval` inputs (see `accumulateFiring`). Each ADE whose observed rate is
|
||||
## above/below `targetRate` has its threshold raised/lowered by `step`. The
|
||||
## counters are then reset, ready for the next interval.
|
||||
##
|
||||
## This is a deterministic integer controller: given the same inputs and the
|
||||
## same starting thresholds, it always produces the same thresholds.
|
||||
doAssert interval > 0, "interval must be positive"
|
||||
let expected = targetRate * float(interval)
|
||||
for i in 0 ..< ad.nAde:
|
||||
let excess = float(ad.fireCounts[i]) - expected
|
||||
if excess > 0.5:
|
||||
ad.thresholds[i] += int32(step)
|
||||
elif excess < -0.5:
|
||||
ad.thresholds[i] -= int32(step)
|
||||
ad.fireCounts[i] = 0'i32
|
||||
|
||||
proc memoryBytes*(ad: AddressDecoder): int =
|
||||
## Bytes held by this AD (synapse codes, thresholds, firing counters).
|
||||
ad.codes.len * sizeof(int32) +
|
||||
ad.thresholds.len * sizeof(int32) +
|
||||
ad.fireCounts.len * sizeof(int32)
|
||||
@@ -0,0 +1,136 @@
|
||||
## BitBrain container — ADs + SBC memories + a counting readout.
|
||||
##
|
||||
## Clean-room implementation of the classification pipeline described in
|
||||
##
|
||||
## "BitBrain and Sparse Binary Coincidence (SBC) memories",
|
||||
## Frontiers in Neuroinformatics 17:1125844, 2023.
|
||||
##
|
||||
## See `ade.nim` for the clean-room note (the reference C is GPL-3.0 and was not
|
||||
## copied).
|
||||
##
|
||||
## The container holds several Address Decoders (possibly of different widths)
|
||||
## and several SBC memories, each built from a pair of ADs. `learn` populates the
|
||||
## SBCs from a labelled sample; `infer` drives every AD, reads every SBC with the
|
||||
## same coincidence rule and sums the per-class set-bit counts. The class with the
|
||||
## highest total wins. Everything is per-sample and order-free: the memory is
|
||||
## monotone, so `learn` and `infer` may be interleaved arbitrarily.
|
||||
##
|
||||
## Reference setup used by the paper and the MNIST acceptance harness: 4 ADs of
|
||||
## 2048 ADEs with widths {6, 8, 10, 12} and 6 cross-AD SBCs. The paper's
|
||||
## 10-SBC variant adds 4 within-AD ("half-size") SBCs; `withinPairs` builds those
|
||||
## (this implementation stores them full-size — the half-size packing is a
|
||||
## separate memory optimisation).
|
||||
|
||||
import std/random
|
||||
import ade, sbc
|
||||
|
||||
export ade, sbc
|
||||
|
||||
type
|
||||
SbcSpec* = object
|
||||
## Which pair of ADs feeds one SBC memory.
|
||||
row*: int
|
||||
col*: int
|
||||
|
||||
BitBrain* = object
|
||||
ades*: seq[AddressDecoder]
|
||||
sbcs*: seq[Sbc]
|
||||
specs*: seq[SbcSpec]
|
||||
nClasses*: int
|
||||
|
||||
proc crossPairs*(nAdes: int): seq[SbcSpec] =
|
||||
## All unordered pairs of distinct ADs. For 4 ADs this is the paper's 6 SBCs.
|
||||
for a in 0 ..< nAdes:
|
||||
for b in a + 1 ..< nAdes:
|
||||
result.add SbcSpec(row: a, col: b)
|
||||
|
||||
proc withinPairs*(nAdes: int): seq[SbcSpec] =
|
||||
## The square/within-AD pairs (the paper's "half-size" SBCs).
|
||||
for a in 0 ..< nAdes:
|
||||
result.add SbcSpec(row: a, col: a)
|
||||
|
||||
proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec],
|
||||
nClasses: int): BitBrain =
|
||||
## Build the container. Every AD must have the same number of ADEs because an
|
||||
## SBC's axes are both `w` long (the paper's setup).
|
||||
doAssert ades.len > 0, "need at least one AD"
|
||||
doAssert nClasses > 0, "need at least one class"
|
||||
let w = ades[0].nAde
|
||||
for ad in ades:
|
||||
doAssert ad.nAde == w, "all ADs must have the same number of ADEs"
|
||||
result.ades = ades
|
||||
result.nClasses = nClasses
|
||||
result.specs = specs
|
||||
result.sbcs = newSeq[Sbc](specs.len)
|
||||
for s in 0 ..< specs.len:
|
||||
doAssert specs[s].row >= 0 and specs[s].row < ades.len
|
||||
doAssert specs[s].col >= 0 and specs[s].col < ades.len
|
||||
result.sbcs[s] = initSbc(w, nClasses)
|
||||
|
||||
proc buildRandomBitBrain*(widths: seq[int], nAde, inputWidth, nClasses: int,
|
||||
seed: int64,
|
||||
specs: seq[SbcSpec] = @[]): BitBrain =
|
||||
## Initialise a whole BitBrain with random ADs. If `specs` is empty the
|
||||
## paper's 6 cross-AD SBCs are used. Fully deterministic given `seed`.
|
||||
var rng = initRand(seed)
|
||||
var ades: seq[AddressDecoder]
|
||||
for w in widths:
|
||||
ades.add initRandomAddressDecoder(nAde, w, inputWidth, rng)
|
||||
let pairs = if specs.len == 0: crossPairs(widths.len) else: specs
|
||||
result = initBitBrain(ades, pairs, nClasses)
|
||||
|
||||
proc nAdes*(bb: BitBrain): int {.inline.} =
|
||||
bb.ades.len
|
||||
|
||||
proc resetLearning*(bb: var BitBrain) =
|
||||
## Wipe every SBC. The ADs (and their thresholds) are left untouched.
|
||||
for s in 0 ..< bb.sbcs.len:
|
||||
bb.sbcs[s].clear()
|
||||
|
||||
proc fireInto*[T: SomeInteger](bb: BitBrain, input: openArray[T],
|
||||
lists: var seq[seq[int32]]) =
|
||||
## Compute the sparse firing pattern of every AD for `input`. `lists` is resized
|
||||
## to `nAdes` and each entry is filled with that AD's firing ADE indices.
|
||||
if lists.len != bb.ades.len:
|
||||
lists.setLen(bb.ades.len)
|
||||
for a in 0 ..< bb.ades.len:
|
||||
bb.ades[a].activeList(input, lists[a])
|
||||
|
||||
proc learn*[T: SomeInteger](bb: var BitBrain, input: openArray[T], class: int) =
|
||||
## One online supervised step: set the `class` bit of every observed
|
||||
## coincidence. Idempotent — repeating the same sample is a no-op.
|
||||
var lists: seq[seq[int32]]
|
||||
bb.fireInto(input, lists)
|
||||
for s in 0 ..< bb.sbcs.len:
|
||||
let spec = bb.specs[s]
|
||||
discard bb.sbcs[s].learn(lists[spec.row], lists[spec.col], class)
|
||||
|
||||
proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
|
||||
tuple[label: int, counts: seq[int]] =
|
||||
## Drive every AD, count set class bits in every SBC, and return the argmax
|
||||
## class plus the aggregated per-class counts. Ties go to the lowest class
|
||||
## index (matching the reference's `>` scan that keeps the first maximum).
|
||||
var lists: seq[seq[int32]]
|
||||
bb.fireInto(input, lists)
|
||||
result.counts = newSeq[int](bb.nClasses)
|
||||
for s in 0 ..< bb.sbcs.len:
|
||||
let spec = bb.specs[s]
|
||||
bb.sbcs[s].infer(lists[spec.row], lists[spec.col], result.counts)
|
||||
result.label = 0
|
||||
var best = result.counts[0]
|
||||
for k in 1 ..< bb.nClasses:
|
||||
if result.counts[k] > best:
|
||||
best = result.counts[k]
|
||||
result.label = k
|
||||
|
||||
proc memoryBytes*(bb: BitBrain): int =
|
||||
## Total bytes: AD synapse codes + thresholds + counters + SBC bit tensors.
|
||||
for ad in bb.ades:
|
||||
result += ad.memoryBytes
|
||||
for s in bb.sbcs:
|
||||
result += s.memoryBytes
|
||||
|
||||
proc sbcMemoryBytes*(bb: BitBrain): int =
|
||||
## Just the SBC bit tensors — the dominant term for realistic configurations.
|
||||
for s in bb.sbcs:
|
||||
result += s.memoryBytes
|
||||
@@ -0,0 +1,114 @@
|
||||
## Sparse Binary Coincidence (SBC) memory — clean-room implementation.
|
||||
##
|
||||
## Implements the supervised half of the BitBrain algorithm as described in
|
||||
##
|
||||
## "BitBrain and Sparse Binary Coincidence (SBC) memories",
|
||||
## Frontiers in Neuroinformatics 17:1125844, 2023.
|
||||
##
|
||||
## Written from the published algorithm only; see `ade.nim` for the clean-room
|
||||
## note. The reference C is GPL-3.0 (c) University of Manchester and was not
|
||||
## copied.
|
||||
##
|
||||
## Mechanism
|
||||
## ---------
|
||||
## Two Address Decoders (ADs) sit on the two axes of a 2-D memory. A pair of
|
||||
## *simultaneously firing* ADEs `(i, j)` is a **coincidence** and addresses one
|
||||
## memory cell. That cell holds a class bitmask with one bit per class
|
||||
## (`nClasses` bits, one-hot encoding in the paper's default).
|
||||
##
|
||||
## Learning is **idempotent**: `learn` *sets* the bit for the observed class;
|
||||
## setting it again is a no-op. There is no clearing, no learning rate, no decay
|
||||
## and no epoch — one pass through the data is a complete supervised training run,
|
||||
## and a second pass changes nothing.
|
||||
##
|
||||
## Inference uses the *same* address decoding: for each observed coincidence every
|
||||
## class bit is read and the set bits are **counted** per class. Counts are summed
|
||||
## across SBCs by the `bitbrain` container and the argmax wins.
|
||||
##
|
||||
## Bit layout
|
||||
## ----------
|
||||
## The bit for `(i, j, class)` lives at `((i * nAde) + j) * nClasses + class`, so
|
||||
## all `nClasses` bits of one coincidence are contiguous. This differs from the
|
||||
## reference C's layout but is a bijection onto the same set of triples; the
|
||||
## learned rule is identical.
|
||||
|
||||
import std/bitops
|
||||
|
||||
type
|
||||
Sbc* = object
|
||||
## A 2-D coincidence memory with a class-bit depth.
|
||||
nAde*: int ## number of ADEs on each axis (the paper's `w`)
|
||||
nClasses*: int ## number of classes (the paper's `D`)
|
||||
bits*: seq[uint32] ## packed bit tensor, nAde*nAde*nClasses bits
|
||||
|
||||
proc initSbc*(nAde, nClasses: int): Sbc =
|
||||
## Allocate a zeroed SBC (nothing is known yet).
|
||||
doAssert nAde > 0, "nAde must be positive"
|
||||
doAssert nClasses > 0, "nClasses must be positive"
|
||||
result.nAde = nAde
|
||||
result.nClasses = nClasses
|
||||
let nbits = nAde * nAde * nClasses
|
||||
result.bits = newSeq[uint32]((nbits + 31) div 32)
|
||||
|
||||
proc clear*(sbc: var Sbc) =
|
||||
## Forget everything. This is how a monotone memory is wiped (e.g. when a bot
|
||||
## switches enemy and must not carry state across battles).
|
||||
for i in 0 ..< sbc.bits.len:
|
||||
sbc.bits[i] = 0'u32
|
||||
|
||||
proc bitIndex(sbc: Sbc, i, j, class: int): int {.inline.} =
|
||||
((i * sbc.nAde) + j) * sbc.nClasses + class
|
||||
|
||||
proc bitAt*(sbc: Sbc, i, j, class: int): bool {.inline.} =
|
||||
## Read one memory bit. Exposed mainly so harnesses can inspect the exact rule.
|
||||
doAssert i >= 0 and i < sbc.nAde
|
||||
doAssert j >= 0 and j < sbc.nAde
|
||||
doAssert class >= 0 and class < sbc.nClasses
|
||||
let bit = bitIndex(sbc, i, j, class)
|
||||
(sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32
|
||||
|
||||
proc learn*(sbc: var Sbc, rowActive, colActive: openArray[int32],
|
||||
class: int): int =
|
||||
## Set the `class` bit for every coincidence between a firing row ADE and a
|
||||
## firing column ADE. Returns the number of bits that were newly set (0 if the
|
||||
## sample added no information, e.g. it was already learned).
|
||||
doAssert class >= 0 and class < sbc.nClasses, "class out of range"
|
||||
let D = sbc.nClasses
|
||||
for r in rowActive:
|
||||
let i = int(r)
|
||||
for c in colActive:
|
||||
let j = int(c)
|
||||
let bit = ((i * sbc.nAde) + j) * D + class
|
||||
let w = bit shr 5
|
||||
let m = 1'u32 shl (bit and 31)
|
||||
if (sbc.bits[w] and m) == 0'u32:
|
||||
sbc.bits[w] = sbc.bits[w] or m
|
||||
inc result
|
||||
|
||||
proc infer*(sbc: Sbc, rowActive, colActive: openArray[int32],
|
||||
counts: var seq[int]) =
|
||||
## Count, per class, how many observed coincidences have their class bit set.
|
||||
## `counts` is *accumulated into* (not reset), so a container can sum several
|
||||
## SBCs. It must be at least `nClasses` long.
|
||||
doAssert counts.len >= sbc.nClasses, "counts buffer too small"
|
||||
let D = sbc.nClasses
|
||||
for r in rowActive:
|
||||
let i = int(r)
|
||||
for c in colActive:
|
||||
let j = int(c)
|
||||
let base = ((i * sbc.nAde) + j) * D
|
||||
for k in 0 ..< D:
|
||||
let bit = base + k
|
||||
if (sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32:
|
||||
inc counts[k]
|
||||
|
||||
proc memoryBytes*(sbc: Sbc): int =
|
||||
## Bytes held by the packed bit tensor.
|
||||
sbc.bits.len * sizeof(uint32)
|
||||
|
||||
proc occupancy*(sbc: Sbc): float =
|
||||
## Fraction of the bit tensor that is set (diagnostic).
|
||||
var setBits = 0
|
||||
for w in sbc.bits:
|
||||
setBits += countSetBits(w)
|
||||
result = float(setBits) / float(sbc.nAde * sbc.nAde * sbc.nClasses)
|
||||
@@ -0,0 +1,253 @@
|
||||
## Unit / sanity tests for the generic BitBrain (ADE + SBC) library.
|
||||
##
|
||||
## No Java, no battles, no I/O: everything here is deterministic and seeded.
|
||||
## Run:
|
||||
## nim c -r --nimcache:/tmp/nc_j90 common_libs/tests/test_bitbrain.nim
|
||||
##
|
||||
## Covers:
|
||||
## * ADE scoring / thresholding against a hand-computed example,
|
||||
## * idempotent SBC learning (learn one sample 1000x -> identical memory),
|
||||
## * a planted rule learned near-perfectly,
|
||||
## * an online (incremental) learning curve that improves monotonically,
|
||||
## * a shuffled-label control that degrades to chance,
|
||||
## * correct, non-crashing behaviour on an unseen input,
|
||||
## * homeostatic threshold adaptation moves the firing rate toward target.
|
||||
|
||||
import std/[random, math, strutils]
|
||||
import bitbrain/ade
|
||||
import bitbrain/sbc
|
||||
import bitbrain/bitbrain
|
||||
|
||||
var checks = 0
|
||||
var failures = 0
|
||||
proc check(name: string, ok: bool) =
|
||||
inc checks
|
||||
if ok: echo "PASS: ", name
|
||||
else: echo "FAIL: ", name; inc failures
|
||||
|
||||
proc randInput(rng: var Rand, n: int): seq[int] =
|
||||
result = newSeq[int](n)
|
||||
for i in 0 ..< n:
|
||||
result[i] = rng.rand(255)
|
||||
|
||||
# ── 1. ADE scoring / thresholding ────────────────────────────────────────────
|
||||
|
||||
proc testAdeScoring() =
|
||||
# One ADE with synapses: +index 2 (code +3), -index 0 (code -1).
|
||||
var ad = initAddressDecoder(1, 2, scale = 1, center = 127)
|
||||
ad.codes[0] = 3 # +1 * (2 + 1): position 2, excitatory
|
||||
ad.codes[1] = -1 # -1 * (0 + 1): position 0, inhibitory
|
||||
let input = @[10, 0, 200] # raw = (200-127) - (10-127) = 73 + 117 = 190
|
||||
check "ADE hand-computed score is 190", ad.score(input, 0) == 190
|
||||
ad.thresholds[0] = 190
|
||||
check "ADE fires at score == threshold (>= comparison)", ad.fires(input, 0)
|
||||
ad.thresholds[0] = 191
|
||||
check "ADE does not fire one above score", not ad.fires(input, 0)
|
||||
|
||||
# Two-ADE active list.
|
||||
var ad2 = initAddressDecoder(3, 1, scale = 1, center = 0)
|
||||
ad2.codes = @[1'i32, 2'i32, 3'i32] # positions 0,1,2, all excitatory
|
||||
ad2.thresholds = @[5'i32, 5'i32, 5'i32]
|
||||
var active: seq[int32]
|
||||
ad2.activeList(@[10, 1, 10], active)
|
||||
check "active list picks exactly the firing ADEs", active == @[0'i32, 2'i32]
|
||||
|
||||
# ── 2. Idempotent SBC learning ───────────────────────────────────────────────
|
||||
|
||||
proc testIdempotentLearning() =
|
||||
var s = initSbc(64, 5)
|
||||
let row = @[1'i32, 3'i32, 7'i32]
|
||||
let col = @[2'i32, 4'i32]
|
||||
let firstAdded = s.learn(row, col, 2)
|
||||
check "first learn sets 3*2 = 6 bits", firstAdded == 6
|
||||
let snapshot = s.bits
|
||||
var allNoop = true
|
||||
for _ in 0 ..< 1000:
|
||||
if s.learn(row, col, 2) != 0: allNoop = false
|
||||
check "learning the same sample 1000x is a no-op", allNoop
|
||||
check "memory is bit-identical after 1000 repeats", s.bits == snapshot
|
||||
|
||||
# Setting a *different* class bit is still new information.
|
||||
check "a different class sets new bits", s.learn(row, col, 3) == 6
|
||||
check "learned bit is readable", s.bitAt(1, 2, 3)
|
||||
check "unlearned bit is clear", not s.bitAt(1, 2, 4)
|
||||
|
||||
# ── 3. Synthetic planted-rule dataset ────────────────────────────────────────
|
||||
|
||||
const
|
||||
SynthInputWidth = 32
|
||||
SynthClasses = 4
|
||||
SynthBlock = 8 # positions owned by each class
|
||||
SynthHot = 4 # of the 8 block positions set per sample
|
||||
|
||||
proc synthSample(rng: var Rand, cls: int): seq[int] =
|
||||
## Input is zero everywhere except `SynthHot` randomly chosen positions inside
|
||||
## class `cls`'s block, set to 255. The class rules are disjoint in position
|
||||
## space, so a rule of "which positions are hot" is perfectly learnable.
|
||||
result = newSeq[int](SynthInputWidth)
|
||||
var poss: seq[int]
|
||||
for p in 0 ..< SynthBlock:
|
||||
poss.add cls * SynthBlock + p
|
||||
rng.shuffle(poss)
|
||||
for k in 0 ..< SynthHot:
|
||||
result[poss[k]] = 255
|
||||
|
||||
proc synthAd(nAde: int, rng: var Rand): AddressDecoder =
|
||||
## All-excitatory width-2 ADEs that fire iff *both* sampled positions are hot
|
||||
## (score 2*255 == threshold 510).
|
||||
result = initRandomAddressDecoder(nAde, 2, SynthInputWidth, rng,
|
||||
scale = 1, center = 0, threshold = 510)
|
||||
for k in 0 ..< result.codes.len:
|
||||
result.codes[k] = abs(result.codes[k])
|
||||
|
||||
proc synthBrain(nAde: int, seed: int64): BitBrain =
|
||||
var rng = initRand(seed)
|
||||
var ades: seq[AddressDecoder]
|
||||
for _ in 0 ..< 3: # a few ADs -> several cross SBCs
|
||||
ades.add synthAd(nAde, rng)
|
||||
initBitBrain(ades, crossPairs(ades.len), SynthClasses)
|
||||
|
||||
proc accuracy(bb: var BitBrain, xs: seq[seq[int]], ys: seq[int]): float =
|
||||
var right = 0
|
||||
for i in 0 ..< xs.len:
|
||||
if bb.infer(xs[i]).label == ys[i]:
|
||||
inc right
|
||||
result = float(right) / float(xs.len)
|
||||
|
||||
proc makeSynthSet(n: int, seed: int64): (seq[seq[int]], seq[int]) =
|
||||
var rng = initRand(seed)
|
||||
for i in 0 ..< n:
|
||||
let cls = i mod SynthClasses
|
||||
result[0].add synthSample(rng, cls)
|
||||
result[1].add cls
|
||||
|
||||
proc testPlantedRuleAndOnlineCurve() =
|
||||
var bb = synthBrain(256, seed = 20240924)
|
||||
var (testX, testY) = makeSynthSet(400, seed = 777)
|
||||
var (trainX, trainY) = makeSynthSet(600, seed = 111)
|
||||
|
||||
# Incremental online learning: accuracy on held-out data must not go backwards
|
||||
# as samples arrive.
|
||||
var accs: seq[float]
|
||||
for i in 0 ..< trainX.len:
|
||||
bb.learn(trainX[i], trainY[i])
|
||||
if (i + 1) mod 50 == 0:
|
||||
accs.add accuracy(bb, testX, testY)
|
||||
|
||||
for i in 1 ..< accs.len:
|
||||
check "online accuracy is monotone at checkpoint " & $((i + 1) * 50) &
|
||||
" (" & formatFloat(accs[i - 1], ffDecimal, 3) & " -> " &
|
||||
formatFloat(accs[i], ffDecimal, 3) & ")",
|
||||
accs[i] >= accs[i - 1]
|
||||
check "planted rule is learned near-perfectly (" &
|
||||
formatFloat(accs[^1], ffDecimal, 3) & " >= 0.95)",
|
||||
accs[^1] >= 0.95
|
||||
|
||||
# ── 4. Shuffled-label control ────────────────────────────────────────────────
|
||||
|
||||
proc testShuffledControl() =
|
||||
var bb = synthBrain(256, seed = 99)
|
||||
var (trainX, _) = makeSynthSet(3000, seed = 5)
|
||||
var (testX, testY) = makeSynthSet(400, seed = 6)
|
||||
var rng = initRand(4242)
|
||||
var shufY = newSeq[int](trainX.len)
|
||||
for i in 0 ..< trainX.len:
|
||||
shufY[i] = rng.rand(SynthClasses - 1)
|
||||
for i in 0 ..< trainX.len:
|
||||
bb.learn(trainX[i], shufY[i])
|
||||
let acc = accuracy(bb, testX, testY)
|
||||
check "shuffled-label control is near chance (" &
|
||||
formatFloat(acc, ffDecimal, 3) & " <= 0.45, chance = 0.25)",
|
||||
acc <= 0.45
|
||||
|
||||
# ── 5. Unseen input ──────────────────────────────────────────────────────────
|
||||
|
||||
proc testUnseenInput() =
|
||||
var bb = synthBrain(128, seed = 33)
|
||||
var (trainX, trainY) = makeSynthSet(400, seed = 44)
|
||||
for i in 0 ..< trainX.len:
|
||||
bb.learn(trainX[i], trainY[i])
|
||||
let blank = newSeq[int](SynthInputWidth)
|
||||
let r = bb.infer(blank)
|
||||
var total = 0
|
||||
for c in r.counts: total += c
|
||||
check "unseen all-zero input fires no ADE and has zero counts", total == 0
|
||||
check "unseen all-zero input still returns a valid class",
|
||||
r.label >= 0 and r.label < SynthClasses
|
||||
# A single hot position inside a class block should not crash and should
|
||||
# produce valid counts.
|
||||
var rng = initRand(1)
|
||||
var one = newSeq[int](SynthInputWidth)
|
||||
one[3 * SynthBlock + 2] = 255
|
||||
let r2 = bb.infer(one)
|
||||
check "unseen partial input returns a valid class/counts",
|
||||
r2.label >= 0 and r2.label < SynthClasses and r2.counts.len == SynthClasses
|
||||
|
||||
# ── 6. Homeostatic threshold adaptation ──────────────────────────────────────
|
||||
proc testHomeostasis() =
|
||||
var rng = initRand(7)
|
||||
const N = 64
|
||||
const Interval = 400
|
||||
const Target = 0.10
|
||||
|
||||
# All thresholds so high nothing fires -> controller must lower them.
|
||||
var hot = initRandomAddressDecoder(N, 4, 200, rng, scale = 1, center = 127,
|
||||
threshold = 100_000)
|
||||
var cold = initRandomAddressDecoder(N, 4, 200, rng, scale = 1, center = 127,
|
||||
threshold = -100_000)
|
||||
# A naturally-initialised AD (threshold 0) is used to test convergence.
|
||||
var mid = initRandomAddressDecoder(N, 4, 200, rng, scale = 1, center = 127,
|
||||
threshold = 0)
|
||||
for round in 0 ..< 300:
|
||||
for _ in 0 ..< Interval:
|
||||
let x = randInput(rng, 200)
|
||||
hot.accumulateFiring(x)
|
||||
cold.accumulateFiring(x)
|
||||
mid.accumulateFiring(x)
|
||||
hot.adaptThresholds(Interval, Target, step = 5)
|
||||
cold.adaptThresholds(Interval, Target, step = 5)
|
||||
mid.adaptThresholds(Interval, Target, step = 5)
|
||||
check "too-cold thresholds are driven down", hot.thresholds[0] < 100_000
|
||||
check "too-hot thresholds are driven up", cold.thresholds[0] > -100_000
|
||||
|
||||
# Measure the converged firing rate of the naturally-initialised AD.
|
||||
var fired = 0
|
||||
for _ in 0 ..< 2000:
|
||||
fired += mid.fireCount(randInput(rng, 200))
|
||||
let rate = float(fired) / float(2000 * N)
|
||||
check "mid AD converges near the 10% target (got " &
|
||||
formatFloat(rate, ffDecimal, 3) & ")",
|
||||
rate > 0.05 and rate < 0.20
|
||||
|
||||
# ── 7. Memory accounting ─────────────────────────────────────────────────────
|
||||
|
||||
proc testMemoryAccounting() =
|
||||
# A gun-sized configuration: 4 ADs x 512 ADEs, widths {6,8,10,12}, 6 SBCs,
|
||||
# 8 classes. SBC tensors dominate: 6 * 512*512*8 bits = 1,572,864 bytes.
|
||||
var ades: seq[AddressDecoder]
|
||||
var rng = initRand(1)
|
||||
for w in [6, 8, 10, 12]:
|
||||
ades.add initRandomAddressDecoder(512, w, 256, rng)
|
||||
let bb = initBitBrain(ades, crossPairs(4), 8)
|
||||
check "gun-sized SBC memory = 1,572,864 bytes",
|
||||
bb.sbcMemoryBytes == 1_572_864
|
||||
check "gun-sized total model = SBC + AD codes/thresholds",
|
||||
bb.memoryBytes > bb.sbcMemoryBytes
|
||||
# The SBC tensor is exactly nAde*nAde*nClasses bits, rounded up to uint32.
|
||||
check "SBC tensor bit count is exact",
|
||||
bb.sbcMemoryBytes * 8 == 6 * 512 * 512 * 8
|
||||
|
||||
# ── driver ───────────────────────────────────────────────────────────────────
|
||||
|
||||
testAdeScoring()
|
||||
testIdempotentLearning()
|
||||
testPlantedRuleAndOnlineCurve()
|
||||
testShuffledControl()
|
||||
testUnseenInput()
|
||||
testHomeostasis()
|
||||
testMemoryAccounting()
|
||||
|
||||
echo "\n", checks, " checks, ", failures, " failure(s)"
|
||||
if failures > 0:
|
||||
quit(1)
|
||||
echo "All bitbrain sanity checks passed."
|
||||
@@ -0,0 +1,179 @@
|
||||
## MNIST acceptance harness for the generic BitBrain library.
|
||||
##
|
||||
## This is the strongest correctness check available: it loads the *exact*
|
||||
## pretrained ADs and thresholds shipped with the reference C program, runs THIS
|
||||
## library's SBC learning and inference over the full MNIST train/test sets, and
|
||||
## reports top-1 accuracy against the two published reference numbers:
|
||||
##
|
||||
## * 96.540% — the reference C as shipped, which has a `uint8_t` truncation bug
|
||||
## in `read_from_sbc` (only ADEs with `i % 32 < 8` are counted),
|
||||
## * 97.210% — the reference C once that bug is fixed.
|
||||
##
|
||||
## A correct clean-room port must reproduce the CORRECTED number (~97.2%). The
|
||||
## harness also re-runs inference in "bug-compatible" mode (filtering row ADE
|
||||
## indices to those with `i % 32 < 8`) to demonstrate it reproduces the shipped
|
||||
## 96.540% too.
|
||||
##
|
||||
## No data is committed: fixtures are read from `/tmp` (or `$BITBRAIN_FIXTURES`).
|
||||
## Expected files (headerless, native-endian, row-major):
|
||||
## AD{1..4}_2048 int32[2048][width], widths = 6,8,10,12
|
||||
## thresh{1..4}_2048 int32[2048]
|
||||
## train_data uint8[60000][784] train_label uint8[60000]
|
||||
## test_data uint8[10000][784] test_label uint8[10000]
|
||||
##
|
||||
## Run:
|
||||
## nim c -r --nimcache:/tmp/nc_j90 -d:release \
|
||||
## --path:common_libs -o:/tmp/acceptance_bitbrain_mnist \
|
||||
## common_libs/tests/test_bitbrain_mnist.nim
|
||||
|
||||
import std/[os, times, strutils]
|
||||
import bitbrain/ade
|
||||
import bitbrain/sbc
|
||||
import bitbrain/bitbrain
|
||||
|
||||
const
|
||||
FixtureDefault = "/tmp/bitbrain/BitBrain_C_code"
|
||||
W = 2048 # ADEs per AD in the reference setup
|
||||
InputSz = 784
|
||||
TrainSz = 60000
|
||||
TestSz = 10000
|
||||
NClasses = 10
|
||||
Widths = [6, 8, 10, 12]
|
||||
|
||||
var failures = 0
|
||||
proc check(name: string, ok: bool) =
|
||||
if ok: echo "PASS: ", name
|
||||
else: echo "FAIL: ", name; inc failures
|
||||
|
||||
proc loadInt32(path: string, n: int): seq[int32] =
|
||||
let f = open(path, fmRead)
|
||||
defer: f.close()
|
||||
result = newSeq[int32](n)
|
||||
if n == 0: return
|
||||
let got = f.readBuffer(addr result[0], n * 4)
|
||||
doAssert got == n * 4, "short read from " & path
|
||||
|
||||
proc loadU8(path: string, n: int): seq[uint8] =
|
||||
let f = open(path, fmRead)
|
||||
defer: f.close()
|
||||
result = newSeq[uint8](n)
|
||||
if n == 0: return
|
||||
let got = f.readBuffer(addr result[0], n)
|
||||
doAssert got == n, "short read from " & path
|
||||
|
||||
proc buildFromFixtures(dir: string): BitBrain =
|
||||
var ades: seq[AddressDecoder]
|
||||
for k in 0 ..< Widths.len:
|
||||
var ad = initAddressDecoder(W, Widths[k])
|
||||
ad.codes = loadInt32(dir / ("AD" & $(k + 1) & "_2048"), W * Widths[k])
|
||||
ad.thresholds = loadInt32(dir / ("thresh" & $(k + 1) & "_2048"), W)
|
||||
ades.add ad
|
||||
# The reference C uses exactly the 6 cross-AD SBCs.
|
||||
result = initBitBrain(ades, crossPairs(ades.len), NClasses)
|
||||
|
||||
proc inferBuggy(bb: BitBrain, input: openArray[uint8]): seq[int] =
|
||||
## Emulate the reference C's `uint8_t bit_test` truncation: only row ADEs whose
|
||||
## index satisfies `i % 32 < 8` are ever counted at read time.
|
||||
var lists: seq[seq[int32]]
|
||||
bb.fireInto(input, lists)
|
||||
result = newSeq[int](bb.nClasses)
|
||||
for s in 0 ..< bb.sbcs.len:
|
||||
let spec = bb.specs[s]
|
||||
let sbc = bb.sbcs[s]
|
||||
for r in lists[spec.row]:
|
||||
if (int(r) and 31) >= 8: continue
|
||||
for c in lists[spec.col]:
|
||||
for k in 0 ..< bb.nClasses:
|
||||
if sbc.bitAt(int(r), int(c), k):
|
||||
inc result[k]
|
||||
|
||||
proc main() =
|
||||
let dir =
|
||||
if getEnv("BITBRAIN_FIXTURES").len > 0: getEnv("BITBRAIN_FIXTURES")
|
||||
else: FixtureDefault
|
||||
if not dirExists(dir):
|
||||
echo "Skipping: BitBrain fixtures not found at ", dir
|
||||
echo " (set BITBRAIN_FIXTURES to override)"
|
||||
return
|
||||
|
||||
echo "Fixtures: ", dir
|
||||
echo "Building BitBrain from pretrained ADs (widths 6,8,10,12) + 6 SBCs ..."
|
||||
var bb = buildFromFixtures(dir)
|
||||
|
||||
echo "ADs: ", bb.nAdes, " SBCs: ", bb.sbcs.len,
|
||||
" classes: ", bb.nClasses
|
||||
echo "Model bytes: ", bb.memoryBytes,
|
||||
" (SBC tensors: ", bb.sbcMemoryBytes, ")"
|
||||
|
||||
echo "Loading MNIST ..."
|
||||
let trainData = loadU8(dir / "train_data", TrainSz * InputSz)
|
||||
let trainLabel = loadU8(dir / "train_label", TrainSz)
|
||||
let testData = loadU8(dir / "test_data", TestSz * InputSz)
|
||||
let testLabel = loadU8(dir / "test_label", TestSz)
|
||||
|
||||
# ── training (one online single pass) ──────────────────────────────────────
|
||||
echo "Training on ", TrainSz, " samples (single online pass) ..."
|
||||
let tTrain0 = cpuTime()
|
||||
for i in 0 ..< TrainSz:
|
||||
bb.learn(toOpenArray(trainData, i * InputSz, i * InputSz + InputSz - 1),
|
||||
int(trainLabel[i]))
|
||||
let trainSec = cpuTime() - tTrain0
|
||||
echo " train: ", formatFloat(trainSec, ffDecimal, 3), " s total, ",
|
||||
formatFloat(trainSec * 1000.0 / float(TrainSz), ffDecimal, 4),
|
||||
" ms/sample"
|
||||
|
||||
# ── inference ──────────────────────────────────────────────────────────────
|
||||
echo "Inferring on ", TestSz, " samples ..."
|
||||
var correct = 0
|
||||
var buggyCorrect = 0
|
||||
var correctCountsNonzero = 0
|
||||
let tInfer0 = cpuTime()
|
||||
for i in 0 ..< TestSz:
|
||||
let r = bb.infer(toOpenArray(testData, i * InputSz,
|
||||
i * InputSz + InputSz - 1))
|
||||
if r.label == int(testLabel[i]): inc correct
|
||||
var tot = 0
|
||||
for c in r.counts: tot += c
|
||||
if tot > 0: inc correctCountsNonzero
|
||||
let inferSec = cpuTime() - tInfer0
|
||||
echo " infer: ", formatFloat(inferSec, ffDecimal, 3), " s total, ",
|
||||
formatFloat(inferSec * 1000.0 / float(TestSz), ffDecimal, 4),
|
||||
" ms/sample"
|
||||
|
||||
let tBuggy0 = cpuTime()
|
||||
for i in 0 ..< TestSz:
|
||||
let counts = inferBuggy(bb, toOpenArray(testData, i * InputSz,
|
||||
i * InputSz + InputSz - 1))
|
||||
var label = 0
|
||||
for k in 1 ..< counts.len:
|
||||
if counts[k] > counts[label]: label = k
|
||||
if label == int(testLabel[i]): inc buggyCorrect
|
||||
let buggySec = cpuTime() - tBuggy0
|
||||
|
||||
let acc = 100.0 * float(correct) / float(TestSz)
|
||||
let buggyAcc = 100.0 * float(buggyCorrect) / float(TestSz)
|
||||
|
||||
echo ""
|
||||
echo "============================================================"
|
||||
echo " Corrected reader (clean-room): ", formatFloat(acc, ffDecimal, 3),
|
||||
" pct (reference 97.210)"
|
||||
echo " Bug-compat reader: ",
|
||||
formatFloat(buggyAcc, ffDecimal, 3), " pct (reference 96.540)"
|
||||
echo " empty-count test samples: ", TestSz - correctCountsNonzero
|
||||
echo " bug-compat infer time: ",
|
||||
formatFloat(buggySec * 1000.0 / float(TestSz), ffDecimal, 4), " ms/sample"
|
||||
echo "============================================================"
|
||||
|
||||
check "corrected accuracy reproduces the reference (~97.2, >= 97.0)",
|
||||
acc >= 97.0
|
||||
check "bug-compat accuracy reproduces the shipped reference (~96.5)",
|
||||
abs(buggyAcc - 96.540) < 0.5
|
||||
check "corrected reader strictly improves on the buggy one",
|
||||
acc > buggyAcc
|
||||
|
||||
if failures > 0:
|
||||
echo "\n", failures, " acceptance check(s) FAILED"
|
||||
quit(1)
|
||||
echo "\nMNIST acceptance checks passed."
|
||||
|
||||
main()
|
||||
Reference in New Issue
Block a user