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
+14
View File
@@ -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 width, ADE count and class count — nothing is hardcoded to 784/10. Deterministic
given the seed. Dependencies: `std/` only. 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 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; 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 this implementation stores them full-size (half-size packing is a separate memory
+82 -11
View File
@@ -20,8 +20,14 @@
## 10-SBC variant adds 4 within-AD ("half-size") SBCs; `withinPairs` builds those ## 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 ## (this implementation stores them full-size — the half-size packing is a
## separate memory optimisation). ## 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 import ade, sbc
export ade, sbc export ade, sbc
@@ -50,9 +56,15 @@ proc withinPairs*(nAdes: int): seq[SbcSpec] =
result.add SbcSpec(row: a, col: a) result.add SbcSpec(row: a, col: a)
proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec], 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 ## 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). ## 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 ades.len > 0, "need at least one AD"
doAssert nClasses > 0, "need at least one class" doAssert nClasses > 0, "need at least one class"
let w = ades[0].nAde let w = ades[0].nAde
@@ -65,11 +77,17 @@ proc initBitBrain*(ades: seq[AddressDecoder], specs: seq[SbcSpec],
for s in 0 ..< specs.len: for s in 0 ..< specs.len:
doAssert specs[s].row >= 0 and specs[s].row < ades.len doAssert specs[s].row >= 0 and specs[s].row < ades.len
doAssert specs[s].col >= 0 and specs[s].col < 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, proc buildRandomBitBrain*(widths: seq[int], nAde, inputWidth, nClasses: int,
seed: int64, 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 ## Initialise a whole BitBrain with random ADs. If `specs` is empty the
## paper's 6 cross-AD SBCs are used. Fully deterministic given `seed`. ## paper's 6 cross-AD SBCs are used. Fully deterministic given `seed`.
var rng = initRand(seed) var rng = initRand(seed)
@@ -77,7 +95,35 @@ proc buildRandomBitBrain*(widths: seq[int], nAde, inputWidth, nClasses: int,
for w in widths: for w in widths:
ades.add initRandomAddressDecoder(nAde, w, inputWidth, rng) ades.add initRandomAddressDecoder(nAde, w, inputWidth, rng)
let pairs = if specs.len == 0: crossPairs(widths.len) else: specs 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.} = proc nAdes*(bb: BitBrain): int {.inline.} =
bb.ades.len 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) = proc learn*[T: SomeInteger](bb: var BitBrain, input: openArray[T], class: int) =
## One online supervised step: set the `class` bit of every observed ## 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]] var lists: seq[seq[int32]]
bb.fireInto(input, lists) bb.fireInto(input, lists)
for s in 0 ..< bb.sbcs.len: 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]): proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
tuple[label: int, counts: seq[int]] = tuple[label: int, counts: seq[int]] =
## Drive every AD, count set class bits in every SBC, and return the argmax ## Drive every AD, accumulate per-class integer evidence in every SBC (set-bit
## class plus the aggregated per-class counts. Ties go to the lowest class ## count, or summed counters), and return the argmax class plus the aggregated
## index (matching the reference's `>` scan that keeps the first maximum). ## 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]] var lists: seq[seq[int32]]
bb.fireInto(input, lists) bb.fireInto(input, lists)
result.counts = newSeq[int](bb.nClasses) result.counts = newSeq[int](bb.nClasses)
@@ -123,14 +171,37 @@ proc infer*[T: SomeInteger](bb: BitBrain, input: openArray[T]):
best = result.counts[k] best = result.counts[k]
result.label = 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 = 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: for ad in bb.ades:
result += ad.memoryBytes result += ad.memoryBytes
for s in bb.sbcs: for s in bb.sbcs:
result += s.memoryBytes result += s.memoryBytes
proc sbcMemoryBytes*(bb: BitBrain): int = 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: for s in bb.sbcs:
result += s.memoryBytes 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
+210 -41
View File
@@ -16,15 +16,18 @@
## memory cell. That cell holds a class bitmask with one bit per class ## memory cell. That cell holds a class bitmask with one bit per class
## (`nClasses` bits, one-hot encoding in the paper's default). ## (`nClasses` bits, one-hot encoding in the paper's default).
## ##
## Learning is **idempotent**: `learn` *sets* the bit for the observed class; ## In the default **bitset** mode learning is **idempotent**: `learn` *sets* the
## setting it again is a no-op. There is no clearing, no learning rate, no decay ## bit for the observed class; setting it again is a no-op. There is no clearing,
## and no epoch — one pass through the data is a complete supervised training run, ## no learning rate, no decay and no epoch — one pass through the data is a
## and a second pass changes nothing. ## complete supervised training run, and a second pass changes nothing.
## ##
## Inference uses the *same* address decoding: for each observed coincidence every ## 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 ## 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. ## 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 ## Bit layout
## ---------- ## ----------
## The bit for `(i, j, class)` lives at `((i * nAde) + j) * nClasses + class`, so ## 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 ## reference C's layout but is a bijection onto the same set of triples; the
## learned rule is identical. ## 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 import std/bitops
type type
SbcMode* = enum
smBitset ## default: idempotent set-bit memory (reference-compatible)
smCounted ## saturating per-cell class counters with global fractional decay
Sbc* = object 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`) nAde*: int ## number of ADEs on each axis (the paper's `w`)
nClasses*: int ## number of classes (the paper's `D`) 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 = 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 nAde > 0, "nAde must be positive"
doAssert nClasses > 0, "nClasses must be positive" doAssert nClasses > 0, "nClasses must be positive"
result.nAde = nAde result.nAde = nAde
result.nClasses = nClasses result.nClasses = nClasses
result.mode = smBitset
let nbits = nAde * nAde * nClasses let nbits = nAde * nAde * nClasses
result.bits = newSeq[uint32]((nbits + 31) div 32) 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) = proc clear*(sbc: var Sbc) =
## Forget everything. This is how a monotone memory is wiped (e.g. when a bot ## Forget everything. This is how a monotone memory is wiped (e.g. when a bot
## switches enemy and must not carry state across battles). ## switches enemy and must not carry state across battles).
for i in 0 ..< sbc.bits.len: for i in 0 ..< sbc.bits.len:
sbc.bits[i] = 0'u32 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.} = proc bitIndex(sbc: Sbc, i, j, class: int): int {.inline.} =
((i * sbc.nAde) + j) * sbc.nClasses + class ((i * sbc.nAde) + j) * sbc.nClasses + class
proc bitAt*(sbc: Sbc, i, j, class: int): bool {.inline.} = 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 i >= 0 and i < sbc.nAde
doAssert j >= 0 and j < sbc.nAde doAssert j >= 0 and j < sbc.nAde
doAssert class >= 0 and class < sbc.nClasses doAssert class >= 0 and class < sbc.nClasses
let bit = bitIndex(sbc, i, j, class) let idx = bitIndex(sbc, i, j, class)
(sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32 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], proc learn*(sbc: var Sbc, rowActive, colActive: openArray[int32],
class: int): int = class: int): int =
## Set the `class` bit for every coincidence between a firing row ADE and a ## Bitset mode: set the `class` bit for every coincidence between a firing row
## firing column ADE. Returns the number of bits that were newly set (0 if the ## ADE and a firing column ADE, returning the number of bits newly set (0 if
## sample added no information, e.g. it was already learned). ## 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" doAssert class >= 0 and class < sbc.nClasses, "class out of range"
let D = sbc.nClasses let D = sbc.nClasses
for r in rowActive: case sbc.mode
let i = int(r) of smBitset:
for c in colActive: for r in rowActive:
let j = int(c) let i = int(r)
let bit = ((i * sbc.nAde) + j) * D + class for c in colActive:
let w = bit shr 5 let j = int(c)
let m = 1'u32 shl (bit and 31) let bit = ((i * sbc.nAde) + j) * D + class
if (sbc.bits[w] and m) == 0'u32: let w = bit shr 5
sbc.bits[w] = sbc.bits[w] or m 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 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], proc infer*(sbc: Sbc, rowActive, colActive: openArray[int32],
counts: var seq[int]) = counts: var seq[int]) =
## Count, per class, how many observed coincidences have their class bit set. ## Bitset mode: count, per class, how many observed coincidences have their
## `counts` is *accumulated into* (not reset), so a container can sum several ## class bit set. Counted mode: sum the `class` counters over the observed
## SBCs. It must be at least `nClasses` long. ## 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" doAssert counts.len >= sbc.nClasses, "counts buffer too small"
let D = sbc.nClasses let D = sbc.nClasses
for r in rowActive: case sbc.mode
let i = int(r) of smBitset:
for c in colActive: for r in rowActive:
let j = int(c) let i = int(r)
let base = ((i * sbc.nAde) + j) * D for c in colActive:
for k in 0 ..< D: let j = int(c)
let bit = base + k let base = ((i * sbc.nAde) + j) * D
if (sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32: for k in 0 ..< D:
inc counts[k] 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 = proc memoryBytes*(sbc: Sbc): int =
## Bytes held by the packed bit tensor. ## Bytes held by this memory: the packed bit tensor (bitset) or the counter
sbc.bits.len * sizeof(uint32) ## tensor (counted).
case sbc.mode
of smBitset: sbc.bits.len * sizeof(uint32)
of smCounted: sbc.counters.len * sizeof(uint8)
proc occupancy*(sbc: Sbc): float = proc occupancy*(sbc: Sbc): float =
## Fraction of the bit tensor that is set (diagnostic). ## Fraction of the memory that is non-empty (diagnostic).
var setBits = 0 case sbc.mode
for w in sbc.bits: of smBitset:
setBits += countSetBits(w) var setBits = 0
result = float(setBits) / float(sbc.nAde * sbc.nAde * sbc.nClasses) 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)
+306
View File
@@ -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."
+101 -1
View File
@@ -13,7 +13,7 @@
## * correct, non-crashing behaviour on an unseen input, ## * correct, non-crashing behaviour on an unseen input,
## * homeostatic threshold adaptation moves the firing rate toward target. ## * 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/ade
import bitbrain/sbc import bitbrain/sbc
import bitbrain/bitbrain import bitbrain/bitbrain
@@ -237,6 +237,102 @@ proc testMemoryAccounting() =
check "SBC tensor bit count is exact", check "SBC tensor bit count is exact",
bb.sbcMemoryBytes * 8 == 6 * 512 * 512 * 8 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 ─────────────────────────────────────────────────────────────────── # ── driver ───────────────────────────────────────────────────────────────────
testAdeScoring() testAdeScoring()
@@ -246,6 +342,10 @@ testShuffledControl()
testUnseenInput() testUnseenInput()
testHomeostasis() testHomeostasis()
testMemoryAccounting() testMemoryAccounting()
testCountedMemory()
testCountedForgetting()
testCountedProbability()
testCountedEnv()
echo "\n", checks, " checks, ", failures, " failure(s)" echo "\n", checks, " checks, ", failures, " failure(s)"
if failures > 0: if failures > 0:
+208
View File
@@ -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.