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:
2026-09-24 20:52:45 +02:00
parent cc11ede824
commit 77e6dace01
6 changed files with 1040 additions and 0 deletions
+253
View File
@@ -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."
+179
View File
@@ -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()