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
|
||||
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 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
|
||||
|
||||
@@ -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
|
||||
|
||||
+210
-41
@@ -16,15 +16,18 @@
|
||||
## memory cell. That cell holds a class bitmask with one bit per class
|
||||
## (`nClasses` bits, one-hot encoding in the paper's default).
|
||||
##
|
||||
## 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.
|
||||
## 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
|
||||
@@ -32,83 +35,249 @@
|
||||
## 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-bit depth.
|
||||
## 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`)
|
||||
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 =
|
||||
## 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 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 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 j >= 0 and j < sbc.nAde
|
||||
doAssert class >= 0 and class < sbc.nClasses
|
||||
let bit = bitIndex(sbc, i, j, class)
|
||||
(sbc.bits[bit shr 5] and (1'u32 shl (bit and 31))) != 0'u32
|
||||
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 =
|
||||
## Set the `class` bit for every coincidence between a firing row ADE and a
|
||||
## firing column ADE. Returns the number of bits that were newly set (0 if the
|
||||
## sample added no information, e.g. it was already learned).
|
||||
## 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
|
||||
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
|
||||
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]) =
|
||||
## Count, per class, how many observed coincidences have their class bit set.
|
||||
## `counts` is *accumulated into* (not reset), so a container can sum several
|
||||
## SBCs. It must be at least `nClasses` long.
|
||||
## 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
|
||||
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]
|
||||
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 the packed bit tensor.
|
||||
sbc.bits.len * sizeof(uint32)
|
||||
## 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 bit tensor that is set (diagnostic).
|
||||
var setBits = 0
|
||||
for w in sbc.bits:
|
||||
setBits += countSetBits(w)
|
||||
result = float(setBits) / float(sbc.nAde * sbc.nAde * sbc.nClasses)
|
||||
## 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)
|
||||
|
||||
Reference in New Issue
Block a user