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:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user