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:
@@ -41,3 +41,12 @@ worktrees/
|
||||
# Test fixtures are data, not logs - the *.jsonl rule above was written for
|
||||
# training logs and silently excluded the entire gun-range fixture set.
|
||||
!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
@@ -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) ─────
|
||||
|
||||
const
|
||||
TM_FRAME_BITS = 83
|
||||
TM_SELF_BITS = 40
|
||||
TM_WINDOW_SIZE = 10
|
||||
TM_FRAME_BITS* = 83
|
||||
TM_SELF_BITS* = 40
|
||||
TM_WINDOW_SIZE* = 10
|
||||
TM_TOTAL_BITS* = TM_FRAME_BITS * TM_WINDOW_SIZE + TM_SELF_BITS # 870
|
||||
TM_MAX_DISTANCE = 1414.0 # diagonal of 1000x1000 arena
|
||||
|
||||
type
|
||||
TmBinaryVector = array[TM_TOTAL_BITS, uint8]
|
||||
TmFrameEncoded = array[TM_FRAME_BITS, uint8]
|
||||
TmSelfEncoded = array[TM_SELF_BITS, uint8]
|
||||
TmBinaryVector* = array[TM_TOTAL_BITS, uint8]
|
||||
TmFrameEncoded* = array[TM_FRAME_BITS, uint8]
|
||||
TmSelfEncoded* = array[TM_SELF_BITS, uint8]
|
||||
|
||||
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:
|
||||
result[bits - 1 - i] = uint8((gray shr i) and 1)
|
||||
|
||||
proc tmEncodeFrame(bearing, distance, velocity, heading,
|
||||
wallN, wallS, wallE, wallW, energy: float): TmFrameEncoded =
|
||||
proc tmEncodeFrame*(bearing, distance, velocity, heading,
|
||||
wallN, wallS, wallE, wallW, energy: float): TmFrameEncoded =
|
||||
var offset = 0
|
||||
|
||||
# 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)
|
||||
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
|
||||
for wall in [wallN, wallS, wallE, wallW]:
|
||||
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
|
||||
result[offset] = if canFire: 1'u8 else: 0'u8
|
||||
|
||||
proc tmEncodeFullVector(window: array[TM_WINDOW_SIZE, TmFrameEncoded],
|
||||
self: TmSelfEncoded): TmBinaryVector =
|
||||
proc tmEncodeFullVector*(window: array[TM_WINDOW_SIZE, TmFrameEncoded],
|
||||
self: TmSelfEncoded): TmBinaryVector =
|
||||
var offset = 0
|
||||
for i in 0..<TM_WINDOW_SIZE:
|
||||
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.} =
|
||||
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] =
|
||||
for i in 0..<TM_N_IN:
|
||||
result[i] = input[i]
|
||||
result[i + TM_N_IN] = 1'u8 - input[i]
|
||||
|
||||
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
|
||||
for lit in 0..<TM_N_LITERALS:
|
||||
let s = net.states[tmStateIdx(outIdx, clause, lit)]
|
||||
if s > 0:
|
||||
hasIncluded = true
|
||||
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,
|
||||
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)
|
||||
var vx = 0.0; var vy = 0.0
|
||||
for c in 0..<TM_N_CLAUSES:
|
||||
let o = tmEvalClause(net, 0, c, lits)
|
||||
cache[c] = o
|
||||
cache[c] = tmEvalClause(net, 0, c, lits, learning = true)
|
||||
vx += tmPolarity(c) * float(o)
|
||||
for c in 0..<TM_N_CLAUSES:
|
||||
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)
|
||||
vx = clamp(vx, -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],
|
||||
cache: TmClauseCache, residual: float) =
|
||||
var vote = 0.0
|
||||
for c in 0..<TM_N_CLAUSES:
|
||||
vote += tmPolarity(c) * float(cache[outIdx * TM_N_CLAUSES + c])
|
||||
vote = clamp(vote, -TM_T, TM_T)
|
||||
cache: TmClauseCache, vote: float, residual: float) =
|
||||
## One faithful Granmo Table 2/3 update against a continuous residual target.
|
||||
##
|
||||
## `cache` holds clause outputs under LEARNING semantics (empty = 1, Eq. 6).
|
||||
## `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 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:
|
||||
if pFeedback <= 0.0: continue
|
||||
if rand(1.0) >= pFeedback: continue
|
||||
let pol = tmPolarity(c)
|
||||
let cOut = cache[outIdx * TM_N_CLAUSES + c]
|
||||
|
||||
if (error > 0.0 and pol > 0.0) or (error < 0.0 and pol < 0.0):
|
||||
# Type I / Ib feedback
|
||||
if pol * d > 0.0:
|
||||
# Type I (Table 2), collapsed to the resulting state move:
|
||||
# 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:
|
||||
let si = tmStateIdx(outIdx, c, lit)
|
||||
var st = int(net.states[si])
|
||||
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:
|
||||
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)
|
||||
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:
|
||||
for lit in 0..<TM_N_LITERALS:
|
||||
if lits[lit] == 0'u8:
|
||||
let si = tmStateIdx(outIdx, c, lit)
|
||||
var st = int(net.states[si])
|
||||
if st > 0:
|
||||
net.states[si] = int16(max(st - 1, -TM_N_STATES))
|
||||
if net.states[si] <= 0:
|
||||
net.states[si] = int16(min(int(net.states[si]) + 1, TM_N_STATES))
|
||||
|
||||
# ── TsetlinGun public type ────────────────────────────────────────────────────
|
||||
|
||||
@@ -200,9 +242,11 @@ type
|
||||
TmTrace = object
|
||||
fireTick: int # key part: tick the bullet was fired
|
||||
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
|
||||
cache: TmClauseCache
|
||||
cache: TmClauseCache # LEARNING clause outputs (empty = 1)
|
||||
alive: bool
|
||||
|
||||
TsetlinGun* = object
|
||||
@@ -214,6 +258,13 @@ type
|
||||
shotCount: int ## total onResult calls received
|
||||
trainedShots*: int ## onResult calls that found and trained their exact trace
|
||||
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
|
||||
|
||||
proc tmBinForSpeed(spd: float): int {.inline.} =
|
||||
@@ -239,6 +290,28 @@ proc initTsetlinGun*(): TsetlinGun =
|
||||
proc isWarmedUp*(g: TsetlinGun): bool {.inline.} =
|
||||
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 =
|
||||
# Encode current frame and push into window
|
||||
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,
|
||||
state.arenaHeight - state.enemyY, state.enemyY,
|
||||
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
|
||||
# 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)
|
||||
|
||||
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 predY = clamp(linearY + cy, 0.0, state.arenaHeight)
|
||||
@@ -295,6 +372,10 @@ proc predict*(g: var TsetlinGun, state: WorldState, bulletSpeed: float): GunPred
|
||||
powerBin: binIdx,
|
||||
predX: predX,
|
||||
predY: predY,
|
||||
linearX: linearX,
|
||||
linearY: linearY,
|
||||
voteX: vx,
|
||||
voteY: vy,
|
||||
input: vec,
|
||||
cache: cache,
|
||||
alive: true,
|
||||
@@ -319,14 +400,16 @@ proc onResult*(g: var TsetlinGun, e: FeedbackEvent) =
|
||||
inc g.traceMisses
|
||||
return
|
||||
|
||||
# Directional residual: actual enemy pos minus our prediction
|
||||
# On hit residual is 0 (we were right); on miss we push toward actual position.
|
||||
let rx = if e.hit: 0.0 else: clamp(e.actualX - t.predX, -TM_RESID_MAX, TM_RESID_MAX)
|
||||
let ry = if e.hit: 0.0 else: clamp(e.actualY - t.predY, -TM_RESID_MAX, TM_RESID_MAX)
|
||||
# Fixes 4+5: train on delta = actual - LINEAR baseline. The old code trained
|
||||
# on actual - predX = delta - cx, so tmLearnOne's error became delta - 2*cx and
|
||||
# the fixed point was cx = delta/2. Hits are NOT zeroed: a hit means
|
||||
# |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)
|
||||
g.net.tmLearnOne(0, lits, t.cache, rx)
|
||||
g.net.tmLearnOne(1, lits, t.cache, ry)
|
||||
g.net.tmLearnOne(0, lits, t.cache, t.voteX, dx)
|
||||
g.net.tmLearnOne(1, lits, t.cache, t.voteY, dy)
|
||||
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
|
||||
inc g.trainedShots
|
||||
|
||||
@@ -40,3 +40,20 @@ proc buildAllGunDrivers*(seed = -1): seq[GunDriver] =
|
||||
makeDriver("DecayGF", initDecayGFGun()),
|
||||
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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user