77e6dace01
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.
180 lines
6.9 KiB
Nim
180 lines
6.9 KiB
Nim
## 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()
|