fix(tsetlin): make the TM actually learn - saturation 714 -> 13.8 literals/clause

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.
This commit is contained in:
2026-09-20 23:59:32 +02:00
parent e0bfa9b5e1
commit 89370008da
6 changed files with 804 additions and 44 deletions
+9
View File
@@ -41,3 +41,12 @@ worktrees/
# Test fixtures are data, not logs - the *.jsonl rule above was written for # Test fixtures are data, not logs - the *.jsonl rule above was written for
# training logs and silently excluded the entire gun-range fixture set. # training logs and silently excluded the entire gun-range fixture set.
!tools/fixtures/**/*.jsonl !tools/fixtures/**/*.jsonl
# Compiled test harnesses have no extension; the test_* rule misses them
common_libs/tests/tm_measure
common_libs/tests/run_range
common_libs/tests/gen_synthetic_fixtures
common_libs/tests/acceptance_offline_vs_online
common_libs/tests/test_tsetlin_gun
common_libs/tests/test_tsetlin_live
common_libs/tests/test_tm_pattern_learning
+127 -44
View File
@@ -9,16 +9,16 @@ import gun_harness/virtual_bullets as vb # PowerBins (power-bin count for trace
# ── Binary encoding (adapted from BNNBot_garage/src/binary_encoding.nim) ───── # ── Binary encoding (adapted from BNNBot_garage/src/binary_encoding.nim) ─────
const const
TM_FRAME_BITS = 83 TM_FRAME_BITS* = 83
TM_SELF_BITS = 40 TM_SELF_BITS* = 40
TM_WINDOW_SIZE = 10 TM_WINDOW_SIZE* = 10
TM_TOTAL_BITS* = TM_FRAME_BITS * TM_WINDOW_SIZE + TM_SELF_BITS # 870 TM_TOTAL_BITS* = TM_FRAME_BITS * TM_WINDOW_SIZE + TM_SELF_BITS # 870
TM_MAX_DISTANCE = 1414.0 # diagonal of 1000x1000 arena TM_MAX_DISTANCE = 1414.0 # diagonal of 1000x1000 arena
type type
TmBinaryVector = array[TM_TOTAL_BITS, uint8] TmBinaryVector* = array[TM_TOTAL_BITS, uint8]
TmFrameEncoded = array[TM_FRAME_BITS, uint8] TmFrameEncoded* = array[TM_FRAME_BITS, uint8]
TmSelfEncoded = array[TM_SELF_BITS, uint8] TmSelfEncoded* = array[TM_SELF_BITS, uint8]
proc tmToGray(value: int): int = value xor (value shr 1) proc tmToGray(value: int): int = value xor (value shr 1)
@@ -28,8 +28,8 @@ proc tmToBits(value: int, bits: int): seq[uint8] =
for i in 0..<bits: for i in 0..<bits:
result[bits - 1 - i] = uint8((gray shr i) and 1) result[bits - 1 - i] = uint8((gray shr i) and 1)
proc tmEncodeFrame(bearing, distance, velocity, heading, proc tmEncodeFrame*(bearing, distance, velocity, heading,
wallN, wallS, wallE, wallW, energy: float): TmFrameEncoded = wallN, wallS, wallE, wallW, energy: float): TmFrameEncoded =
var offset = 0 var offset = 0
# bearing sin: 8 bits # bearing sin: 8 bits
@@ -72,7 +72,7 @@ proc tmEncodeFrame(bearing, distance, velocity, heading,
let eBits = tmToBits(clamp(int(energy * 10.0), 0, 1500), 11) let eBits = tmToBits(clamp(int(energy * 10.0), 0, 1500), 11)
for i in 0..<11: result[offset + i] = eBits[i] for i in 0..<11: result[offset + i] = eBits[i]
proc tmEncodeSelf(wallN, wallS, wallE, wallW, energy: float, canFire: bool): TmSelfEncoded = proc tmEncodeSelf*(wallN, wallS, wallE, wallW, energy: float, canFire: bool): TmSelfEncoded =
var offset = 0 var offset = 0
for wall in [wallN, wallS, wallE, wallW]: for wall in [wallN, wallS, wallE, wallW]:
let wBits = tmToBits(clamp(int(wall / 1000.0 * 99.0), 0, 99), 7) let wBits = tmToBits(clamp(int(wall / 1000.0 * 99.0), 0, 99), 7)
@@ -83,8 +83,8 @@ proc tmEncodeSelf(wallN, wallS, wallE, wallW, energy: float, canFire: bool): TmS
offset += 11 offset += 11
result[offset] = if canFire: 1'u8 else: 0'u8 result[offset] = if canFire: 1'u8 else: 0'u8
proc tmEncodeFullVector(window: array[TM_WINDOW_SIZE, TmFrameEncoded], proc tmEncodeFullVector*(window: array[TM_WINDOW_SIZE, TmFrameEncoded],
self: TmSelfEncoded): TmBinaryVector = self: TmSelfEncoded): TmBinaryVector =
var offset = 0 var offset = 0
for i in 0..<TM_WINDOW_SIZE: for i in 0..<TM_WINDOW_SIZE:
for j in 0..<TM_FRAME_BITS: for j in 0..<TM_FRAME_BITS:
@@ -119,72 +119,114 @@ proc tmStateIdx(outIdx, clause, lit: int): int {.inline.} =
proc tmPolarity(clause: int): float {.inline.} = proc tmPolarity(clause: int): float {.inline.} =
if clause < TM_HALF: 1.0 else: -1.0 if clause < TM_HALF: 1.0 else: -1.0
type
TmClauseStats* = object
## Include-count statistics over both clause teams (TM_N_OUT * TM_N_CLAUSES).
## "Active" = a clause with at least one included literal at the sampled
## moment. A healthy TM settles into a SPARSE regime (single digits to low
## tens of literals per clause). Saturation at hundreds of literals/clause
## (measured mean ≈714 on the energy-threshold fixture before the fix) is the
## failure signature of the broken Type I feedback rule.
nClauses*: int
nActive*: int
minIncluded*: int
maxIncluded*: int
meanIncluded*: float ## mean include count over ACTIVE clauses
meanIncludedAll*: float ## mean include count over ALL clauses (incl. empty)
proc tmMakeLiterals(input: TmBinaryVector): array[TM_N_LITERALS, uint8] = proc tmMakeLiterals(input: TmBinaryVector): array[TM_N_LITERALS, uint8] =
for i in 0..<TM_N_IN: for i in 0..<TM_N_IN:
result[i] = input[i] result[i] = input[i]
result[i + TM_N_IN] = 1'u8 - input[i] result[i + TM_N_IN] = 1'u8 - input[i]
proc tmEvalClause(net: TmNet, outIdx, clause: int, proc tmEvalClause(net: TmNet, outIdx, clause: int,
lits: array[TM_N_LITERALS, uint8]): uint8 = lits: array[TM_N_LITERALS, uint8],
learning = false): uint8 =
var hasIncluded = false var hasIncluded = false
for lit in 0..<TM_N_LITERALS: for lit in 0..<TM_N_LITERALS:
let s = net.states[tmStateIdx(outIdx, clause, lit)] let s = net.states[tmStateIdx(outIdx, clause, lit)]
if s > 0: if s > 0:
hasIncluded = true hasIncluded = true
if lits[lit] == 0: return 0'u8 if lits[lit] == 0: return 0'u8
return if hasIncluded: 1'u8 else: 0'u8 if hasIncluded: return 1'u8
# Granmo Eq. 6: an all-Exclude clause evaluates to 1 during LEARNING (the
# empty conjunction is vacuously true) and 0 during classification. The
# learning value is what bootstraps the clauses: at initialisation every
# automaton sits at the Exclude boundary, and with the corrected Type I rule
# (c=0, lk=1 -> toward Exclude) a classification-only 0 would deadlock every
# clause at empty forever.
return if learning: 1'u8 else: 0'u8
proc tmForwardWithCache(net: TmNet, input: TmBinaryVector, proc tmForwardWithCache(net: TmNet, input: TmBinaryVector,
cache: var TmClauseCache): (float, float) = cache: var TmClauseCache): (float, float, float, float) =
## Returns (correctionX, correctionY, voteX, voteY). `cache` receives the
## clause outputs under LEARNING semantics (empty clause = 1) for tmLearnOne;
## the returned correction and votes use classification semantics (empty = 0).
let lits = tmMakeLiterals(input) let lits = tmMakeLiterals(input)
var vx = 0.0; var vy = 0.0 var vx = 0.0; var vy = 0.0
for c in 0..<TM_N_CLAUSES: for c in 0..<TM_N_CLAUSES:
let o = tmEvalClause(net, 0, c, lits) let o = tmEvalClause(net, 0, c, lits)
cache[c] = o cache[c] = tmEvalClause(net, 0, c, lits, learning = true)
vx += tmPolarity(c) * float(o) vx += tmPolarity(c) * float(o)
for c in 0..<TM_N_CLAUSES: for c in 0..<TM_N_CLAUSES:
let o = tmEvalClause(net, 1, c, lits) let o = tmEvalClause(net, 1, c, lits)
cache[TM_N_CLAUSES + c] = o cache[TM_N_CLAUSES + c] = tmEvalClause(net, 1, c, lits, learning = true)
vy += tmPolarity(c) * float(o) vy += tmPolarity(c) * float(o)
vx = clamp(vx, -TM_T, TM_T) vx = clamp(vx, -TM_T, TM_T)
vy = clamp(vy, -TM_T, TM_T) vy = clamp(vy, -TM_T, TM_T)
(vx / TM_T * TM_RESID_MAX, vy / TM_T * TM_RESID_MAX) (vx / TM_T * TM_RESID_MAX, vy / TM_T * TM_RESID_MAX, vx, vy)
proc tmLearnOne(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8], proc tmLearnOne(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8],
cache: TmClauseCache, residual: float) = cache: TmClauseCache, vote: float, residual: float) =
var vote = 0.0 ## One faithful Granmo Table 2/3 update against a continuous residual target.
for c in 0..<TM_N_CLAUSES: ##
vote += tmPolarity(c) * float(cache[outIdx * TM_N_CLAUSES + c]) ## `cache` holds clause outputs under LEARNING semantics (empty = 1, Eq. 6).
vote = clamp(vote, -TM_T, TM_T) ## `vote` is the classification-semantics clause sum at prediction time,
## already clamped to [-TM_T, TM_T]. `residual` is the correction target
## delta = actual - linear baseline (see onResult), NOT actual - prediction.
##
## Regression adaptation: the "label" direction is d = sign(residual -
## predicted), i.e. which way the correction must move. The resource
## allocation of Granmo Eq. 8-11 collapses to a single probability
## p = (T - d*clip(v,-T,T)) / (2T), applied both to the aligned clauses
## (Type I) and the opposed ones (Type II), exactly as Algorithm 1 lines 11-22.
let predicted = vote / TM_T * TM_RESID_MAX let predicted = vote / TM_T * TM_RESID_MAX
let error = residual - predicted let error = residual - predicted
let pFeedback = min(1.0, abs(error) / (2.0 * TM_RESID_MAX)) let d = if error > 0.0: 1.0 elif error < 0.0: -1.0 else: return
let pFeedback = (TM_T - d * vote) / (2.0 * TM_T)
if pFeedback <= 0.0: return
for c in 0..<TM_N_CLAUSES: for c in 0..<TM_N_CLAUSES:
if pFeedback <= 0.0: continue
if rand(1.0) >= pFeedback: continue if rand(1.0) >= pFeedback: continue
let pol = tmPolarity(c) let pol = tmPolarity(c)
let cOut = cache[outIdx * TM_N_CLAUSES + c] let cOut = cache[outIdx * TM_N_CLAUSES + c]
if pol * d > 0.0:
if (error > 0.0 and pol > 0.0) or (error < 0.0 and pol < 0.0): # Type I (Table 2), collapsed to the resulting state move:
# Type I / Ib feedback # c=1, lk=1 -> +1 (toward Include) w.p. (s-1)/s
# c=0, lk=1 -> -1 (toward Exclude) w.p. 1/s <- the missing counter-force
# lk=0 -> -1 (toward Exclude) w.p. 1/s
for lit in 0..<TM_N_LITERALS: for lit in 0..<TM_N_LITERALS:
let si = tmStateIdx(outIdx, c, lit) let si = tmStateIdx(outIdx, c, lit)
var st = int(net.states[si]) var st = int(net.states[si])
if lits[lit] == 1'u8: if lits[lit] == 1'u8:
if rand(1.0) < (TM_S - 1.0) / TM_S: st = min(st + 1, TM_N_STATES) if cOut == 1'u8:
if rand(1.0) < (TM_S - 1.0) / TM_S: st = min(st + 1, TM_N_STATES)
else:
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_N_STATES)
else: else:
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_N_STATES) if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_N_STATES)
net.states[si] = int16(st) net.states[si] = int16(st)
else: else:
# Type II: shrink false literals in include range # Type II (Table 3): the only non-Inaction cell is
# c=1, lk=0, Exclude action -> Penalty -> toward Include.
# (The old code penalised INCLUDED false literals, which is unreachable
# when c=1 and the wrong direction.)
if cOut == 1'u8: if cOut == 1'u8:
for lit in 0..<TM_N_LITERALS: for lit in 0..<TM_N_LITERALS:
if lits[lit] == 0'u8: if lits[lit] == 0'u8:
let si = tmStateIdx(outIdx, c, lit) let si = tmStateIdx(outIdx, c, lit)
var st = int(net.states[si]) if net.states[si] <= 0:
if st > 0: net.states[si] = int16(min(int(net.states[si]) + 1, TM_N_STATES))
net.states[si] = int16(max(st - 1, -TM_N_STATES))
# ── TsetlinGun public type ──────────────────────────────────────────────────── # ── TsetlinGun public type ────────────────────────────────────────────────────
@@ -200,9 +242,11 @@ type
TmTrace = object TmTrace = object
fireTick: int # key part: tick the bullet was fired fireTick: int # key part: tick the bullet was fired
powerBin: int # key part: power bin the bullet belonged to powerBin: int # key part: power bin the bullet belonged to
predX, predY: float # stored prediction, for the directional residual predX, predY: float # stored prediction (for debug/provenance)
linearX, linearY: float # baseline the correction was added to (fix 4)
voteX, voteY: float # classification clause sum at fire time
input: TmBinaryVector input: TmBinaryVector
cache: TmClauseCache cache: TmClauseCache # LEARNING clause outputs (empty = 1)
alive: bool alive: bool
TsetlinGun* = object TsetlinGun* = object
@@ -214,6 +258,13 @@ type
shotCount: int ## total onResult calls received shotCount: int ## total onResult calls received
trainedShots*: int ## onResult calls that found and trained their exact trace trainedShots*: int ## onResult calls that found and trained their exact trace
traceMisses*: int ## onResult calls whose trace was gone (integrity counter) traceMisses*: int ## onResult calls whose trace was gone (integrity counter)
# Instrumentation: is the TM actually producing a nonzero correction? Before
# the feedback-rule fix vHits was effectively identical to Linear because
# the clauses over-specified to a mean of ~714 included literals and almost
# never fired (correction was nonzero on only 8 of 764 predict calls).
predictCalls*: int ## predict() invocations (4/tick in the harness)
correctionsNonzero*: int ## predict() calls whose |cx|+|cy| > 1e-9
lastCx*, lastCy*: float ## last correction emitted (debugging/assertions)
debugGraphics*: bool debugGraphics*: bool
proc tmBinForSpeed(spd: float): int {.inline.} = proc tmBinForSpeed(spd: float): int {.inline.} =
@@ -239,6 +290,28 @@ proc initTsetlinGun*(): TsetlinGun =
proc isWarmedUp*(g: TsetlinGun): bool {.inline.} = proc isWarmedUp*(g: TsetlinGun): bool {.inline.} =
g.bufferCount >= TM_WINDOW_SIZE g.bufferCount >= TM_WINDOW_SIZE
proc tmClauseStats*(g: TsetlinGun): TmClauseStats =
## Sample the current include-count distribution of the automata. O(outputs ×
## clauses × literals); call on demand, not per-tick.
result.nClauses = TM_N_OUT * TM_N_CLAUSES
result.minIncluded = high(int)
var total = 0
for o in 0..<TM_N_OUT:
for c in 0..<TM_N_CLAUSES:
var inc = 0
for lit in 0..<TM_N_LITERALS:
if g.net.states[tmStateIdx(o, c, lit)] > 0: inc += 1
total += inc
if inc > 0:
inc result.nActive
result.minIncluded = min(result.minIncluded, inc)
result.maxIncluded = max(result.maxIncluded, inc)
if result.nActive > 0:
result.meanIncluded = total.float / result.nActive.float
else:
result.minIncluded = 0
result.meanIncludedAll = total.float / result.nClauses.float
proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPrediction = proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPrediction =
# Encode current frame and push into window # Encode current frame and push into window
let dist = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY) let dist = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY)
@@ -248,7 +321,7 @@ proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPred
bearing, dist, state.enemySpeed, state.enemyHeading, bearing, dist, state.enemySpeed, state.enemyHeading,
state.arenaHeight - state.enemyY, state.enemyY, state.arenaHeight - state.enemyY, state.enemyY,
state.arenaWidth - state.enemyX, state.enemyX, state.arenaWidth - state.enemyX, state.enemyX,
state.selfEnergy, # use self energy as proxy (enemy energy not in WorldState) state.enemyEnergy, # fix 6: enemy energy, not a duplicate of self energy
) )
# Shift window: index 0 = newest. Do this at most once per tick — the harness # Shift window: index 0 = newest. Do this at most once per tick — the harness
# calls predict() 4-5x/tick (once per power bin), which used to shift the # calls predict() 4-5x/tick (once per power bin), which used to shift the
@@ -280,7 +353,11 @@ proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPred
let vec = tmEncodeFullVector(g.frameBuffer, selfState) let vec = tmEncodeFullVector(g.frameBuffer, selfState)
var cache: TmClauseCache var cache: TmClauseCache
let (cx, cy) = tmForwardWithCache(g.net, vec, cache) let (cx, cy, vx, vy) = tmForwardWithCache(g.net, vec, cache)
inc g.predictCalls
g.lastCx = cx
g.lastCy = cy
if abs(cx) + abs(cy) > 1e-9: inc g.correctionsNonzero
let predX = clamp(linearX + cx, 0.0, state.arenaWidth) let predX = clamp(linearX + cx, 0.0, state.arenaWidth)
let predY = clamp(linearY + cy, 0.0, state.arenaHeight) let predY = clamp(linearY + cy, 0.0, state.arenaHeight)
@@ -295,6 +372,10 @@ proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPred
powerBin: binIdx, powerBin: binIdx,
predX: predX, predX: predX,
predY: predY, predY: predY,
linearX: linearX,
linearY: linearY,
voteX: vx,
voteY: vy,
input: vec, input: vec,
cache: cache, cache: cache,
alive: true, alive: true,
@@ -319,14 +400,16 @@ proc onResult*(g: var TsetlinGun, e: FeedbackEvent) =
inc g.traceMisses inc g.traceMisses
return return
# Directional residual: actual enemy pos minus our prediction # Fixes 4+5: train on delta = actual - LINEAR baseline. The old code trained
# On hit residual is 0 (we were right); on miss we push toward actual position. # on actual - predX = delta - cx, so tmLearnOne's error became delta - 2*cx and
let rx = if e.hit: 0.0 else: clamp(e.actualX - t.predX, -TM_RESID_MAX, TM_RESID_MAX) # the fixed point was cx = delta/2. Hits are NOT zeroed: a hit means
let ry = if e.hit: 0.0 else: clamp(e.actualY - t.predY, -TM_RESID_MAX, TM_RESID_MAX) # |miss| < BotRadius, not residual == 0.
let dx = clamp(e.actualX - t.linearX, -TM_RESID_MAX, TM_RESID_MAX)
let dy = clamp(e.actualY - t.linearY, -TM_RESID_MAX, TM_RESID_MAX)
let lits = tmMakeLiterals(t.input) let lits = tmMakeLiterals(t.input)
g.net.tmLearnOne(0, lits, t.cache, rx) g.net.tmLearnOne(0, lits, t.cache, t.voteX, dx)
g.net.tmLearnOne(1, lits, t.cache, ry) g.net.tmLearnOne(1, lits, t.cache, t.voteY, dy)
when DebugTM: when DebugTM:
echo fmt"[tm-dbg] shot={g.shotCount} tick={e.fireTick} bin={binIdx} miss={e.missDistance:.1f}px predicted=({t.predX:.0f},{t.predY:.0f}) actual=({e.actualX:.0f},{e.actualY:.0f}) rx={rx:.1f} ry={ry:.1f} hit={e.hit}" echo fmt"[tm-dbg] shot={g.shotCount} tick={e.fireTick} bin={binIdx} miss={e.missDistance:.1f}px predicted=({t.predX:.0f},{t.predY:.0f}) actual=({e.actualX:.0f},{e.actualY:.0f}) dx={dx:.1f} dy={dy:.1f} hit={e.hit}"
t.alive = false t.alive = false
inc g.trainedShots inc g.trainedShots
+17
View File
@@ -40,3 +40,20 @@ proc buildAllGunDrivers*(seed = -1): seq[GunDriver] =
makeDriver("DecayGF", initDecayGFGun()), makeDriver("DecayGF", initDecayGFGun()),
makeDriver("KNN", initKNNGun()), makeDriver("KNN", initKNNGun()),
] ]
proc makeTsetlinDriver*(seed = -1): tuple[driver: GunDriver, gun: ref TsetlinGun] =
## Same as makeDriver("Tsetlin", ...) but keeps a handle to the concrete gun,
## so a test can read its clause-sparsity / correction instrumentation after a
## replay. `makeDriver` heap-boxes a copy internally and drops the handle.
let g = new(TsetlinGun)
g[] = initTsetlinGun()
if seed >= 0:
randomize(seed)
result.gun = g
result.driver = GunDriver(
name: "Tsetlin",
predictCb: proc(state: WorldState, bulletSpeed: float): GunPrediction =
g[].predict(state, bulletSpeed),
resultCb: proc(e: FeedbackEvent) = g[].onResult(e),
readyCb: proc(): bool = g[].isWarmedUp(),
)
@@ -0,0 +1,504 @@
## 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()
+78
View File
@@ -0,0 +1,78 @@
## L1 / L3 tests for the repaired Tsetlin gun.
##
## L1 (clause sparsity + nonzero correction): before the fix the gun's Type I
## feedback never conditioned on the clause output, so every frequently-true
## literal ratcheted toward Include with no counter-force. Measured on
## `energy-threshold-turner` the clauses saturated at mean ≈714 included
## literals/clause (max 747, 100/100 active) and the correction was nonzero on
## only 8 of 764 predict calls — i.e. `Tsetlin.vHits` was effectively `Linear`.
##
## L3 (divergence): with the corrected Granmo Table 2/3 rules, the learning core
## is seeded (`makeTsetlinDriver(seed=1)`) so these numbers are reproducible.
##
## Run: nim c -r common_libs/tests/test_tsetlin_gun.nim
import std/strformat
import gun_harness/offline_range
import range_guns
import guns/tsetlin
import guns/linear
var failures = 0
proc check(name: string, ok: bool) =
if ok: echo "PASS: ", name
else: echo "FAIL: ", name; inc failures
proc main() =
# ── L1 ────────────────────────────────────────────────────────────────────
block:
let fx = synthesizeEnergyThresholdTurner()
let (drv, gun) = makeTsetlinDriver(seed = 1)
let reps = replayFixture(fx, @[drv])
let st = gun[].tmClauseStats()
echo &"L1 energy-threshold-turner (ticks={fx.states.len}):"
echo &" clauses: mean={st.meanIncluded:.1f} max={st.maxIncluded} " &
&"active={st.nActive}/{st.nClauses} meanAll={st.meanIncludedAll:.1f}"
echo &" correction nonzero on {gun[].correctionsNonzero}/{gun[].predictCalls} predict calls"
echo &" trainedShots={gun[].trainedShots} traceMisses={gun[].traceMisses} hits={reps[0].hits}/{reps[0].shots}"
check "L1: mean included literals/clause fell to a sparse regime (< 40; was ~714)", st.meanIncluded < 40.0
check "L1: at least some clauses are active (the TM is doing something)", st.nActive > 0
check "L1: the learned correction is nonzero", gun[].correctionsNonzero > 0
check "L1: correction is nonzero on most predicts", gun[].correctionsNonzero * 2 > gun[].predictCalls
check "L1: trace pairing intact (trainedShots > 0, traceMisses == 0)",
gun[].trainedShots > 0 and gun[].traceMisses == 0
echo ""
# ── L3 ────────────────────────────────────────────────────────────────────
block:
var diverged = false
var energyT, energyL: GunReport
echo "L3 divergence over synthetic range fixtures (seed=1):"
for name in SyntheticFixtureNames:
let fx = synthesizeByName(name)
let (drv, gun) = makeTsetlinDriver(seed = 1)
let reps = replayFixture(fx, @[drv, makeDriver("Linear", LinearGun())])
let t = reps[0]
let l = reps[1]
echo &" {name:<26} Tsetlin {t.hits:>4}/{t.shots:<4} Linear {l.hits:>4}/{l.shots:<4} " &
&"nzCorr={gun[].correctionsNonzero}/{gun[].predictCalls}"
if t.hits != l.hits: diverged = true
if name == "energy-threshold-turner":
energyT = t
energyL = l
echo ""
echo &"L3 energy-threshold-turner: Tsetlin {energyT.hits}/{energyT.shots} " &
&"({energyT.hitRate()*100:.1f}%) vs Linear {energyL.hits}/{energyL.shots} " &
&"({energyL.hitRate()*100:.1f}%)"
check "L3: Tsetlin.vHits diverges from Linear.vHits on at least one range fixture", diverged
check "L3: Tsetlin.vHits != Linear.vHits on energy-threshold-turner",
energyT.hits != energyL.hits
if failures > 0:
echo "\n", failures, " check(s) FAILED"
quit(1)
echo "\nAll Tsetlin-gun checks passed."
when isMainModule:
main()
+69
View File
@@ -0,0 +1,69 @@
## L4: one live gauntlet proving `Tsetlin.vHits != Linear.vHits`.
##
## Runs ModularBot against a chosen adversary through the real Tank Royale server
## and reads the per-gun virtual-bullet window from /tmp/gun_stats.jsonl (written
## by ModularBot.onRoundEnded, independent of the RecordWorldState flag).
##
## Usage: nim c -r common_libs/tests/check_tsetlin_live.nim [AdversaryName] [rounds]
## default adversary RandomMover, 3 rounds.
## Skips cleanly when the TR JARs are unavailable.
import std/[os, strformat, json, strutils]
import test_framework/test_framework
const
repoRoot = currentSourcePath().parentDir.parentDir.parentDir
modularBotDir = repoRoot / "ModularBot_garage"
adversariesDir = repoRoot / "common_libs" / "test_framework" / "adversaries"
statsPath = "/tmp/gun_stats.jsonl"
proc main() =
let adversary = if paramCount() >= 1: paramStr(1) else: "RandomMover"
let rounds = if paramCount() >= 2: parseInt(paramStr(2)) else: 3
let advDir = adversariesDir / adversary
if not dirExists(advDir):
echo "Skipping: adversary not found: ", advDir
quit(0)
if getEnv("TR_SERVER_JAR", "").len == 0 and
not fileExists("/home/davide/Projects/tank-royale/server/build/libs/robocode-tankroyale-server-0.35.5-all.jar"):
echo "Skipping: TR server JAR not found"
quit(0)
if fileExists(statsPath): removeFile(statsPath)
echo &"live battle: ModularBot vs {adversary}, {rounds} round(s)"
let battle = runBattle(@[modularBotDir, advDir], rounds = rounds,
timeout = 300000, maxSpeed = true)
for r in battle.results:
echo &" {r.name:<14} rank={r.rank} score={r.totalScore}"
var tHits, tShots, lHits, lShots, vDrop, vStarve = 0
var rows = 0
if fileExists(statsPath):
for line in lines(statsPath):
let s = line.strip()
if s.len == 0: continue
let d = parseJson(s)
if not d.hasKey("guns"): continue
inc rows
for g in d["guns"]:
case g["name"].getStr()
of "Tsetlin": tHits += g["vHits"].getInt(); tShots += g["vShots"].getInt()
of "Linear": lHits += g["vHits"].getInt(); lShots += g["vShots"].getInt()
else: discard
vDrop = d["vDropped"].getInt()
vStarve = d["vStarved"].getInt()
echo ""
echo &"rounds logged: {rows}"
echo &"Tsetlin vHits={tHits}/{tShots} Linear vHits={lHits}/{lShots}"
echo &"vDropped={vDrop} vStarved={vStarve}"
if tHits != lHits:
echo "VERDICT: DIVERGED — Tsetlin.vHits != Linear.vHits"
else:
echo "VERDICT: IDENTICAL — the fix did not change the live outcome for this opponent"
if vDrop != 0 or vStarve != 0:
echo "WARNING: vDropped/vStarved are non-zero"
quit(1)
when isMainModule:
main()