## 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()