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
+82 -11
View File
@@ -20,8 +20,14 @@
## 10-SBC variant adds 4 within-AD ("half-size") SBCs; `withinPairs` builds those
## (this implementation stores them full-size — the half-size packing is a
## separate memory optimisation).
##
## Storage modes: the container can build every SBC in the default `smBitset`
## mode or in `smCounted` mode (saturating counters + optional global decay).
## `infer` returns integer per-class evidence (set-bit count, or summed
## counters); `inferProb` returns the per-cell posterior sum, which is the
## scale-free probability readout. Both are additive across SBCs.
import std/random
import std/[random, os, strutils]
import ade, sbc
export ade, sbc
@@ -50,9 +56,15 @@ proc withinPairs*(nAdes: int): seq[SbcSpec] =
result.add SbcSpec(row: a, col: a)
proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec],
nClasses: int): BitBrain =
nClasses: int,
mode: SbcMode = smBitset,
decayEvery: int = DefaultDecayEvery,
decayShift: int = DefaultDecayShift): BitBrain =
## Build the container. Every AD must have the same number of ADEs because an
## SBC's axes are both `w` long (the paper's setup).
##
## `mode`/`decayEvery`/`decayShift` select the SBC storage: the defaults are
## the reference-compatible bitset SBC (decay knobs ignored).
doAssert ades.len > 0, "need at least one AD"
doAssert nClasses > 0, "need at least one class"
let w = ades[0].nAde
@@ -65,11 +77,17 @@ proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec],
for s in 0 ..< specs.len:
doAssert specs[s].row >= 0 and specs[s].row < ades.len
doAssert specs[s].col >= 0 and specs[s].col < ades.len
result.sbcs[s] = initSbc(w, nClasses)
if mode == smCounted:
result.sbcs[s] = initCountedSbc(w, nClasses, decayEvery, decayShift)
else:
result.sbcs[s] = initSbc(w, nClasses)
proc buildRandomBitBrain*(widths: seq[int], nAde, inputWidth, nClasses: int,
seed: int64,
specs: seq[SbcSpec] = @[]): BitBrain =
specs: seq[SbcSpec] = @[],
mode: SbcMode = smBitset,
decayEvery: int = DefaultDecayEvery,
decayShift: int = DefaultDecayShift): BitBrain =
## Initialise a whole BitBrain with random ADs. If `specs` is empty the
## paper's 6 cross-AD SBCs are used. Fully deterministic given `seed`.
var rng = initRand(seed)
@@ -77,7 +95,35 @@ proc buildRandomBitBrain*(widths: seq[int], nAde, inputWidth, nClasses: int,
for w in widths:
ades.add initRandomAddressDecoder(nAde, w, inputWidth, rng)
let pairs = if specs.len == 0: crossPairs(widths.len) else: specs
result = initBitBrain(ades, pairs, nClasses)
result = initBitBrain(ades, pairs, nClasses, mode, decayEvery, decayShift)
# ── runtime configuration (env; compile-time defaults are the constants above) ─
const
BB_MODE_ENV* = "TR_BITBRAIN_MODE"
## `bitset` (default) | `counted`.
BB_DECAY_EVERY_ENV* = "TR_BITBRAIN_DECAY_EVERY"
## learns between global decay passes (counted mode only).
BB_DECAY_SHIFT_ENV* = "TR_BITBRAIN_DECAY_SHIFT"
## fractional decay strength (counted mode only); `0` disables decay.
proc envSbcMode*(default = smBitset): SbcMode =
## Runtime mode knob. Unknown/empty values fall back to `default` (the shipped
## default is the unchanged bitset path).
case getEnv(BB_MODE_ENV, "").strip().toLowerAscii()
of "counted", "counter", "counters", "count": smCounted
of "bitset", "bit", "bits": smBitset
else: default
proc envDecayEvery*(default = DefaultDecayEvery): int =
let v = getEnv(BB_DECAY_EVERY_ENV, "").strip()
if v.len == 0: return default
try: parseInt(v) except ValueError: default
proc envDecayShift*(default = DefaultDecayShift): int =
let v = getEnv(BB_DECAY_SHIFT_ENV, "").strip()
if v.len == 0: return default
try: parseInt(v) except ValueError: default
proc nAdes*(bb: BitBrain): int {.inline.} =
bb.ades.len
@@ -98,7 +144,8 @@ proc fireInto*[T: SomeInteger](bb: BitBrain, input: openArray[T],
proc learn*[T: SomeInteger](bb: var BitBrain, input: openArray[T], class: int) =
## One online supervised step: set the `class` bit of every observed
## coincidence. Idempotent — repeating the same sample is a no-op.
## coincidence (bitset mode, idempotent) or increment/saturate the `class`
## counter of every observed coincidence (counted mode).
var lists: seq[seq[int32]]
bb.fireInto(input, lists)
for s in 0 ..< bb.sbcs.len:
@@ -107,9 +154,10 @@ proc learn*[T: SomeInteger](bb: var BitBrain, input: openArray[T], class: int) =
proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
tuple[label: int, counts: seq[int]] =
## Drive every AD, count set class bits in every SBC, and return the argmax
## class plus the aggregated per-class counts. Ties go to the lowest class
## index (matching the reference's `>` scan that keeps the first maximum).
## Drive every AD, accumulate per-class integer evidence in every SBC (set-bit
## count, or summed counters), and return the argmax class plus the aggregated
## per-class counts. Ties go to the lowest class index (matching the reference's
## `>` scan that keeps the first maximum).
var lists: seq[seq[int32]]
bb.fireInto(input, lists)
result.counts = newSeq[int](bb.nClasses)
@@ -123,14 +171,37 @@ proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
best = result.counts[k]
result.label = k
proc inferProb*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
tuple[label: int, scores: seq[float]] =
## Drive every AD and accumulate the scale-free probability readout: each SBC
## adds `P(class | coincidence cell)` per observed coincidence. The argmax is
## the label. This is the readout that a rare-but-predictable class needs (see
## `inferProb` on the SBC); `infer` is the plain vote/count readout.
var lists: seq[seq[int32]]
bb.fireInto(input, lists)
result.scores = newSeq[float](bb.nClasses)
for s in 0 ..< bb.sbcs.len:
let spec = bb.specs[s]
bb.sbcs[s].inferProb(lists[spec.row], lists[spec.col], result.scores)
result.label = 0
var best = result.scores[0]
for k in 1 ..< bb.nClasses:
if result.scores[k] > best:
best = result.scores[k]
result.label = k
proc memoryBytes*(bb: BitBrain): int =
## Total bytes: AD synapse codes + thresholds + counters + SBC bit tensors.
## Total bytes: AD synapse codes + thresholds + counters + SBC tensors.
for ad in bb.ades:
result += ad.memoryBytes
for s in bb.sbcs:
result += s.memoryBytes
proc sbcMemoryBytes*(bb: BitBrain): int =
## Just the SBC bit tensors — the dominant term for realistic configurations.
## Just the SBC tensors — the dominant term for realistic configurations.
for s in bb.sbcs:
result += s.memoryBytes
proc sbcMode*(bb: BitBrain): SbcMode =
## Storage mode of this container (all SBCs share one mode).
if bb.sbcs.len == 0: smBitset else: bb.sbcs[0].mode