BitBrain SBC: counted mode + global decay (forgetting, probabilities)

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
This commit is contained in:
2026-09-25 08:39:10 +02:00
parent 39e06719fb
commit 40ba96f649
6 changed files with 921 additions and 53 deletions
+101 -1
View File
@@ -13,7 +13,7 @@
## * correct, non-crashing behaviour on an unseen input,
## * homeostatic threshold adaptation moves the firing rate toward target.
import std/[random, math, strutils]
import std/[random, math, strutils, os]
import bitbrain/ade
import bitbrain/sbc
import bitbrain/bitbrain
@@ -237,6 +237,102 @@ proc testMemoryAccounting() =
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()
@@ -246,6 +342,10 @@ testShuffledControl()
testUnseenInput()
testHomeostasis()
testMemoryAccounting()
testCountedMemory()
testCountedForgetting()
testCountedProbability()
testCountedEnv()
echo "\n", checks, " checks, ", failures, " failure(s)"
if failures > 0: