d5061ee215
Task 1 - the acceptance proof was unrunnable because RecordWorldState was a
compile-time const set to false. It is now a RUNTIME switch
(let RecordWorldState* = existsEnv("TR_RECORD_WORLDSTATE")), default OFF, so
ordinary runs write no fixture, and acceptance_offline_vs_online.nim enables
it for the battle it spawns and clears it afterwards. Restored and run twice:
12/12 deterministic guns match exactly (128-tick and 546-tick battles), with
Tsetlin reported separately as stochastic. Both nimble build variants clean.
Task 2 - does a compact encoding turn the TM's 99.35% into a READABLE rule?
Measured across window sizes (fixed seed, no tuning):
frames TEST acc eff.lits/clause firing clauses counterfactual low/high/mean
10 99.35% 152.8 37 100/24/62.4%
3 95.94% 54.9 35 96/20/58.6%
2 99.48% 39.6 38 95/25/60.7%
1 98.30% 19.2 45 100/24/62.3%
So 2 frames is strictly better than 10 on BOTH axes: +0.13 accuracy for 4x
smaller clauses. The 3-frame dip is non-monotonic and left unexplained rather
than smoothed over.
A readable rule WAS partially recovered. Five clauses carry the exact Gray
form !g10 ^ !g9 ^ !g8; g10 is inert in this data, so the effective rule is the
2-literal proposition !g9 ^ !g8, i.e. energy < 25.6. That is a genuine
threshold in readable propositional form - but at 25.6, NOT the labelled 30,
because 256 is a power-of-two Gray boundary expressible in two literals while
300 needs a longer conjunction. The TM found the nearest SIMPLE threshold.
The honest caveat: that threshold is not the ensemble's decision mechanism.
The counterfactual follow rate (high 24%, mean 62.3%) is statistically
identical at 1, 2 and 10 frames, so compactness did not make the model read
energy - its vote is carried by co-occurring bearing/velocity/heading/wall
literals. Also identified: clauses containing all 11 Gray energy bits are
satisfied at exactly one raw value (50, the dataset floor), so they are
'energy has hit the floor' detectors, not thresholds.
Methodological fix worth keeping: the earlier single-frame counterfactual
wrote energy into all 10 frame slots including the zeroed ones, reviving dead
clauses and producing a spurious 2% high-follow rate. setEnergyFrames now
rewrites only the exposed frames; the corrected figure is 24%.
725 lines
30 KiB
Nim
725 lines
30 KiB
Nim
## L2 — can a Tsetlin Machine learn a high-level pattern?
|
||
##
|
||
## A STANDALONE Granmo binary classifier (Table 2/3, resource allocation
|
||
## Eq. 8-11, empty clause = 1 during learning per Eq. 6) over the SAME
|
||
## frame-stacked binary encoding the Tsetlin gun uses (870 bits = 10 frames × 83
|
||
## + 40 self bits). It is independent of the gun: its own clause teams, its own
|
||
## training loop.
|
||
##
|
||
## The source of truth is `synthesizeEnergyThresholdTurner` in
|
||
## `gun_harness/offline_range.nim`:
|
||
## RULE: energy(t) = max(5, e0 - decay*t)
|
||
## straight while energy >= 30, hard turn each tick below 30.
|
||
## So the label is literally `energy(t) < 30`, a propositional predicate over the
|
||
## 11-bit Gray-coded energy field of the current frame. The experiment asks
|
||
## whether the TM recovers it, and — the payoff — whether the learned clauses are
|
||
## readable as that rule.
|
||
##
|
||
## Train and test use DIFFERENT trajectories (different e0/decay/speed/start
|
||
## position/hardTurnDeg, same threshold=30). A model that memorises position or
|
||
## the training trajectory cannot generalise; one that reads the energy bits can.
|
||
##
|
||
## Run: nim c -r common_libs/tests/test_tm_pattern_learning.nim
|
||
|
||
import std/[math, strformat, random, algorithm, strutils]
|
||
|
||
import gun_harness/offline_range
|
||
import guns/tsetlin # reuse tmEncodeFrame / tmEncodeSelf / tmEncodeFullVector
|
||
|
||
# ── standalone Granmo binary TM ──────────────────────────────────────────────
|
||
|
||
const
|
||
C_N_IN = TM_TOTAL_BITS # 870
|
||
C_N_LITS = C_N_IN * 2 # 1740 literals (each bit + its negation)
|
||
C_N_POS = 24
|
||
C_N_NEG = 24
|
||
C_N_CLAUSES = C_N_POS + C_N_NEG
|
||
C_N_STATES = 64 # automaton state range [-64, 64]
|
||
C_S = 3.9 # Granmo specificity
|
||
C_T = float(C_N_POS) # summation target
|
||
|
||
type
|
||
ClsNet = object
|
||
states: array[C_N_CLAUSES * C_N_LITS, int16]
|
||
|
||
Sample = object
|
||
lits: array[C_N_LITS, uint8]
|
||
y: int
|
||
|
||
# Fast local xorshift so the experiment is reproducible without touching the
|
||
# gun's std/random stream.
|
||
var rngState: uint64 = 0x9E3779B97F4A7C15'u64
|
||
|
||
proc seedRng(s: uint64) = rngState = (if s == 0: 1'u64 else: s)
|
||
proc nextU64(): uint64 {.inline.} =
|
||
rngState = rngState xor (rngState shl 13)
|
||
rngState = rngState xor (rngState shr 7)
|
||
rngState = rngState xor (rngState shl 17)
|
||
rngState
|
||
proc rand01(): float {.inline.} = float(nextU64() shr 40) / 16777216.0 # 24-bit
|
||
|
||
proc litIdx(c, lit: int): int {.inline.} = c * C_N_LITS + lit
|
||
proc polarity(c: int): float {.inline.} =
|
||
if c < C_N_POS: 1.0 else: -1.0
|
||
|
||
proc evalClause(net: ClsNet, c: int, lits: array[C_N_LITS, uint8],
|
||
learning: bool): uint8 =
|
||
var hasInc = false
|
||
for lit in 0..<C_N_LITS:
|
||
if net.states[litIdx(c, lit)] > 0:
|
||
hasInc = true
|
||
if lits[lit] == 0'u8: return 0'u8
|
||
if hasInc: return 1'u8
|
||
if learning: return 1'u8
|
||
return 0'u8
|
||
|
||
proc forward(net: ClsNet, lits: array[C_N_LITS, uint8], learning: bool): float =
|
||
var v = 0.0
|
||
for c in 0..<C_N_CLAUSES:
|
||
v += polarity(c) * float(evalClause(net, c, lits, learning))
|
||
clamp(v, -C_T, C_T)
|
||
|
||
proc predictLabel(net: ClsNet, lits: array[C_N_LITS, uint8]): int =
|
||
if forward(net, lits, false) >= 0.0: 1 else: 0
|
||
|
||
proc typeIFeedback(net: var ClsNet, c: int, lits: array[C_N_LITS, uint8]) =
|
||
## Granmo Table 2, collapsed to the resulting state move.
|
||
let cOut = evalClause(net, c, lits, learning = true)
|
||
for lit in 0..<C_N_LITS:
|
||
let si = litIdx(c, lit)
|
||
var st = int(net.states[si])
|
||
if lits[lit] == 1'u8:
|
||
if cOut == 1'u8:
|
||
if rand01() < (C_S - 1.0) / C_S: st = min(st + 1, C_N_STATES)
|
||
else:
|
||
if rand01() < 1.0 / C_S: st = max(st - 1, -C_N_STATES)
|
||
else:
|
||
if rand01() < 1.0 / C_S: st = max(st - 1, -C_N_STATES)
|
||
net.states[si] = int16(st)
|
||
|
||
proc typeIIFeedback(net: var ClsNet, c: int, lits: array[C_N_LITS, uint8]) =
|
||
## Granmo Table 3: penalise the exclusion of a zero literal when c=1.
|
||
let cOut = evalClause(net, c, lits, learning = true)
|
||
if cOut == 1'u8:
|
||
for lit in 0..<C_N_LITS:
|
||
if lits[lit] == 0'u8:
|
||
let si = litIdx(c, lit)
|
||
if net.states[si] <= 0:
|
||
net.states[si] = int16(min(int(net.states[si]) + 1, C_N_STATES))
|
||
|
||
proc trainOne(net: var ClsNet, lits: array[C_N_LITS, uint8], y: int) =
|
||
let v = forward(net, lits, learning = true)
|
||
let p = if y == 1: (C_T - v) / (2.0 * C_T)
|
||
else: (C_T + v) / (2.0 * C_T)
|
||
if p <= 0.0: return
|
||
for c in 0..<C_N_CLAUSES:
|
||
if rand01() >= p: continue
|
||
let isPos = c < C_N_POS
|
||
if (y == 1 and isPos) or (y == 0 and not isPos):
|
||
typeIFeedback(net, c, lits)
|
||
else:
|
||
typeIIFeedback(net, c, lits)
|
||
|
||
# ── dataset: frame-stacked encoding of energy-threshold-turner trajectories ──
|
||
|
||
proc makeLiterals(vec: TmBinaryVector): array[C_N_LITS, uint8] =
|
||
for i in 0..<C_N_IN:
|
||
result[i] = vec[i]
|
||
result[i + C_N_IN] = 1'u8 - vec[i]
|
||
|
||
proc samplesFromFixture(fx: Fixture, threshold: float,
|
||
windowFrames = TM_WINDOW_SIZE): seq[Sample] =
|
||
## Replicates the gun's window bookkeeping exactly: index 0 = newest frame,
|
||
## shifted once per tick. Skips the first 10 ticks (window warm-up).
|
||
## `windowFrames` controls how many of the MOST RECENT frames are exposed to
|
||
## the classifier; older frames stay zeroed (constant, hence inert). Shrinking
|
||
## it isolates how much of the clause bloat is 10-frame redundancy rather than
|
||
## a property of the Tsetlin Machine itself.
|
||
var window: array[TM_WINDOW_SIZE, TmFrameEncoded]
|
||
var count = 0
|
||
for t in 0..<fx.states.len:
|
||
let s = fx.states[t]
|
||
let dist = hypot(s.enemyX - s.selfX, s.enemyY - s.selfY)
|
||
let bearing = radToDeg(arctan2(s.enemyY - s.selfY, s.enemyX - s.selfX))
|
||
let frame = tmEncodeFrame(
|
||
bearing, dist, s.enemySpeed, s.enemyHeading,
|
||
s.arenaHeight - s.enemyY, s.enemyY,
|
||
s.arenaWidth - s.enemyX, s.enemyX,
|
||
s.enemyEnergy) # fix 6: enemy energy is encoded
|
||
for i in countdown(TM_WINDOW_SIZE - 1, 1): window[i] = window[i - 1]
|
||
window[0] = frame
|
||
if count < TM_WINDOW_SIZE: inc count
|
||
if count < TM_WINDOW_SIZE: continue
|
||
let selfState = tmEncodeSelf(
|
||
s.arenaHeight - s.selfY, s.selfY,
|
||
s.arenaWidth - s.selfX, s.selfX,
|
||
s.selfEnergy, true)
|
||
var vec: TmBinaryVector
|
||
var w1: array[TM_WINDOW_SIZE, TmFrameEncoded]
|
||
for i in 0 ..< min(windowFrames, TM_WINDOW_SIZE):
|
||
w1[i] = window[i]
|
||
vec = tmEncodeFullVector(w1, selfState)
|
||
result.add Sample(lits: makeLiterals(vec),
|
||
y: (if s.enemyEnergy < threshold: 1 else: 0))
|
||
|
||
type
|
||
RuleCfg = object
|
||
e0, decay, speed, ex, ey, turn: float
|
||
|
||
const
|
||
Threshold = 30.0
|
||
TrainCfgs = [
|
||
RuleCfg(e0: 60.0, decay: 0.40, speed: 4.0, ex: 120.0, ey: 250.0, turn: 15.0),
|
||
RuleCfg(e0: 50.0, decay: 0.50, speed: 5.0, ex: 100.0, ey: 300.0, turn: 20.0),
|
||
RuleCfg(e0: 70.0, decay: 0.60, speed: 3.0, ex: 150.0, ey: 200.0, turn: 12.0),
|
||
RuleCfg(e0: 42.0, decay: 0.30, speed: 6.0, ex: 240.0, ey: 420.0, turn: 25.0),
|
||
RuleCfg(e0: 55.0, decay: 0.45, speed: 4.5, ex: 80.0, ey: 150.0, turn: 18.0),
|
||
RuleCfg(e0: 66.0, decay: 0.55, speed: 5.5, ex: 300.0, ey: 480.0, turn: 22.0),
|
||
RuleCfg(e0: 48.0, decay: 0.38, speed: 3.5, ex: 180.0, ey: 360.0, turn: 14.0),
|
||
RuleCfg(e0: 62.0, decay: 0.48, speed: 5.2, ex: 60.0, ey: 460.0, turn: 24.0),
|
||
]
|
||
TestCfgs = [
|
||
RuleCfg(e0: 58.0, decay: 0.42, speed: 4.2, ex: 130.0, ey: 280.0, turn: 16.0),
|
||
RuleCfg(e0: 64.0, decay: 0.52, speed: 4.8, ex: 90.0, ey: 220.0, turn: 19.0),
|
||
RuleCfg(e0: 45.0, decay: 0.35, speed: 5.8, ex: 200.0, ey: 330.0, turn: 23.0),
|
||
RuleCfg(e0: 52.0, decay: 0.44, speed: 3.8, ex: 280.0, ey: 180.0, turn: 17.0),
|
||
]
|
||
|
||
proc makeFixture(c: RuleCfg): Fixture =
|
||
synthesizeEnergyThresholdTurner(ticks = 200, ex = c.ex, ey = c.ey,
|
||
e0 = c.e0, decay = c.decay, threshold = Threshold, speed = c.speed,
|
||
hardTurnDeg = c.turn)
|
||
|
||
proc gather(cfgs: openArray[RuleCfg], windowFrames = TM_WINDOW_SIZE): seq[Sample] =
|
||
for c in cfgs: result.add samplesFromFixture(makeFixture(c), Threshold, windowFrames)
|
||
|
||
# ── clause decoding: literal index -> readable propositional logic ───────────
|
||
|
||
proc bitField(bitIdx: int): string =
|
||
## Decode a bit index of the 870-bit frame-stacked vector to "f<frame>.<field>[<b>]".
|
||
if bitIdx < TM_FRAME_BITS * TM_WINDOW_SIZE:
|
||
let f = bitIdx div TM_FRAME_BITS
|
||
let off = bitIdx mod TM_FRAME_BITS
|
||
if off < 8: return &"f{f}.bearingSin[{off}]"
|
||
if off < 16: return &"f{f}.bearingCos[{off-8}]"
|
||
if off < 23: return &"f{f}.distance[{off-16}]"
|
||
if off < 28: return &"f{f}.velocity[{off-23}]"
|
||
if off < 36: return &"f{f}.headingSin[{off-28}]"
|
||
if off < 44: return &"f{f}.headingCos[{off-36}]"
|
||
if off < 51: return &"f{f}.wallN[{off-44}]"
|
||
if off < 58: return &"f{f}.wallS[{off-51}]"
|
||
if off < 65: return &"f{f}.wallE[{off-58}]"
|
||
if off < 72: return &"f{f}.wallW[{off-65}]"
|
||
return &"f{f}.energy[{off-72}]"
|
||
else:
|
||
let off = bitIdx - TM_FRAME_BITS * TM_WINDOW_SIZE
|
||
if off < 7: return &"self.wallN[{off}]"
|
||
if off < 14: return &"self.wallS[{off-7}]"
|
||
if off < 21: return &"self.wallE[{off-14}]"
|
||
if off < 28: return &"self.wallW[{off-21}]"
|
||
if off < 39: return &"self.energy[{off-28}]"
|
||
return "self.canFire"
|
||
|
||
proc describeLiteral(lit: int): string =
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
let negated = lit >= C_N_IN
|
||
(if negated: "!" else: "") & bitField(bitIdx)
|
||
|
||
proc isEnergyField(lit: int): bool =
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if bitIdx < TM_FRAME_BITS * TM_WINDOW_SIZE:
|
||
let off = bitIdx mod TM_FRAME_BITS
|
||
return off >= 72 and off < 83
|
||
let off = bitIdx - TM_FRAME_BITS * TM_WINDOW_SIZE
|
||
return off >= 28 and off < 39
|
||
|
||
proc included(net: ClsNet, c: int): seq[int] =
|
||
for lit in 0..<C_N_LITS:
|
||
if net.states[litIdx(c, lit)] > 0: result.add lit
|
||
|
||
proc clauseString(net: ClsNet, c: int, maxLits = 10): string =
|
||
let inc = included(net, c)
|
||
if inc.len == 0: return "TRUE (empty)"
|
||
var parts: seq[string]
|
||
for i, lit in inc:
|
||
if i >= maxLits:
|
||
parts.add &"… (+{inc.len - maxLits} more)"
|
||
break
|
||
parts.add describeLiteral(lit)
|
||
result = parts.join(" ∧ ")
|
||
|
||
# ── leakage controls: which channel does the TM actually read? ──────────────
|
||
|
||
type
|
||
FieldChannel = enum
|
||
fcBearing, fcDistance, fcVelocity, fcHeading, fcWalls, fcEnergy, fcSelf
|
||
|
||
proc channelOf(bitIdx: int): FieldChannel =
|
||
if bitIdx < TM_FRAME_BITS * TM_WINDOW_SIZE:
|
||
let off = bitIdx mod TM_FRAME_BITS
|
||
if off < 16: return fcBearing
|
||
if off < 23: return fcDistance
|
||
if off < 28: return fcVelocity
|
||
if off < 44: return fcHeading
|
||
if off < 72: return fcWalls
|
||
return fcEnergy
|
||
return fcSelf
|
||
|
||
proc maskChannel(s: Sample, ch: FieldChannel): Sample =
|
||
## Zero out every bit of one channel (the literal becomes the constant 0, its
|
||
## negation the constant 1). NOTE: this also breaks any clause that happens to
|
||
## INCLUDE an inert literal on that channel, so a large drop is not proof of
|
||
## causal use — the counterfactual probe below is the reliable test.
|
||
result = s
|
||
for i in 0..<C_N_IN:
|
||
if channelOf(i) == ch:
|
||
result.lits[i] = 0'u8
|
||
result.lits[i + C_N_IN] = 1'u8
|
||
|
||
proc setEnergyFrames(s: var Sample, energy: float, frames: int) =
|
||
## Counterfactual: rewrite the 11-bit Gray-coded energy field of the first
|
||
## `frames` frames to `energy`, leaving every other channel untouched. Only
|
||
## the exposed frames are written: compact encodings zero frames
|
||
## `frames..<TM_WINDOW_SIZE`, and those constant bits must stay constant or the
|
||
## probe would revive clauses that never fire in the real encoding.
|
||
let raw = clamp(int(energy * 10.0), 0, 1500)
|
||
let gray = raw xor (raw shr 1)
|
||
for f in 0..<min(frames, TM_WINDOW_SIZE):
|
||
for b in 0..<11:
|
||
let bit = uint8((gray shr (10 - b)) and 1)
|
||
let idx = f * TM_FRAME_BITS + 72 + b
|
||
s.lits[idx] = bit
|
||
s.lits[idx + C_N_IN] = 1'u8 - bit
|
||
|
||
proc varyingBits(data: seq[Sample]): seq[bool] =
|
||
## Which of the 870 input bits actually vary across the dataset. Included
|
||
## literals on constant bits are INERT: they inflate the nominal clause width
|
||
## without changing when the clause fires.
|
||
result = newSeq[bool](C_N_IN)
|
||
for i in 0..<C_N_IN:
|
||
let v0 = data[0].lits[i]
|
||
for s in data:
|
||
if s.lits[i] != v0:
|
||
result[i] = true
|
||
break
|
||
|
||
proc clauseStringVarying(net: ClsNet, c: int, varying: seq[bool],
|
||
maxLits = 12): string =
|
||
var parts: seq[string]
|
||
var totalVary = 0
|
||
for lit in included(net, c):
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if varying[bitIdx]:
|
||
inc totalVary
|
||
if parts.len < maxLits: parts.add describeLiteral(lit)
|
||
if totalVary == 0: return "(no varying literals)"
|
||
result = parts.join(" ∧ ")
|
||
if totalVary > maxLits: result.add &" … (+{totalVary - maxLits} varying)"
|
||
|
||
# ── main experiment ──────────────────────────────────────────────────────────
|
||
|
||
proc accuracy(net: ClsNet, data: seq[Sample]): float =
|
||
if data.len == 0: return 0.0
|
||
var ok = 0
|
||
for s in data:
|
||
if predictLabel(net, s.lits) == s.y: inc ok
|
||
ok.float / data.len.float
|
||
|
||
proc confusion(net: ClsNet, data: seq[Sample]): tuple[tp, tn, fp, fn: int] =
|
||
for s in data:
|
||
let p = predictLabel(net, s.lits)
|
||
if p == 1 and s.y == 1: inc result.tp
|
||
elif p == 0 and s.y == 0: inc result.tn
|
||
elif p == 1 and s.y == 0: inc result.fp
|
||
else: inc result.fn
|
||
|
||
type
|
||
ClauseStats = object
|
||
meanNominal*: float
|
||
meanEffective*: float ## incl. literals on bits that actually vary (excludes inert)
|
||
maxInc*: int
|
||
active*: int ## clauses with >= 1 included literal
|
||
firing*: int ## clauses that output 1 on >= 1 data sample (classify semantics)
|
||
firingPerSample*: float
|
||
energyFrac*: float ## energy-field literals / total included literals
|
||
|
||
proc clauseStats(net: ClsNet, data: seq[Sample],
|
||
varying: seq[bool]): ClauseStats =
|
||
## Width + firing statistics for one trained model. `firing` counts clauses
|
||
## under classification semantics (empty clause = 0), i.e. clauses that
|
||
## actually contribute a vote on at least one sample.
|
||
var totalNominal, totalEffective = 0
|
||
var energyLits, totalLits = 0
|
||
for c in 0..<C_N_CLAUSES:
|
||
let inc = included(net, c)
|
||
if inc.len == 0: continue
|
||
inc result.active
|
||
totalNominal += inc.len
|
||
result.maxInc = max(result.maxInc, inc.len)
|
||
for lit in inc:
|
||
inc totalLits
|
||
if isEnergyField(lit): inc energyLits
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if varying[bitIdx]: inc totalEffective
|
||
let denom = max(result.active, 1)
|
||
result.meanNominal = float(totalNominal) / float(denom)
|
||
result.meanEffective = float(totalEffective) / float(denom)
|
||
result.energyFrac = if totalLits > 0: float(energyLits) / float(totalLits) else: 0.0
|
||
var everFired: array[C_N_CLAUSES, bool]
|
||
var fires = 0
|
||
for s in data:
|
||
for c in 0..<C_N_CLAUSES:
|
||
if evalClause(net, c, s.lits, learning = false) == 1'u8:
|
||
everFired[c] = true
|
||
inc fires
|
||
for b in everFired:
|
||
if b: inc result.firing
|
||
result.firingPerSample =
|
||
if data.len > 0: float(fires) / float(data.len) else: 0.0
|
||
|
||
type
|
||
EnergyProfile = object
|
||
bits*: int ## number of current-frame (f0) energy literals in the clause
|
||
count*: int ## raw energy values (0..1500) satisfying that sub-conjunction
|
||
maxRun*: int ## longest contiguous run of satisfying raw values (threshold => large)
|
||
lo*, hi*: int ## min/max satisfying raw value (energy = raw/10)
|
||
spec*: string ## the literal spec, e.g. "!g10 !g9 !g8"
|
||
|
||
proc grayBits11(raw: int): array[11, uint8] =
|
||
## Same encoding as tmEncodeFrame's energy field: 11-bit Gray code, MSB first.
|
||
let g = raw xor (raw shr 1)
|
||
for b in 0..<11:
|
||
result[b] = uint8((g shr (10 - b)) and 1)
|
||
|
||
proc clauseEnergyProfile(net: ClsNet, c: int): EnergyProfile =
|
||
## Evaluate ONLY the current-frame (f0) energy literals of a clause over every
|
||
## raw energy value 0..1500. A faithful threshold rule would satisfy the
|
||
## conjunction over one long contiguous run ending near raw 300 (energy 30); a
|
||
## scattered pattern with maxRun == 1 is a Gray-code coincidence, not a
|
||
## threshold. bits == 0 means the clause has no f0-energy literals.
|
||
var lits: seq[tuple[b: int, want: uint8]]
|
||
var parts: seq[string]
|
||
for lit in included(net, c):
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if bitIdx >= 72 and bitIdx < 83:
|
||
let b = bitIdx - 72
|
||
let pos = lit < C_N_IN
|
||
lits.add (b, (if pos: 1'u8 else: 0'u8))
|
||
parts.add (if pos: "g" else: "!g") & $(10 - b)
|
||
result.bits = lits.len
|
||
if lits.len == 0: return
|
||
result.spec = parts.join(" ")
|
||
result.lo = -1
|
||
var run = 0
|
||
for raw in 0..1500:
|
||
let g = grayBits11(raw)
|
||
var ok = true
|
||
for (b, want) in lits:
|
||
if g[b] != want: ok = false; break
|
||
if ok:
|
||
inc result.count
|
||
inc run
|
||
result.maxRun = max(result.maxRun, run)
|
||
if result.lo < 0: result.lo = raw
|
||
result.hi = raw
|
||
else:
|
||
run = 0
|
||
|
||
proc shuffleRun(train: seq[Sample], labels: seq[int], seedState: uint64,
|
||
epochs: int): ClsNet =
|
||
seedRng(seedState)
|
||
var order: seq[int]
|
||
for i in 0..<train.len: order.add i
|
||
for epoch in 0..<epochs:
|
||
for i in countdown(order.high, 1):
|
||
let j = int(nextU64() mod uint64(i + 1))
|
||
swap(order[i], order[j])
|
||
for idx in order:
|
||
trainOne(result, train[idx].lits, labels[idx])
|
||
|
||
var failures = 0
|
||
proc check(name: string, ok: bool) =
|
||
if ok: echo "PASS: ", name
|
||
else: echo "FAIL: ", name; inc failures
|
||
|
||
proc main() =
|
||
const Epochs = 25
|
||
echo "=== L2: Tsetlin Machine pattern-learning benchmark ==="
|
||
echo &"encoding: {C_N_IN} bits ({TM_WINDOW_SIZE} frames × {TM_FRAME_BITS} + {TM_SELF_BITS} self), " &
|
||
&"{C_N_LITS} literals"
|
||
echo &"classifier: {C_N_POS}+{C_N_NEG} clauses, T={C_T:.0f}, s={C_S}, states=[-{C_N_STATES},{C_N_STATES}]"
|
||
echo &"rule: label = (enemy energy < {Threshold:.0f}); threshold fixed across train/test"
|
||
|
||
let train = gather(TrainCfgs)
|
||
let test = gather(TestCfgs)
|
||
var trainPos, testPos = 0
|
||
for s in train: trainPos += s.y
|
||
for s in test: testPos += s.y
|
||
let trainMajority = max(trainPos, train.len - trainPos).float / train.len.float
|
||
let testMajority = max(testPos, test.len - testPos).float / test.len.float
|
||
echo &"train samples={train.len} (pos={trainPos}, neg={train.len-trainPos}), majority={trainMajority*100:.1f}%"
|
||
echo &"test samples={test.len} (pos={testPos}, neg={test.len-testPos}), majority={testMajority*100:.1f}%"
|
||
echo ""
|
||
|
||
var trainLabels: seq[int]
|
||
for s in train: trainLabels.add s.y
|
||
var net = shuffleRun(train, trainLabels, 20250920'u64, Epochs)
|
||
|
||
let trAcc = accuracy(net, train)
|
||
let teAcc = accuracy(net, test)
|
||
let cm = confusion(net, test)
|
||
echo &"train accuracy = {trAcc*100:.2f}% (majority {trainMajority*100:.2f}%)"
|
||
echo &"TEST accuracy = {teAcc*100:.2f}% (majority {testMajority*100:.2f}%)"
|
||
echo &"TEST confusion: tp={cm.tp} tn={cm.tn} fp={cm.fp} fn={cm.fn}"
|
||
echo ""
|
||
|
||
# Control: identical pipeline on shuffled labels. If the TM were memorising
|
||
# trajectory structure rather than the rule, this would also score high.
|
||
var shuffled = trainLabels
|
||
for i in countdown(shuffled.high, 1):
|
||
let j = int(nextU64() mod uint64(i + 1))
|
||
swap(shuffled[i], shuffled[j])
|
||
let ctrl = shuffleRun(train, shuffled, 777'u64, Epochs)
|
||
let ctrlAcc = accuracy(ctrl, test)
|
||
echo &"control (shuffled labels): TEST accuracy = {ctrlAcc*100:.2f}% (should be ~majority)"
|
||
echo ""
|
||
|
||
# Nominal vs effective clause width (constant bits are inert).
|
||
let varying = varyingBits(train)
|
||
var nVary = 0
|
||
for v in varying: (if v: inc nVary)
|
||
var totalInc, varyingInc, energyInc, activeClauses, maxInc = 0
|
||
for c in 0..<C_N_CLAUSES:
|
||
let inc = included(net, c)
|
||
if inc.len > 0: inc activeClauses
|
||
totalInc += inc.len
|
||
maxInc = max(maxInc, inc.len)
|
||
for lit in inc:
|
||
if isEnergyField(lit): inc energyInc
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if varying[bitIdx]: inc varyingInc
|
||
echo &"input bits varying across dataset: {nVary}/{C_N_IN}"
|
||
echo &"learned clauses: {activeClauses}/{C_N_CLAUSES} non-empty, " &
|
||
&"mean {float(totalInc)/float(max(activeClauses,1)):.1f} literals/clause (nominal), " &
|
||
&"mean {float(varyingInc)/float(max(activeClauses,1)):.1f} effective (varying), max {maxInc}"
|
||
if totalInc > 0:
|
||
echo &"energy-field literals: {energyInc}/{totalInc} ({energyInc.float/totalInc.float*100:.0f}%)"
|
||
echo ""
|
||
|
||
# Counterfactual energy probe: force all frames' energy below / above the
|
||
# threshold and see how often the prediction follows. A model that reads the
|
||
# energy threshold is near 100%; one that reads a proxy (e.g. heading) stays
|
||
# near its default rate.
|
||
var cfLow = test
|
||
var cfHigh = test
|
||
for i in 0..<test.len:
|
||
cfLow[i] = test[i]; setEnergyFrames(cfLow[i], 10.0, TM_WINDOW_SIZE) # label would be 1
|
||
cfHigh[i] = test[i]; setEnergyFrames(cfHigh[i], 50.0, TM_WINDOW_SIZE) # label would be 0
|
||
var followLow, followHigh = 0
|
||
for i in 0..<test.len:
|
||
if predictLabel(net, cfLow[i].lits) == 1: inc followLow
|
||
if predictLabel(net, cfHigh[i].lits) == 0: inc followHigh
|
||
let energyFollow = (followLow.float + followHigh.float) / (2.0 * test.len.float)
|
||
echo &"counterfactual energy flip (all frames -> energy 10 / 50): " &
|
||
&"follows low {followLow*100 div test.len}%, high {followHigh*100 div test.len}%, mean {energyFollow*100:.1f}%"
|
||
echo ""
|
||
|
||
echo "field-masking probe (TEST accuracy with one channel zeroed; confounded by inert literals):"
|
||
for ch in [fcEnergy, fcHeading, fcVelocity, fcDistance, fcBearing, fcWalls, fcSelf]:
|
||
var masked = test
|
||
for i in 0..<masked.len: masked[i] = maskChannel(masked[i], ch)
|
||
echo &" mask {($ch):<10} -> {accuracy(net, masked)*100:6.2f}%"
|
||
echo ""
|
||
|
||
echo "── learned clauses (positive class: predicts TURN) ──"
|
||
echo " [varying literals only; inert constant-bit literals omitted]"
|
||
for c in 0..<C_N_POS:
|
||
if included(net, c).len > 0:
|
||
echo &" C{c:>2}: {clauseStringVarying(net, c, varying)}"
|
||
echo ""
|
||
echo "── learned clauses (negative class: predicts STRAIGHT) ──"
|
||
for c in C_N_POS..<C_N_CLAUSES:
|
||
if included(net, c).len > 0:
|
||
echo &" C{c:>2}: {clauseStringVarying(net, c, varying)}"
|
||
echo ""
|
||
|
||
# ── encoding comparison: how many frames does the rule actually need? ──
|
||
# Keep the full-stack model above as the reference, then train fresh models on
|
||
# windows that expose only the N most recent frames (older frames zeroed). The
|
||
# label and every other channel are unchanged, so this isolates the effect of
|
||
# 10-frame temporal redundancy on clause width and on whether the energy rule
|
||
# is read. Hyperparameters are identical to the reference run (no tuning).
|
||
echo ""
|
||
echo "── encoding comparison (label = energy < threshold, unchanged) ──"
|
||
echo " frames TEST acc eff.lits/cl firing cl energy% cf(low/high/mean)"
|
||
echo " -------------------------------------------------------------------------"
|
||
|
||
type EncResult = object
|
||
frames: int
|
||
net: ClsNet
|
||
trainData: seq[Sample]
|
||
testData: seq[Sample]
|
||
varying: seq[bool]
|
||
teAcc: float
|
||
stats: ClauseStats
|
||
cfLow*: int ## % of test samples whose prediction follows energy -> 10
|
||
cfHigh*: int ## % whose prediction follows energy -> 50 (straight)
|
||
cfMean*: float
|
||
|
||
var encs: seq[EncResult]
|
||
for frames in [TM_WINDOW_SIZE, 3, 2, 1]:
|
||
var e = EncResult(frames: frames)
|
||
if frames == TM_WINDOW_SIZE:
|
||
e.trainData = train
|
||
e.testData = test
|
||
e.net = net
|
||
else:
|
||
e.trainData = gather(TrainCfgs, windowFrames = frames)
|
||
e.testData = gather(TestCfgs, windowFrames = frames)
|
||
var labels: seq[int]
|
||
for s in e.trainData: labels.add s.y
|
||
e.net = shuffleRun(e.trainData, labels, 20250920'u64, Epochs)
|
||
e.varying = varyingBits(e.trainData)
|
||
e.teAcc = accuracy(e.net, e.testData)
|
||
e.stats = clauseStats(e.net, e.testData, e.varying)
|
||
var cfL = e.testData
|
||
var cfH = e.testData
|
||
for i in 0..<e.testData.len:
|
||
cfL[i] = e.testData[i]; setEnergyFrames(cfL[i], 10.0, e.frames)
|
||
cfH[i] = e.testData[i]; setEnergyFrames(cfH[i], 50.0, e.frames)
|
||
var fL, fH = 0
|
||
for i in 0..<e.testData.len:
|
||
if predictLabel(e.net, cfL[i].lits) == 1: inc fL
|
||
if predictLabel(e.net, cfH[i].lits) == 0: inc fH
|
||
e.cfLow = fL * 100 div e.testData.len
|
||
e.cfHigh = fH * 100 div e.testData.len
|
||
e.cfMean = (fL.float + fH.float) / (2.0 * e.testData.len.float) * 100.0
|
||
encs.add e
|
||
echo &" {frames:>4} {e.teAcc*100:7.2f}% {e.stats.meanEffective:10.1f} " &
|
||
&"{e.stats.firing:>9} {e.stats.energyFrac*100:8.1f}% " &
|
||
&"{e.cfLow:>4}/{e.cfHigh:>4}/{e.cfMean:>5.1f}"
|
||
|
||
# "Keeps accuracy" = TEST accuracy within 1.5 percentage points of the full
|
||
# 10-frame model. Pick the smallest such window (fewest frames = most readable).
|
||
# The full table above is printed regardless, so the tradeoff is transparent.
|
||
let fullAcc = encs[0].teAcc
|
||
var best = encs[0]
|
||
for e in encs:
|
||
if e.frames < best.frames and e.teAcc >= fullAcc - 0.015:
|
||
best = e
|
||
echo ""
|
||
if best.frames == TM_WINDOW_SIZE:
|
||
echo &"NOTE: no compact window stayed within 1.5 points of full ({fullAcc*100:.2f}%); reporting the full stack."
|
||
echo &"most compact encoding within 1.5 points of full ({fullAcc*100:.2f}%): " &
|
||
&"{best.frames} frame(s) -> TEST {best.teAcc*100:.2f}%, " &
|
||
&"{best.stats.meanEffective:.1f} effective literals/clause, " &
|
||
&"{best.stats.firing} firing clauses"
|
||
|
||
echo ""
|
||
echo &"── learned clauses, compact {best.frames}-frame encoding (positive: TURN) ──"
|
||
for c in 0..<C_N_POS:
|
||
if included(best.net, c).len > 0:
|
||
echo &" C{c:>2}: {clauseStringVarying(best.net, c, best.varying)}"
|
||
echo ""
|
||
echo &"── learned clauses, compact {best.frames}-frame encoding (negative: STRAIGHT) ──"
|
||
for c in C_N_POS..<C_N_CLAUSES:
|
||
if included(best.net, c).len > 0:
|
||
echo &" C{c:>2}: {clauseStringVarying(best.net, c, best.varying)}"
|
||
|
||
# Positive-clause audit: for each positive clause that ever fires, how well
|
||
# does "fires" agree with the true label (energy < threshold) and how much of
|
||
# its effective width is energy-field literals? A clause that recovered the
|
||
# threshold should fire mostly on positives and be energy-dominated.
|
||
var posN = 0
|
||
for s in best.testData:
|
||
if s.y == 1: inc posN
|
||
echo ""
|
||
echo &"── compact {best.frames}-frame positive-clause audit " &
|
||
&"(label = energy < {Threshold:.0f}) ──"
|
||
echo " clause fires/764 agree(+)% energyLits/effLits recall%"
|
||
for c in 0..<C_N_POS:
|
||
let inc = included(best.net, c)
|
||
var vLits, eLits = 0
|
||
for lit in inc:
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if best.varying[bitIdx]:
|
||
inc vLits
|
||
if isEnergyField(lit): inc eLits
|
||
if vLits == 0: continue
|
||
var fires, agree = 0
|
||
for s in best.testData:
|
||
if evalClause(best.net, c, s.lits, learning = false) == 1'u8:
|
||
inc fires
|
||
if s.y == 1: inc agree
|
||
if fires == 0: continue
|
||
echo &" C{c:>2} {fires:>4}/{best.testData.len:<5} {agree*100 div fires:>8}% " &
|
||
&"{eLits:>6}/{vLits:<6} {agree*100 div max(posN,1):>6}%"
|
||
|
||
# Energy-literal audit: for each positive clause with current-frame energy
|
||
# literals, evaluate that sub-conjunction ALONE across the whole energy range.
|
||
# A faithful threshold rule lights up one long contiguous run ending near raw
|
||
# 300 (energy 30). A scattered result (longestRun small) is a Gray-code
|
||
# coincidence, not a threshold -- however well the full clause predicts.
|
||
echo ""
|
||
echo &"── compact {best.frames}-frame energy-literal audit " &
|
||
&"(f0 energy bits alone, all raw values 0..1500) ──"
|
||
echo " clause E-lits satisfying longestRun raw range energy-bit form verdict"
|
||
for c in 0..<C_N_POS:
|
||
let p = clauseEnergyProfile(best.net, c)
|
||
if p.bits == 0: continue
|
||
let verdict =
|
||
if p.maxRun == 1: "scattered"
|
||
elif p.lo == 0 and p.hi >= 250 and p.hi <= 350 and p.maxRun == p.hi + 1:
|
||
&"threshold-shaped (energy < {(p.hi + 1).float / 10.0:.1f})"
|
||
else: "partial/other"
|
||
echo &" C{c:>2} {p.bits:>4} {p.count:>7} {p.maxRun:>8} " &
|
||
&"{p.lo:>5}..{p.hi:<5} {p.spec:<38} {verdict}"
|
||
|
||
# Counterfactual energy probe on the compact encoding: rewrite every frame's
|
||
# energy field to a fixed value and see whether the prediction follows. This
|
||
# is the test of whether the model learned the RIGHT reason (energy) or a
|
||
# heading proxy.
|
||
var ccfLow = best.testData
|
||
var ccfHigh = best.testData
|
||
for i in 0..<best.testData.len:
|
||
ccfLow[i] = best.testData[i]; setEnergyFrames(ccfLow[i], 10.0, best.frames)
|
||
ccfHigh[i] = best.testData[i]; setEnergyFrames(ccfHigh[i], 50.0, best.frames)
|
||
var cfLowFollow, cfHighFollow = 0
|
||
for i in 0..<best.testData.len:
|
||
if predictLabel(best.net, ccfLow[i].lits) == 1: inc cfLowFollow
|
||
if predictLabel(best.net, ccfHigh[i].lits) == 0: inc cfHighFollow
|
||
let cfMean = (cfLowFollow.float + cfHighFollow.float) / (2.0 * best.testData.len.float)
|
||
echo ""
|
||
echo &"counterfactual energy flip on the compact {best.frames}-frame encoding: " &
|
||
&"follows low {cfLowFollow*100 div best.testData.len}%, " &
|
||
&"high {cfHighFollow*100 div best.testData.len}%, mean {cfMean*100:.1f}% " &
|
||
&"(full-stack was 100% / 24% / 62.4%)"
|
||
echo ""
|
||
|
||
# The two classes are separable by the rule, so a model that found it should
|
||
# clearly beat the majority baseline and the shuffled-label control.
|
||
if teAcc > testMajority + 0.05:
|
||
echo &"VERDICT: the TM generalised above majority ({teAcc*100:.1f}% vs {testMajority*100:.1f}%), " &
|
||
&"while the shuffled-label control stayed at {ctrlAcc*100:.1f}%."
|
||
else:
|
||
echo &"VERDICT: no evidence the TM learned the rule " &
|
||
&"(TEST {teAcc*100:.1f}% vs majority {testMajority*100:.1f}%)."
|
||
echo ""
|
||
|
||
var energyVaries = false
|
||
for i in 0..<C_N_IN:
|
||
if channelOf(i) == fcEnergy and varying[i]: energyVaries = true
|
||
check "energy channel is present and varying in the dataset", energyVaries
|
||
check "TM test accuracy beats the majority baseline by >5 points",
|
||
teAcc > testMajority + 0.05
|
||
check "shuffled-label control does NOT beat the majority baseline by >5 points",
|
||
ctrlAcc <= testMajority + 0.05
|
||
|
||
if failures > 0:
|
||
echo "\n", failures, " check(s) FAILED"
|
||
quit(1)
|
||
echo "\nAll pattern-learning checks passed."
|
||
|
||
when isMainModule:
|
||
main()
|