40ba96f649
Adds an smCounted storage mode alongside the default smBitset. Each (i,j,class) cell becomes a saturating uint8 counter; learn increments it and a global fractional decay (c -= c shr decayShift every decayEvery learns) makes forgetting possible. infer sums raw counters; new inferProb sums the per-cell posterior P(class|cell) (scale-free, recommended readout). Bitset path is the default and byte-for-byte unchanged: test_bitbrain 56/56 (was 32), and test_bitbrain_mnist reproduces 97.210% corrected / 96.540% bug-compatible exactly. Counted mode configurable at runtime (TR_BITBRAIN_MODE / TR_BITBRAIN_DECAY_*) and compile time (-d:bitbrainDecay*). Measured: forgetting (86.2% vs 48.9% on a permuted-label stream), probabilities (rare-class balanced 0.998 vs 0.500), and the stationary cost (counted hurts MNIST; see docs/bitbrain_counted_sbc.md). Harness: common_libs/tests/measure_counted_sbc.nim
354 lines
15 KiB
Nim
354 lines
15 KiB
Nim
## 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."
|