## 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.. 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..= 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..= 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.. readable propositional logic ─────────── proc bitField(bitIdx: int): string = ## Decode a bit index of the 870-bit frame-stacked vector to "f.[]". 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.. 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.. 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.. 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.. 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.. {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.. 0: echo &" C{c:>2}: {clauseStringVarying(net, c, varying)}" echo "" echo "── learned clauses (negative class: predicts STRAIGHT) ──" for c in C_N_POS.. 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.. 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.. 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..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()