## 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)