## 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, os] 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 # Counted mode is one uint8 per (i,j,class): 8x the bit tensor at gun size. var rng2 = initRand(1) var ades2: seq[AddressDecoder] for w in [6, 8, 10, 12]: ades2.add initRandomAddressDecoder(512, w, 256, rng2) let cbb = initBitBrain(ades2, crossPairs(4), 8, smCounted, 0, 0) check "counted mode is selected in the container", cbb.sbcMode == smCounted check "gun-sized counted SBC memory = 12,582,912 bytes", cbb.sbcMemoryBytes == 12_582_912 check "counted gun memory is exactly 8x the bitset tensor", cbb.sbcMemoryBytes == 8 * bb.sbcMemoryBytes # ── 8. Counted mode: learning, saturation, decay, readouts ──────────────────── proc testCountedMemory() = # Explicit no-decay so this test is independent of the compile-time default. var s = initCountedSbc(64, 2, decayEvery = 0, decayShift = 0) check "counted sbc reports counted mode", s.mode == smCounted check "counted sbc allocates counters, not bits", s.counters.len == 64 * 64 * 2 and s.bits.len == 0 check "counted bytes = cells (1 byte each)", s.memoryBytes == 64 * 64 * 2 check "counted first learn touches one cell", s.learn(@[1'i32], @[2'i32], 0) == 1 check "counted counter increments", s.countAt(1, 2, 0) == 1 # The bitset is idempotent; the counter is not: repetition becomes evidence. for _ in 0 ..< 9: discard s.learn(@[1'i32], @[2'i32], 0) check "counted learning the same sample 10x gives count 10", s.countAt(1, 2, 0) == 10 # Saturation at the uint8 ceiling. for _ in 0 ..< 300: discard s.learn(@[1'i32], @[3'i32], 1) check "counted counters saturate at 255", s.countAt(1, 3, 1) == 255 var raw = newSeq[int](2) s.infer(@[1'i32], @[2'i32, 3'i32], raw) check "counted infer sums raw counters", raw[0] == 10 and raw[1] == 255 s.clear() check "counted clear wipes counters", s.countAt(1, 2, 0) == 0 proc testCountedForgetting() = # No decay: the stale class keeps the majority after a regime change. var nd = initCountedSbc(64, 2, decayEvery = 0, decayShift = 0) for _ in 0 ..< 50: discard nd.learn(@[0'i32], @[0'i32], 0) for _ in 0 ..< 20: discard nd.learn(@[0'i32], @[0'i32], 1) var ct = newSeq[int](2) nd.infer(@[0'i32], @[0'i32], ct) check "no-decay counters keep the stale class (" & $ct[0] & " vs " & $ct[1] & ")", ct[0] > ct[1] # Decay: the recent class wins even though it was seen fewer times. var dd = initCountedSbc(64, 2, decayEvery = 10, decayShift = 1) for _ in 0 ..< 50: discard dd.learn(@[0'i32], @[0'i32], 0) for _ in 0 ..< 20: discard dd.learn(@[0'i32], @[0'i32], 1) var cd = newSeq[int](2) dd.infer(@[0'i32], @[0'i32], cd) check "decay lets the recent class win (" & $cd[0] & " vs " & $cd[1] & ")", cd[1] > cd[0] check "decay keeps counters bounded below the ceiling", cd[0] < 255 and cd[1] < 255 proc testCountedProbability() = # One mixed cell (class0 50x, class1 5x) and three pure class1 cells (5x each). # The raw counter sum favours the common class; the per-cell posterior does not. var s = initCountedSbc(8, 2, decayEvery = 0, decayShift = 0) for _ in 0 ..< 50: discard s.learn(@[0'i32], @[0'i32], 0) for _ in 0 ..< 5: discard s.learn(@[0'i32], @[0'i32], 1) for c in 1 .. 3: for _ in 0 ..< 5: discard s.learn(@[0'i32], @[c.int32], 1) let cols = @[0'i32, 1'i32, 2'i32, 3'i32] var raw = newSeq[int](2) s.infer(@[0'i32], cols, raw) check "raw sum is dominated by the 50-count mixed cell (" & $raw[0] & " vs " & $raw[1] & ")", raw[0] > raw[1] var post = newSeq[float](2) s.inferProb(@[0'i32], cols, post) check "per-cell posterior favours the predictable class (" & formatFloat(post[0], ffDecimal, 3) & " vs " & formatFloat(post[1], ffDecimal, 3) & ")", post[1] > post[0] proc testCountedEnv() = putEnv("TR_BITBRAIN_MODE", "counted") check "env mode parses counted", envSbcMode() == smCounted putEnv("TR_BITBRAIN_MODE", "bitset") check "env mode parses bitset", envSbcMode() == smBitset putEnv("TR_BITBRAIN_MODE", "banana") check "unknown env mode falls back to the default", envSbcMode() == smBitset putEnv("TR_BITBRAIN_DECAY_SHIFT", "5") check "env decay shift parses", envDecayShift() == 5 putEnv("TR_BITBRAIN_DECAY_EVERY", "77") check "env decay interval parses", envDecayEvery() == 77 delEnv("TR_BITBRAIN_MODE") delEnv("TR_BITBRAIN_DECAY_SHIFT") delEnv("TR_BITBRAIN_DECAY_EVERY") check "unset env decay shift falls back to the compile-time default", envDecayShift() == DefaultDecayShift check "unset env decay interval falls back to the compile-time default", envDecayEvery() == DefaultDecayEvery # ── driver ─────────────────────────────────────────────────────────────────── testAdeScoring() testIdempotentLearning() testPlantedRuleAndOnlineCurve() testShuffledControl() testUnseenInput() testHomeostasis() testMemoryAccounting() testCountedMemory() testCountedForgetting() testCountedProbability() testCountedEnv() echo "\n", checks, " checks, ", failures, " failure(s)" if failures > 0: quit(1) echo "All bitbrain sanity checks passed."