89370008da
The gun has never contributed anything: Tsetlin.vHits was byte-for-byte equal to Linear.vHits in every measured round of every run, because its learned correction was always exactly 0. Six diagnosed defects fixed, plus one that was required to make the first one work: 1. Type I now conditions on the clause output. It previously rewarded included true literals unconditionally, omitting Granmo's (c=0, lk=1) -> toward Exclude counter-force, so true literals ratcheted toward Include forever. This was the root cause of the saturation. 2. Type II was unreachable dead code: its guard required cOut==1 AND lits[lit]==0 AND st>0 (included), but cOut==1 guarantees every included literal is 1. Its direction was wrong too - it should increment EXCLUDED false literals when the clause fires. 3. Resource allocation restored: Granmo's (T - clip(v,-T,T))/(2T) target replaces |error|/(2*RESID_MAX); TM_T was only an output normaliser. 4. Label baseline fixed - the factor-2 shrink. predX = linearX + cx, so the label was delta - cx while the learner's output IS cx, giving error = delta - 2cx and a fixed point of cx = delta/2: HALF the needed correction even with perfect feedback. TmTrace now stores linearX/linearY and training uses delta. 5. Hits no longer zero their label (a hit means |miss| < 18px, not 0). 6. The enemy-energy feature was duplicated - tmEncodeFrame passed state.selfEnergy with a stale comment claiming enemyEnergy was absent, while WorldState.enemyEnergy exists. Enemy-energy rules were literally unrepresentable. 7. REQUIRED EXTRA: tmEvalClause now implements Granmo Eq. 6 - an all-Exclude clause outputs 1 during learning and 0 during classification. Without it, fix #1 deadlocks every clause at empty. MEASURED EFFECT (energy-threshold-turner fixture, seed 1): mean included literals per active clause 714.0 -> 13.8 active clauses 100/100 -> 53/100 nonzero corrections 8/764 -> 708/764 Tsetlin virtual hits (Linear = 27/400) 27/400 -> 69/400 Divergence achieved: offline on 7/8 fixtures, and in a live gauntlet (RandomMover: Tsetlin 199/1200 vs Linear 288/1200, vDropped=vStarved=0). Tsetlin now LEARNS but is not yet competitive with Linear - the regression head is untuned, flagged as follow-up rather than claimed as a win. Also ignores compiled test harnesses that have no file extension, which the existing '**/tests/test_*' rule misses.
505 lines
21 KiB
Nim
505 lines
21 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,
|
||
singleFrame = false): 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). If
|
||
## `singleFrame`, only the current frame is exposed (frames 1..9 zeroed); this
|
||
## isolates how much of the clause bloat comes from the 10-frame redundancy.
|
||
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
|
||
if singleFrame:
|
||
var w1: array[TM_WINDOW_SIZE, TmFrameEncoded]
|
||
w1[0] = window[0]
|
||
vec = tmEncodeFullVector(w1, selfState)
|
||
else:
|
||
vec = tmEncodeFullVector(window, 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], singleFrame = false): seq[Sample] =
|
||
for c in cfgs: result.add samplesFromFixture(makeFixture(c), Threshold, singleFrame)
|
||
|
||
# ── 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 setEnergyAllFrames(s: var Sample, energy: float) =
|
||
## Counterfactual: rewrite the 11-bit Gray-coded energy field of every frame to
|
||
## `energy`, leaving every other channel untouched. If the model keys on the
|
||
## energy threshold, its prediction follows this rewrite.
|
||
let raw = clamp(int(energy * 10.0), 0, 1500)
|
||
let gray = raw xor (raw shr 1)
|
||
for f in 0..<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
|
||
|
||
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]; setEnergyAllFrames(cfLow[i], 10.0) # label would be 1
|
||
cfHigh[i] = test[i]; setEnergyAllFrames(cfHigh[i], 50.0) # 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 ""
|
||
|
||
# Ablation: single-frame input removes the 10-frame redundancy. If the clause
|
||
# bloat is a redundancy artefact, effective width should fall here.
|
||
echo ""
|
||
echo "── ablation: single-frame input (current frame only) ──"
|
||
let train1 = gather(TrainCfgs, singleFrame = true)
|
||
let test1 = gather(TestCfgs, singleFrame = true)
|
||
var labels1: seq[int]
|
||
for s in train1: labels1.add s.y
|
||
let net1 = shuffleRun(train1, labels1, 20250920'u64, Epochs)
|
||
let teAcc1 = accuracy(net1, test1)
|
||
let varying1 = varyingBits(train1)
|
||
var total1, varyingInc1, active1 = 0
|
||
for c in 0..<C_N_CLAUSES:
|
||
let inc = included(net1, c)
|
||
if inc.len > 0: inc active1
|
||
total1 += inc.len
|
||
for lit in inc:
|
||
let bitIdx = if lit < C_N_IN: lit else: lit - C_N_IN
|
||
if varying1[bitIdx]: inc varyingInc1
|
||
echo &"single-frame TEST accuracy = {teAcc1*100:.2f}% (majority {testMajority*100:.2f}%), " &
|
||
&"mean {float(total1)/float(max(active1,1)):.1f} nominal / " &
|
||
&"{float(varyingInc1)/float(max(active1,1)):.1f} effective literals/clause"
|
||
for c in 0..<C_N_POS:
|
||
if included(net1, c).len > 0:
|
||
echo &" C{c:>2}: {clauseStringVarying(net1, c, varying1, 8)}"
|
||
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()
|