Files
SirRoboGarage/common_libs/bitbrain/bitbrain.nim
T
SirStone 40ba96f649 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
2026-09-25 08:39:10 +02:00

208 lines
8.4 KiB
Nim

## BitBrain container — ADs + SBC memories + a counting readout.
##
## Clean-room implementation of the classification pipeline described in
##
## "BitBrain and Sparse Binary Coincidence (SBC) memories",
## Frontiers in Neuroinformatics 17:1125844, 2023.
##
## See `ade.nim` for the clean-room note (the reference C is GPL-3.0 and was not
## copied).
##
## The container holds several Address Decoders (possibly of different widths)
## and several SBC memories, each built from a pair of ADs. `learn` populates the
## SBCs from a labelled sample; `infer` drives every AD, reads every SBC with the
## same coincidence rule and sums the per-class set-bit counts. The class with the
## highest total wins. Everything is per-sample and order-free: the memory is
## monotone, so `learn` and `infer` may be interleaved arbitrarily.
##
## Reference setup used by the paper and the MNIST acceptance harness: 4 ADs of
## 2048 ADEs with widths {6, 8, 10, 12} and 6 cross-AD SBCs. The paper's
## 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, os, strutils]
import ade, sbc
export ade, sbc
type
SbcSpec* = object
## Which pair of ADs feeds one SBC memory.
row*: int
col*: int
BitBrain* = object
ades*: seq[AddressDecoder]
sbcs*: seq[Sbc]
specs*: seq[SbcSpec]
nClasses*: int
proc crossPairs*(nAdes: int): seq[SbcSpec] =
## All unordered pairs of distinct ADs. For 4 ADs this is the paper's 6 SBCs.
for a in 0 ..< nAdes:
for b in a + 1 ..< nAdes:
result.add SbcSpec(row: a, col: b)
proc withinPairs*(nAdes: int): seq[SbcSpec] =
## The square/within-AD pairs (the paper's "half-size" SBCs).
for a in 0 ..< nAdes:
result.add SbcSpec(row: a, col: a)
proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec],
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
for ad in ades:
doAssert ad.nAde == w, "all ADs must have the same number of ADEs"
result.ades = ades
result.nClasses = nClasses
result.specs = specs
result.sbcs = newSeq[Sbc](specs.len)
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
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] = @[],
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)
var ades: seq[AddressDecoder]
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, 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
proc resetLearning*(bb: var BitBrain) =
## Wipe every SBC. The ADs (and their thresholds) are left untouched.
for s in 0 ..< bb.sbcs.len:
bb.sbcs[s].clear()
proc fireInto*[T: SomeInteger](bb: BitBrain, input: openArray[T],
lists: var seq[seq[int32]]) =
## Compute the sparse firing pattern of every AD for `input`. `lists` is resized
## to `nAdes` and each entry is filled with that AD's firing ADE indices.
if lists.len != bb.ades.len:
lists.setLen(bb.ades.len)
for a in 0 ..< bb.ades.len:
bb.ades[a].activeList(input, lists[a])
proc learn*[T: SomeInteger](bb: var BitBrain, input: openArray[T], class: int) =
## One online supervised step: set the `class` bit of every observed
## 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:
let spec = bb.specs[s]
discard bb.sbcs[s].learn(lists[spec.row], lists[spec.col], class)
proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
tuple[label: int, counts: seq[int]] =
## 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)
for s in 0 ..< bb.sbcs.len:
let spec = bb.specs[s]
bb.sbcs[s].infer(lists[spec.row], lists[spec.col], result.counts)
result.label = 0
var best = result.counts[0]
for k in 1 ..< bb.nClasses:
if result.counts[k] > best:
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 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 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