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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user