## 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 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.. 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.. 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.. 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 "" # ── 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..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.. 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.. 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..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..= 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.. 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()