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