From 40ba96f64974bd8085b227f3ba8a2712c88c356d Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Fri, 25 Sep 2026 08:39:10 +0200 Subject: [PATCH] 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 --- common_libs/bitbrain/README.md | 14 + common_libs/bitbrain/bitbrain.nim | 93 ++++++- common_libs/bitbrain/sbc.nim | 251 +++++++++++++++--- common_libs/tests/measure_counted_sbc.nim | 306 ++++++++++++++++++++++ common_libs/tests/test_bitbrain.nim | 102 +++++++- docs/bitbrain_counted_sbc.md | 208 +++++++++++++++ 6 files changed, 921 insertions(+), 53 deletions(-) create mode 100644 common_libs/tests/measure_counted_sbc.nim create mode 100644 docs/bitbrain_counted_sbc.md diff --git a/common_libs/bitbrain/README.md b/common_libs/bitbrain/README.md index a3e76bf..32d50ee 100644 --- a/common_libs/bitbrain/README.md +++ b/common_libs/bitbrain/README.md @@ -86,6 +86,20 @@ Generic over the input element type (`openArray[SomeInteger]`) and over the inpu width, ADE count and class count — nothing is hardcoded to 784/10. Deterministic given the seed. Dependencies: `std/` only. +### Counted SBC mode (saturating counters + forgetting) + +An optional `smCounted` mode replaces each bit with a saturating `uint8` counter +and adds global fractional decay (`c -= c shr decayShift` every `decayEvery` +learns). `initBitBrain(..., mode, decayEvery, decayShift)` selects it; +`initCountedSbc` builds a single counted memory. `infer` sums raw counters; +`inferProb` sums the per-cell posterior `P(class | cell)` (the recommended +counted readout). Runtime knobs: `TR_BITBRAIN_MODE` (`bitset` default, +`counted`), `TR_BITBRAIN_DECAY_EVERY`, `TR_BITBRAIN_DECAY_SHIFT`; compile-time +defaults: `-d:bitbrainDecayEvery=N`, `-d:bitbrainDecayShift=N`. The default +bitset path is unchanged and remains the reference-compatible one. Design, +memory cost and the measured forgetting/probability/stationary evidence are in +[`docs/bitbrain_counted_sbc.md`](../../docs/bitbrain_counted_sbc.md). + Reference SBC wiring: `crossPairs(4)` gives the 6 cross-AD SBCs used by the reference C. `withinPairs(n)` adds the paper's 4 within-AD ("half-size") SBCs; this implementation stores them full-size (half-size packing is a separate memory diff --git a/common_libs/bitbrain/bitbrain.nim b/common_libs/bitbrain/bitbrain.nim index 0289020..53a479c 100644 --- a/common_libs/bitbrain/bitbrain.nim +++ b/common_libs/bitbrain/bitbrain.nim @@ -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 diff --git a/common_libs/bitbrain/sbc.nim b/common_libs/bitbrain/sbc.nim index 3ba84ae..1ea2853 100644 --- a/common_libs/bitbrain/sbc.nim +++ b/common_libs/bitbrain/sbc.nim @@ -16,15 +16,18 @@ ## memory cell. That cell holds a class bitmask with one bit per class ## (`nClasses` bits, one-hot encoding in the paper's default). ## -## Learning is **idempotent**: `learn` *sets* the bit for the observed class; -## setting it again is a no-op. There is no clearing, no learning rate, no decay -## and no epoch — one pass through the data is a complete supervised training run, -## and a second pass changes nothing. +## In the default **bitset** mode learning is **idempotent**: `learn` *sets* the +## bit for the observed class; setting it again is a no-op. There is no clearing, +## no learning rate, no decay and no epoch — one pass through the data is a +## complete supervised training run, and a second pass changes nothing. ## ## Inference uses the *same* address decoding: for each observed coincidence every ## class bit is read and the set bits are **counted** per class. Counts are summed ## across SBCs by the `bitbrain` container and the argmax wins. ## +## A second **counted** mode replaces the bit with a saturating counter and adds +## forgetting; see below. +## ## Bit layout ## ---------- ## The bit for `(i, j, class)` lives at `((i * nAde) + j) * nClasses + class`, so @@ -32,83 +35,249 @@ ## reference C's layout but is a bijection onto the same set of triples; the ## learned rule is identical. +## Two storage modes are available: +## +## * `smBitset` (DEFAULT) — the reference-compatible idempotent set-bit memory +## described above. Behaviour is byte-for-byte unchanged. +## * `smCounted` — each cell stores a small saturating `uint8` counter per class +## instead of one bit. `learn` increments the observed class's counter and a +## global fractional decay (`counter -= counter shr decayShift`, run every +## `decayEvery` learns) ages every counter on a schedule. Inference sums the +## counters per class (a frequency estimate) or, with `inferProb`, sums the +## per-cell posterior `count[k] / cellTotal` (a probability estimate). +## +## Counted mode fixes the two defects of the bit memory: a bit records *that* a +## class co-occurred, never *how often* (so the object is a probability); and a +## bit cannot be cleared (so stale associations saturate). The decay makes the +## memory bounded and recency-weighted. + import std/bitops type + SbcMode* = enum + smBitset ## default: idempotent set-bit memory (reference-compatible) + smCounted ## saturating per-cell class counters with global fractional decay + Sbc* = object - ## A 2-D coincidence memory with a class-bit depth. + ## A 2-D coincidence memory with a class depth. nAde*: int ## number of ADEs on each axis (the paper's `w`) nClasses*: int ## number of classes (the paper's `D`) - bits*: seq[uint32] ## packed bit tensor, nAde*nAde*nClasses bits + mode*: SbcMode ## which storage/readout this memory uses + bits*: seq[uint32] ## smBitset: packed bit tensor, nAde*nAde*nClasses bits + counters*: seq[uint8] ## smCounted: one saturating counter per (i,j,class) + decayEvery*: int ## smCounted: learns between global decay passes (0 = never) + decayShift*: int ## smCounted: `c -= c shr decayShift` per decay pass + learnCount*: int ## smCounted: learns since the last decay pass + +const + DefaultDecayEvery* {.intdefine: "bitbrainDecayEvery".} = 1024 + ## learns between global decay passes. Compile-time default; override with + ## `-d:bitbrainDecayEvery=N` or at runtime with `TR_BITBRAIN_DECAY_EVERY`. + DefaultDecayShift* {.intdefine: "bitbrainDecayShift".} = 1 + ## fractional decay strength: `c -= c shr shift`. `1` is a halving; higher + ## values forget more slowly. `0` disables decay. Compile-time default; + ## override with `-d:bitbrainDecayShift=N` or `TR_BITBRAIN_DECAY_SHIFT`. proc initSbc*(nAde, nClasses: int): Sbc = - ## Allocate a zeroed SBC (nothing is known yet). + ## Allocate a zeroed bitset SBC (nothing is known yet). This is the default + ## reference-compatible mode and is unchanged. doAssert nAde > 0, "nAde must be positive" doAssert nClasses > 0, "nClasses must be positive" result.nAde = nAde result.nClasses = nClasses + result.mode = smBitset let nbits = nAde * nAde * nClasses result.bits = newSeq[uint32]((nbits + 31) div 32) +proc initCountedSbc*(nAde, nClasses: int, + decayEvery = DefaultDecayEvery, + decayShift = DefaultDecayShift): Sbc = + ## Allocate a zeroed counted SBC: one saturating `uint8` per `(i, j, class)`. + ## `decayEvery` is the number of learns between global decay passes (0 = no + ## forgetting, i.e. a pure saturating counter) and `decayShift` the fractional + ## decay strength (`c -= c shr decayShift`). + doAssert nAde > 0, "nAde must be positive" + doAssert nClasses > 0, "nClasses must be positive" + doAssert decayEvery >= 0, "decayEvery must be >= 0" + doAssert decayShift >= 0, "decayShift must be >= 0" + result.nAde = nAde + result.nClasses = nClasses + result.mode = smCounted + result.counters = newSeq[uint8](nAde * nAde * nClasses) + result.decayEvery = decayEvery + result.decayShift = decayShift + proc clear*(sbc: var Sbc) = ## Forget everything. This is how a monotone memory is wiped (e.g. when a bot ## switches enemy and must not carry state across battles). for i in 0 ..< sbc.bits.len: sbc.bits[i] = 0'u32 + for i in 0 ..< sbc.counters.len: + sbc.counters[i] = 0'u8 + sbc.learnCount = 0 + +proc applyDecay*(sbc: var Sbc) = + ## One global forgetting pass over every counter: `c -= c shr decayShift`. + ## No-op in bitset mode or when decay is disabled. Amortised cost is + ## O(nAde*nAde*nClasses / decayEvery) per learn, so a per-tick learner only + ## pays the whole pass once every `decayEvery` learns. Unlike a per-cell EMA, + ## it also ages cells that are never visited again (true forgetting). + if sbc.mode != smCounted or sbc.decayShift <= 0: return + let sh = sbc.decayShift + for i in 0 ..< sbc.counters.len: + sbc.counters[i] = sbc.counters[i] - (sbc.counters[i] shr sh) proc bitIndex(sbc: Sbc, i, j, class: int): int {.inline.} = ((i * sbc.nAde) + j) * sbc.nClasses + class proc bitAt*(sbc: Sbc, i, j, class: int): bool {.inline.} = - ## Read one memory bit. Exposed mainly so harnesses can inspect the exact rule. + ## Read one memory entry as a boolean: a set bit in bitset mode, a non-zero + ## counter in counted mode. Exposed mainly so harnesses can inspect the rule. doAssert i >= 0 and i < sbc.nAde doAssert j >= 0 and j < sbc.nAde doAssert class >= 0 and class < sbc.nClasses - let bit = bitIndex(sbc, i, j, class) - (sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32 + let idx = bitIndex(sbc, i, j, class) + case sbc.mode + of smBitset: + (sbc.bits[idx shr 5] and (1'u32 shl (idx and 31))) != 0'u32 + of smCounted: + sbc.counters[idx] != 0'u8 + +proc countAt*(sbc: Sbc, i, j, class: int): int {.inline.} = + ## Read the raw evidence at one cell: `0/1` for a bit, `0..255` for a counter. + doAssert i >= 0 and i < sbc.nAde + doAssert j >= 0 and j < sbc.nAde + doAssert class >= 0 and class < sbc.nClasses + let idx = bitIndex(sbc, i, j, class) + case sbc.mode + of smBitset: + if (sbc.bits[idx shr 5] and (1'u32 shl (idx and 31))) != 0'u32: 1 else: 0 + of smCounted: + int(sbc.counters[idx]) proc learn*(sbc: var Sbc, rowActive, colActive: openArray[int32], class: int): int = - ## Set the `class` bit for every coincidence between a firing row ADE and a - ## firing column ADE. Returns the number of bits that were newly set (0 if the - ## sample added no information, e.g. it was already learned). + ## Bitset mode: set the `class` bit for every coincidence between a firing row + ## ADE and a firing column ADE, returning the number of bits newly set (0 if + ## the sample added no information). + ## + ## Counted mode: increment the `class` counter of every coincidence (saturating + ## at 255) and, once every `decayEvery` learns, apply the global decay. Returns + ## the number of coincidence cells touched. doAssert class >= 0 and class < sbc.nClasses, "class out of range" let D = sbc.nClasses - for r in rowActive: - let i = int(r) - for c in colActive: - let j = int(c) - let bit = ((i * sbc.nAde) + j) * D + class - let w = bit shr 5 - let m = 1'u32 shl (bit and 31) - if (sbc.bits[w] and m) == 0'u32: - sbc.bits[w] = sbc.bits[w] or m + case sbc.mode + of smBitset: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let bit = ((i * sbc.nAde) + j) * D + class + let w = bit shr 5 + let m = 1'u32 shl (bit and 31) + if (sbc.bits[w] and m) == 0'u32: + sbc.bits[w] = sbc.bits[w] or m + inc result + of smCounted: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let idx = ((i * sbc.nAde) + j) * D + class + if sbc.counters[idx] < 255'u8: + inc sbc.counters[idx] inc result + inc sbc.learnCount + if sbc.decayEvery > 0 and sbc.learnCount >= sbc.decayEvery: + sbc.applyDecay() + sbc.learnCount = 0 proc infer*(sbc: Sbc, rowActive, colActive: openArray[int32], counts: var seq[int]) = - ## Count, per class, how many observed coincidences have their class bit set. - ## `counts` is *accumulated into* (not reset), so a container can sum several - ## SBCs. It must be at least `nClasses` long. + ## Bitset mode: count, per class, how many observed coincidences have their + ## class bit set. Counted mode: sum the `class` counters over the observed + ## coincidences (a frequency estimate). `counts` is *accumulated into* (not + ## reset), so a container can sum several SBCs. It must be at least `nClasses`. doAssert counts.len >= sbc.nClasses, "counts buffer too small" let D = sbc.nClasses - for r in rowActive: - let i = int(r) - for c in colActive: - let j = int(c) - let base = ((i * sbc.nAde) + j) * D - for k in 0 ..< D: - let bit = base + k - if (sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32: - inc counts[k] + case sbc.mode + of smBitset: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let base = ((i * sbc.nAde) + j) * D + for k in 0 ..< D: + let bit = base + k + if (sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32: + inc counts[k] + of smCounted: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let base = ((i * sbc.nAde) + j) * D + for k in 0 ..< D: + counts[k] += int(sbc.counters[base + k]) + +proc inferProb*(sbc: Sbc, rowActive, colActive: openArray[int32], + scores: var seq[float]) = + ## Probability readout. For every observed coincidence, form the per-cell + ## posterior `P(class | cell) = count[class] / Σ_k count[k]` (in bitset mode, + ## the uniform posterior over the set bits) and **sum it per class**. This is + ## scale-free in the class marginals: a rare class whose cells are almost + ## always co-labelled with it wins over a common class that merely touches more + ## cells. `scores` is accumulated into and must be at least `nClasses` long. + doAssert scores.len >= sbc.nClasses, "scores buffer too small" + let D = sbc.nClasses + case sbc.mode + of smBitset: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let base = ((i * sbc.nAde) + j) * D + var tot = 0 + for k in 0 ..< D: + if (sbc.bits[(base + k) shr 5] and + (1'u32 shl ((base + k) and 31))) != 0'u32: + inc tot + if tot > 0: + let p = 1.0 / float(tot) + for k in 0 ..< D: + if (sbc.bits[(base + k) shr 5] and + (1'u32 shl ((base + k) and 31))) != 0'u32: + scores[k] += p + of smCounted: + for r in rowActive: + let i = int(r) + for c in colActive: + let j = int(c) + let base = ((i * sbc.nAde) + j) * D + var tot = 0 + for k in 0 ..< D: + tot += int(sbc.counters[base + k]) + if tot > 0: + for k in 0 ..< D: + scores[k] += float(sbc.counters[base + k]) / float(tot) proc memoryBytes*(sbc: Sbc): int = - ## Bytes held by the packed bit tensor. - sbc.bits.len * sizeof(uint32) + ## Bytes held by this memory: the packed bit tensor (bitset) or the counter + ## tensor (counted). + case sbc.mode + of smBitset: sbc.bits.len * sizeof(uint32) + of smCounted: sbc.counters.len * sizeof(uint8) proc occupancy*(sbc: Sbc): float = - ## Fraction of the bit tensor that is set (diagnostic). - var setBits = 0 - for w in sbc.bits: - setBits += countSetBits(w) - result = float(setBits) / float(sbc.nAde * sbc.nAde * sbc.nClasses) + ## Fraction of the memory that is non-empty (diagnostic). + case sbc.mode + of smBitset: + var setBits = 0 + for w in sbc.bits: + setBits += countSetBits(w) + result = float(setBits) / float(sbc.nAde * sbc.nAde * sbc.nClasses) + of smCounted: + var used = 0 + for c in sbc.counters: + if c != 0'u8: inc used + result = float(used) / float(sbc.counters.len) diff --git a/common_libs/tests/measure_counted_sbc.nim b/common_libs/tests/measure_counted_sbc.nim new file mode 100644 index 0000000..3809cad --- /dev/null +++ b/common_libs/tests/measure_counted_sbc.nim @@ -0,0 +1,306 @@ +## Measured evidence for the counted SBC with forgetting. +## +## Three experiments, all deterministic (fixed seeds/permutation): +## +## A. non-stationary MNIST — the same 20k images are streamed twice: pass 1 +## with the true labels, pass 2 with the labels permuted by a fixed π. We +## report accuracy on the "old task" (labels as in pass 1) and the "new +## task" (π labels). This is the forgetting test: can the memory track the +## current mapping? arms: bitset / counted-no-decay / counted+decay. +## +## B. rare-vs-common — a synthetic SBC-level stream with one common/noisy and +## one rare/predictable class. Compares the set-bit vote (`infer`), the raw +## counter sum (`infer`) and the per-cell posterior (`inferProb`). +## +## C. stationary MNIST — counted mode on the reference setup, one pass over the +## full 60k train set. Answers "does counting cost anything when the data is +## stationary?" Compare to the bitset anchor (97.210%). +## +## Fixtures are read-only from `/tmp/bitbrain/BitBrain_C_code` (override with +## `$BITBRAIN_FIXTURES`). Experiment A/B skip cleanly without fixtures; C needs +## them too. Run: +## +## nim c -r --nimcache:/tmp/nc_j102 -d:release --path:common_libs \ +## common_libs/tests/measure_counted_sbc.nim +## +## Env knobs (all optional): BB_EXPERIMENT=A|B|C (default all), +## BB_EPOCH=20000, BB_TESTN=2000, BB_TRAINN=60000, +## BB_DECAY_EVERY=2000, BB_DECAY_SHIFT=1, BB_PRIOR=0.2, BB_SYN_N=20000. + +import std/[os, times, strutils, math, random] +import bitbrain/ade +import bitbrain/sbc +import bitbrain/bitbrain + +const + FixtureDefault = "/tmp/bitbrain/BitBrain_C_code" + W = 2048 + InputSz = 784 + Widths = [6, 8, 10, 12] + NClasses = 10 + +# ── env ────────────────────────────────────────────────────────────────────── + +proc envInt(name: string, default: int): int = + let v = getEnv(name, "").strip() + if v.len == 0: return default + try: parseInt(v) except ValueError: default + +proc envFloat(name: string, default: float): float = + let v = getEnv(name, "").strip() + if v.len == 0: return default + try: parseFloat(v) except ValueError: default + +let + ExpSel = getEnv("BB_EXPERIMENT", "ABC").toUpperAscii() + EpochN = envInt("BB_EPOCH", 20000) + TestN = envInt("BB_TESTN", 2000) + TrainN = envInt("BB_TRAINN", 60000) + DecayEv = envInt("BB_DECAY_EVERY", 2000) + DecaySh = envInt("BB_DECAY_SHIFT", 1) + +# label permutation for the non-stationary task (a fixed derangement) +const Perm = [7, 2, 9, 0, 4, 6, 1, 8, 3, 5] + +# ── fixtures ───────────────────────────────────────────────────────────────── + +proc fixtureDir(): string = + let d = getEnv("BITBRAIN_FIXTURES", "") + if d.len > 0: d else: FixtureDefault + +proc loadInt32(path: string, n: int): seq[int32] = + let f = open(path, fmRead) + defer: f.close() + result = newSeq[int32](n) + if n == 0: return + doAssert f.readBuffer(addr result[0], n * 4) == n * 4, "short read " & path + +proc loadU8(path: string, n: int): seq[uint8] = + let f = open(path, fmRead) + defer: f.close() + result = newSeq[uint8](n) + if n == 0: return + doAssert f.readBuffer(addr result[0], n) == n, "short read " & path + +proc buildFromFixtures(dir: string, mode: SbcMode, + decayEvery, decayShift: int): BitBrain = + var ades: seq[AddressDecoder] + for k in 0 ..< Widths.len: + var ad = initAddressDecoder(W, Widths[k]) + ad.codes = loadInt32(dir / ("AD" & $(k + 1) & "_2048"), W * Widths[k]) + ad.thresholds = loadInt32(dir / ("thresh" & $(k + 1) & "_2048"), W) + ades.add ad + result = initBitBrain(ades, crossPairs(ades.len), NClasses, + mode, decayEvery, decayShift) + +proc inferSlice(bb: BitBrain, data: seq[uint8], i: int): int = + bb.infer(toOpenArray(data, i * InputSz, i * InputSz + InputSz - 1)).label + +proc inferProbSlice(bb: BitBrain, data: seq[uint8], i: int): int = + bb.inferProb(toOpenArray(data, i * InputSz, i * InputSz + InputSz - 1)).label + +# ── Experiment A: non-stationary MNIST ─────────────────────────────────────── + +type NonStatResult = object + name: string + oldE1, newE1, oldE2, newE2: float + newTrace: seq[tuple[at: int, acc: float]] + +proc runNonStatin(dir: string, name: string, mode: SbcMode, + decayEvery, decayShift: int): NonStatResult = + let trainData = loadU8(dir / "train_data", EpochN * InputSz) + let trainLabel = loadU8(dir / "train_label", EpochN) + let testData = loadU8(dir / "test_data", TestN * InputSz) + let testLabel = loadU8(dir / "test_label", TestN) + + var bb = buildFromFixtures(dir, mode, decayEvery, decayShift) + result.name = name + + proc evalOld(): float = + var right = 0 + for i in 0 ..< TestN: + if inferSlice(bb, testData, i) == int(testLabel[i]): inc right + float(right) / float(TestN) + + proc evalNew(): float = + var right = 0 + for i in 0 ..< TestN: + if inferSlice(bb, testData, i) == Perm[int(testLabel[i])]: inc right + float(right) / float(TestN) + + # pass 1: identity labels + for i in 0 ..< EpochN: + bb.learn(toOpenArray(trainData, i * InputSz, i * InputSz + InputSz - 1), + int(trainLabel[i])) + result.oldE1 = evalOld() + result.newE1 = evalNew() + + # pass 2: permuted labels; trace new-task accuracy through the drift + let traceEvery = max(1, EpochN div 8) + for i in 0 ..< EpochN: + bb.learn(toOpenArray(trainData, i * InputSz, i * InputSz + InputSz - 1), + Perm[int(trainLabel[i])]) + if (i + 1) mod traceEvery == 0: + result.newTrace.add (at: i + 1, acc: evalNew()) + result.oldE2 = evalOld() + result.newE2 = evalNew() + +proc experimentA(dir: string) = + if not dirExists(dir): + echo "Skipping experiment A: fixtures not found at ", dir + return + echo "\n=== A. non-stationary MNIST (same ", EpochN, + " images streamed twice; labels permuted in pass 2) ===" + echo "permutation pi = ", Perm + let arms = [ + ("bitset", smBitset, 0, 0), + ("counted-no-decay", smCounted, 0, 0), + ("counted+decay", smCounted, DecayEv, DecaySh)] + var rows: seq[NonStatResult] + for (name, mode, de, ds) in arms: + let t0 = cpuTime() + let r = runNonStatin(dir, name, mode, de, ds) + rows.add r + echo " ran ", name, " in ", formatFloat(cpuTime() - t0, ffDecimal, 1), " s" + echo "" + echo "arm old@pass1 new@pass1 old@pass2 new@pass2 (new = adaptation)" + for r in rows: + echo align(r.name, 20), " ", + align(formatFloat(r.oldE1 * 100, ffDecimal, 2), 8), " ", + align(formatFloat(r.newE1 * 100, ffDecimal, 2), 8), " ", + align(formatFloat(r.oldE2 * 100, ffDecimal, 2), 8), " ", + align(formatFloat(r.newE2 * 100, ffDecimal, 2), 8) + echo "\nnew-task accuracy through pass-2 drift (samples seen in pass 2 -> %):" + for r in rows: + var line = align(r.name, 20) + for (at, acc) in r.newTrace: + line.add " " & $at & ":" & formatFloat(acc * 100, ffDecimal, 1) + echo line + +# ── Experiment B: rare-vs-common (probability vs vote) ─────────────────────── + +proc experimentB() = + let N = envInt("BB_SYN_N", 4000) + let Prior = envFloat("BB_PRIOR", 0.2) + const M = 256 + const SigLen = 16 + const P0 = 0.02 + var rng = initRand(20240924) + let row = @[0'i32] + + proc class0Sample(): seq[int32] = + result = @[] + for c in 0 ..< M: + if rng.rand(1.0) < P0: result.add int32(c) + + proc class1Sample(): seq[int32] = + result = newSeq[int32](SigLen) + for c in 0 ..< SigLen: result[c] = int32(c) + + # four arms over the same generated stream + type Arm = object + name: string + sbc: Sbc + prob: bool + var arms: seq[Arm] = @[ + Arm(name: "bitset/vote", sbc: initSbc(M, 2), prob: false), + Arm(name: "counted/raw-sum", sbc: initCountedSbc(M, 2, 0, 0), prob: false), + Arm(name: "counted/prob", sbc: initCountedSbc(M, 2, 0, 0), prob: true), + Arm(name: "bitset/prob", sbc: initSbc(M, 2), prob: true)] + + # generate the stream once + var stream: seq[seq[int32]] + var labels: seq[int] + for _ in 0 ..< N: + if rng.rand(1.0) < Prior: + stream.add class1Sample(); labels.add 1 + else: + stream.add class0Sample(); labels.add 0 + for i in 0 ..< stream.len: + for a in mitems(arms): + discard a.sbc.learn(row, stream[i], labels[i]) + + # balanced test set: 2000 pure class-1 and 2000 pure class-0 + var testX: seq[seq[int32]] + var testY: seq[int] + for _ in 0 ..< 2000: + testX.add class1Sample(); testY.add 1 + for _ in 0 ..< 2000: + testX.add class0Sample(); testY.add 0 + + echo "\n=== B. rare-vs-common (common class 0 = 80% of a noisy stream, ", + "rare class 1 = ", formatFloat(Prior * 100, ffDecimal, 0), "% but predictable) ===" + echo "arm recall(class0) recall(class1) balanced" + for a in arms: + var r0 = 0 + var r1 = 0 + var n0 = 0 + var n1 = 0 + for i in 0 ..< testX.len: + var label: int + if a.prob: + var sc = newSeq[float](2) + a.sbc.inferProb(row, testX[i], sc) + label = if sc[1] > sc[0]: 1 else: 0 + else: + var ct = newSeq[int](2) + a.sbc.infer(row, testX[i], ct) + label = if ct[1] > ct[0]: 1 else: 0 + if testY[i] == 0: + inc n0 + if label == 0: inc r0 + else: + inc n1 + if label == 1: inc r1 + let acc0 = float(r0) / float(n0) + let acc1 = float(r1) / float(n1) + echo align(a.name, 20), " ", + align(formatFloat(acc0, ffDecimal, 3), 13), " ", + align(formatFloat(acc1, ffDecimal, 3), 13), " ", + align(formatFloat((acc0 + acc1) / 2.0, ffDecimal, 3), 8) + +# ── Experiment C: stationary MNIST counted ─────────────────────────────────── + +proc experimentC(dir: string) = + if not dirExists(dir): + echo "Skipping experiment C: fixtures not found at ", dir + return + let testN = envInt("BB_TESTN", 2000) + let trainData = loadU8(dir / "train_data", TrainN * InputSz) + let trainLabel = loadU8(dir / "train_label", TrainN) + let testData = loadU8(dir / "test_data", testN * InputSz) + let testLabel = loadU8(dir / "test_label", testN) + + echo "\n=== C. stationary MNIST single pass (train ", TrainN, + ", test ", testN, ") ===" + echo "arm vote-argmax% prob-argmax%" + for (name, mode, de, ds) in [ + ("bitset", smBitset, 0, 0), + ("counted-no-decay", smCounted, 0, 0), + ("counted+decay", smCounted, DecayEv, DecaySh)]: + var bb = buildFromFixtures(dir, mode, de, ds) + let t0 = cpuTime() + for i in 0 ..< TrainN: + bb.learn(toOpenArray(trainData, i * InputSz, i * InputSz + InputSz - 1), + int(trainLabel[i])) + let tTrain = cpuTime() - t0 + var cV = 0 + var cP = 0 + for i in 0 ..< testN: + if inferSlice(bb, testData, i) == int(testLabel[i]): inc cV + if inferProbSlice(bb, testData, i) == int(testLabel[i]): inc cP + echo align(name, 20), " ", + align(formatFloat(100.0 * float(cV) / float(testN), ffDecimal, 3), 11), " ", + align(formatFloat(100.0 * float(cP) / float(testN), ffDecimal, 3), 11), + " (train " & formatFloat(tTrain, ffDecimal, 1) & " s, bytes " & + $bb.sbcMemoryBytes & ")" + echo "bitset anchor (from test_bitbrain_mnist, full 60k, all 10k test): 97.210 vote-argmax" + +# ── driver ─────────────────────────────────────────────────────────────────── + +let dir = fixtureDir() +echo "Fixtures: ", dir +if 'A' in ExpSel: experimentA(dir) +if 'B' in ExpSel: experimentB() +if 'C' in ExpSel: experimentC(dir) +echo "\ndone." diff --git a/common_libs/tests/test_bitbrain.nim b/common_libs/tests/test_bitbrain.nim index 1f01d16..0230fd4 100644 --- a/common_libs/tests/test_bitbrain.nim +++ b/common_libs/tests/test_bitbrain.nim @@ -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: diff --git a/docs/bitbrain_counted_sbc.md b/docs/bitbrain_counted_sbc.md new file mode 100644 index 0000000..6e09244 --- /dev/null +++ b/docs/bitbrain_counted_sbc.md @@ -0,0 +1,208 @@ +# Counted SBC with forgetting — design and measured evidence + +Step 1 of 2 of the "give the SBC a counter and a forgetting mechanism" change. +Library: `common_libs/bitbrain/sbc.nim`, `common_libs/bitbrain/bitbrain.nim`. +Harness: `common_libs/tests/measure_counted_sbc.nim`. +Unit tests: `common_libs/tests/test_bitbrain.nim` (56 checks, up from 32). + +All numbers below are **MEASURED** on this machine with `-d:release`, +single-threaded, deterministic (fixed seed / fixed permutation), unless tagged +**INFERRED**. + +--- + +## Why (recap of the diagnosed defects) + +The default SBC is a set-union: `bits: seq[uint32]`, `learn` only *sets* bits. +Two consequences: + +1. A bit records *that* a coincidence went with class `k`, never *how often*. + The target is stochastic, so the honest object is a **probability**. +2. A bit cannot be cleared. `docs/bitbrain_gate_test.md` (`f41cd08`) measured + that the retained-across-rounds regime was the *weakest* precisely because the + idempotent memory only grows and saturates with noise. + +Counted mode replaces the bit with a small counter and adds decay. + +--- + +## Design + +**Counter.** One saturating `uint8` per `(i, j, class)`, `0..255`. Chosen over +packed 4-bit for clarity and speed; the 4-bit cost is reported below as +INFERRED. `learn` increments the observed class's counter (saturating) and +returns the number of coincidence cells touched. Repetition is now evidence: +learning the same sample ten times gives a count of 10, where the bitset is a +no-op after the first. + +**Decay — global fractional decay every `decayEvery` learns.** On a schedule, +every counter is aged once: `c -= c shr decayShift`. `decayShift = 1` is a +halving; higher values forget more slowly; `0` disables decay. Chosen over the +alternatives for three reasons: + +* **Cost.** The decay pass is O(cells) but runs once per `decayEvery` learns, so + the per-learn amortised cost is O(cells / decayEvery) and the hot per-tick + `learn` path stays as cheap as a bit-set. A per-cell EMA decays every touched + cell on every learn — ~`nClasses`× more work per coincidence (at MNIST scale + that is ~10× the learn cost). +* **True forgetting.** It ages cells that are *never visited again*, which a + per-cell EMA cannot (an EMA only decays cells it touches). +* **Simplicity/determinism.** No per-cell timestamps, no extra state beyond a + learn counter, and the same input stream always produces the same memory. + +**Readouts.** + +* `infer` — literal "sum the counters per class" (raw frequency sum). Kept for + the requested semantics and as the baseline. +* `inferProb` — for each observed coincidence, form the per-cell posterior + `P(class | cell) = count[class] / Σ_k count[k]` and sum it per class. This is + the **recommended counted readout**: it is scale-free in the class marginals, + so a single high-count cell cannot dominate a majority of low-count cells. + The bitset `infer` is unchanged. + +**Config (which is which).** The default is and remains `smBitset`; the bitset +path is byte-for-byte unchanged. Counted mode is selected explicitly or at +runtime: + +| knob | kind | values | +|---|---|---| +| `TR_BITBRAIN_MODE` | runtime env | `bitset` (default) / `counted` | +| `TR_BITBRAIN_DECAY_EVERY` | runtime env | learns between decay passes | +| `TR_BITBRAIN_DECAY_SHIFT` | runtime env | decay strength (`0` = off) | +| `-d:bitbrainDecayEvery=N` | compile-time | overrides `DefaultDecayEvery` (1024) | +| `-d:bitbrainDecayShift=N` | compile-time | overrides `DefaultDecayShift` (1) | + +Env is read by `envSbcMode` / `envDecayEvery` / `envDecayShift`; unknown values +fall back to the shipped bitset defaults. The `-d:` defines use `{.intdefine.}` +and were verified to change `DefaultDecayShift` at compile time. + +--- + +## Memory cost + +One `uint8` per `(i, j, class)`, so 8× the packed bit tensor, plus the (small) +AD term. + +| configuration | bitset SBC | counted SBC | ADs | counted total | packed 4-bit (INFERRED) | +|---|---:|---:|---:|---:|---:| +| Reference MNIST: 6 × 2048² × 10 | 30.0 MiB (31,457,280 B) | **240.0 MiB** (251,658,240 B) | 0.34 MiB | 240.3 MiB | 120 MiB | +| Gun-sized: 6 × 512² × 8 | 1.5 MiB (1,572,864 B) | **12.0 MiB** (12,582,912 B) | 90,112 B | 12.1 MiB (12,673,024 B) | 6 MiB | + +The gun-sized byte-per-cell cost is 12 MiB — acceptable for a gun (and the +whole 12.1 MiB model fits comfortably next to the rest of a bot). The MNIST +240 MiB is fine for an offline measurement machine but is 8× the bit tensor; +packed 4-bit (INFERRED, not implemented) would halve both figures. At +MNIST scale, `decayEvery = 2000` costs ≈ 0.2 ms/learn amortised (measured train +30.3 s vs 25.4 s over 60k for the decay arm, i.e. ~0.08 ms/sample extra). + +--- + +## MNIST regression (the proof the library still matches the reference) + +`test_bitbrain_mnist.nim` was run **unchanged** against the pretrained ADs and +full MNIST: + +| reader | this library | reference | +|---|---:|---:| +| Corrected (clean-room) | **97.210 %** | 97.210 | +| Bug-compatible (`i % 32 < 8`) | **96.540 %** | 96.540 | + +Both anchors reproduce to the digit, and the fixed `read_from_sbc` truncation +bug is not regressed. `test_bitbrain` is **56 checks, 0 failures** (the original +32 checks still pass unchanged, plus 24 new counted-mode checks). + +--- + +## Experiment A — non-stationary adaptation (the forgetting proof) + +The same 20,000 MNIST training images are streamed twice: pass 1 with the true +labels, pass 2 with the labels permuted by the fixed `π = [7,2,9,0,4,6,1,8,3,5]`. +Evaluation is on the first 2,000 test images, under the "old" (true) labels and +the "new" (`π`) labels. Because the *same images recur*, the coincidence cells +are revisited under the new mapping — exactly the case where forgetting must +overwrite stale associations. All arms use `infer` (raw argmax). + +| arm | old@pass1 | new@pass1 | old@pass2 | **new@pass2** | +|---|---:|---:|---:|---:| +| bitset (default) | 94.05 | 11.10 | 55.75 | **48.90** | +| counted, no decay | 85.40 | 11.55 | 49.35 | **43.35** | +| counted + decay (`every=2000`, `shift=1`) | 88.45 | 11.10 | 12.60 | **86.20** | + +New-task accuracy as pass 2 proceeds (samples seen in pass 2): + +| arm | 2500 | 5000 | 7500 | 10000 | 12500 | 15000 | 17500 | 20000 | +|---|---:|---:|---:|---:|---:|---:|---:|---:| +| bitset | 11.9 | 12.8 | 13.8 | 15.7 | 17.4 | 22.6 | 31.4 | 48.9 | +| counted, no decay | 11.9 | 12.4 | 13.1 | 13.8 | 14.7 | 17.2 | 21.3 | 43.4 | +| counted + decay | 41.7 | 76.2 | 82.8 | 82.8 | 83.7 | 83.5 | 84.9 | 86.2 | + +The bitset climbs only to ~49 % and the counters without decay do not forget at +all (~43 %): both are stuck with the pass-1 associations. Counted+decay tracks +the change and reaches **86.2 %** on the new task (and its old-task accuracy +falls to 12.6 %, i.e. it genuinely abandoned the old mapping). This is the +measured proof of the mechanism diagnosed in the gate test. + +--- + +## Experiment B — probabilities (rare but predictable vs common but noisy) + +A synthetic SBC-level stream (256 columns, `row = [0]`, so cells are the active +columns). Class 1 is **rare (20 %)** but predictable: its 16-cell signature is +always active. Class 0 is **common (80 %)** but noisy: every sample activates a +random 2 % of all columns, so class 0 slowly sets a bit / lays a small count on +almost every cell. Tested on pure class-1 and pure class-0 inputs. + +| readout | recall(class 0) | recall(class 1) | balanced | +|---|---:|---:|---:| +| bitset / vote (`infer`) | 1.000 | **0.000** | 0.500 | +| counted / raw sum (`infer`) | 0.911 | 1.000 | 0.956 | +| counted / per-cell posterior (`inferProb`) | 0.996 | 1.000 | **0.998** | +| bitset / posterior (`inferProb`) | 1.000 | 0.000 | 0.500 | + +The set-bit vote gives the common class a full vote at every cell it ever +touched, so it never finds the rare class (balanced 0.500). Summing counters +finds it, and the **per-cell posterior beats the raw sum** (0.998 vs 0.956): +the raw sum lets one high-count cell dominate, the posterior normalises per +cell. + +--- + +## Experiment C — does counting cost anything when the data is stationary? + +Reference setup, one online pass over the full 60k MNIST train set, evaluated on +the first 2,000 test images (the bitset arm is run through the same harness so +the comparison is apples-to-apples): + +| arm | vote-argmax % | prob-argmax % | SBC bytes | +|---|---:|---:|---:| +| bitset | 95.700 | 95.850 | 31,457,280 | +| counted, no decay | 87.700 | 93.450 | 251,658,240 | +| counted + decay | 88.150 | 94.900 | 251,658,240 | + +The full-10k bitset anchor is 97.210 % (the 2,000-image subset is simply +harder). **Counting hurts on stationary MNIST.** The raw counter sum loses ~8 +points (87.7 vs 95.7); the per-cell posterior recovers most of it (93.5 / 94.9) +but still trails the bitset by ~1–2 points. This is not a saturation artifact: +the same gap appears at 2,000 training samples (bitset 89.65 vs counted-no-decay +80.35 vote / 88.60 prob) and 8,000 samples (92.90 vs 83.35 / 91.70), where +per-cell counts are far from 255. + +So the honest reading is: counting+decay is a **trade**, not a free win — it buys +forgetting and true probabilities at the cost of roughly a point of stationary +accuracy (with the recommended posterior readout) and 8× the memory. + +--- + +## Direct answer + +* **Forgetting: YES, MEASURED.** After a label permutation the bitset is stuck + at 48.9 % on the new task and no-decay counters at 43.4 %, while + counters+decay reaches 86.2 %. The global decay bound is what makes the memory + adaptive. +* **Probabilities: YES, MEASURED.** Summing counters resolves a rare but + predictable class that the set-bit vote cannot (balanced 0.956–0.998 vs + 0.500). The per-cell posterior (`inferProb`) is the readout to use, and it + beats the raw counter sum. +* **Stationary cost: real, MEASURED.** Counting does not help MNIST; the raw + sum loses ~8 points and the posterior readout ~1–2 points versus the bitset. + Bitset remains the default for exactly this reason.