Files
SirStone 40ba96f649 BitBrain SBC: counted mode + global decay (forgetting, probabilities)
Adds an smCounted storage mode alongside the default smBitset. Each
(i,j,class) cell becomes a saturating uint8 counter; learn increments it and
a global fractional decay (c -= c shr decayShift every decayEvery learns)
makes forgetting possible. infer sums raw counters; new inferProb sums the
per-cell posterior P(class|cell) (scale-free, recommended readout).

Bitset path is the default and byte-for-byte unchanged: test_bitbrain 56/56
(was 32), and test_bitbrain_mnist reproduces 97.210% corrected / 96.540%
bug-compatible exactly.

Counted mode configurable at runtime (TR_BITBRAIN_MODE / TR_BITBRAIN_DECAY_*)
and compile time (-d:bitbrainDecay*). Measured: forgetting (86.2% vs 48.9% on
a permuted-label stream), probabilities (rare-class balanced 0.998 vs 0.500),
and the stationary cost (counted hurts MNIST; see docs/bitbrain_counted_sbc.md).

Harness: common_libs/tests/measure_counted_sbc.nim
2026-09-25 08:39:10 +02:00

284 lines
12 KiB
Nim

## Sparse Binary Coincidence (SBC) memory — clean-room implementation.
##
## Implements the supervised half of the BitBrain algorithm as described in
##
## "BitBrain and Sparse Binary Coincidence (SBC) memories",
## Frontiers in Neuroinformatics 17:1125844, 2023.
##
## Written from the published algorithm only; see `ade.nim` for the clean-room
## note. The reference C is GPL-3.0 (c) University of Manchester and was not
## copied.
##
## Mechanism
## ---------
## Two Address Decoders (ADs) sit on the two axes of a 2-D memory. A pair of
## *simultaneously firing* ADEs `(i, j)` is a **coincidence** and addresses one
## memory cell. That cell holds a class bitmask with one bit per class
## (`nClasses` bits, one-hot encoding in the paper's default).
##
## 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
## all `nClasses` bits of one coincidence are contiguous. This differs from the
## 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 depth.
nAde*: int ## number of ADEs on each axis (the paper's `w`)
nClasses*: int ## number of classes (the paper's `D`)
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 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 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 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 =
## 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
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]) =
## 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
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 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 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)