40ba96f649
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
307 lines
12 KiB
Nim
307 lines
12 KiB
Nim
## 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."
|