From f9f8d846711b79ddc4c50ace911025a6a78dfa06 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Tue, 22 Sep 2026 21:24:27 +0200 Subject: [PATCH] TM diagnostics kit: VALIDATED (finds a known dead input), and it found a real bug Built `common_libs/tm_diag/` as a first-class offline diagnostics kit for Tsetlin work, BEFORE writing the new gun - because we hit two data problems tonight that no amount of reading the TM's clauses would have revealed (a 38.8% majority answer, and 36-58% mislabelled training samples). WHAT IT PROVIDES - `feature_spec.nim`: a NAMED feature container, so a learned clause prints as a sentence (`IF near-wall AND bullet-dead-on AND turn-left(t-2) THEN class=3`) instead of "feature 17". Includes the 49-bit draft spec from the design session and the shipped 40-bit encoding. - `tm_core.nim`: a compact deterministic Granmo multiclass TM with an INTROSPECTABLE clause layout (mirrors the tm_pattern core). - `diagnostics.nim`, six groups: (1) pre-flight DATA checks + shuffled-label control, (2) clause introspection (readable dump, per-clause vote counts, empty and never-fired clauses, length distribution, per-class balance), (3) per-feature contribution with an explicit DEAD-INPUT LIST and a ranked most-valuable list, (4) accuracy vs the majority baseline with per-class precision/recall and pred-majority share, (5) learning curve, (6) ablation hooks (drop a block / scramble a bit). === TASK 3: THE VALIDATION THAT GATES EVERYTHING - PASSED WITH NUMBERS === A diagnostic we never checked is worthless, so the kit was tested on a synthetic set with a PLANTED RULE (class2 = A and B, class1 = A and not B, class0 = not A), a deliberately IRRELEVANT block (US, 9 bits) and a PURE-NOISE bit (17). - majority baseline 60.63% (class0); over-30% correctly flagged - **the planted rule is recovered EXACTLY** via `necessaryLiterals`: class0 IF NOT dist-wall<50 | class1 IF dist-wall<50 AND NOT lat DEAD-ON | class2 IF dist-wall<50 AND lat DEAD-ON - **DEAD-INPUT LIST = all 9 US bits AND the noise bit 17**, while the planted bits 0 and 45 are correctly NOT listed - top contributors: bit0 w=1241.7, bit45 w=583.3, then 49.8 - a 12-25x gap, so the relevant bits are unmistakable - **ABLATION: drop WALLS -39.47pp, drop BULLETS -19.33pp, drop US 0.00pp**, scramble A -42.00pp, scramble the noise bit 0.00pp - shuffled-label control 60.40% vs majority 60.63% = -0.23pp -> no leak So the kit reliably finds a known dead input and a known relevant one. === TASK 4: THE REAL READING, AND A BUG IN THE SHIPPED GUN === `tm_pattern` GF head, 6 DrussGT fixtures, pooled 250,745 samples: - label balance c2 = **34.4%** (majority-heavy, flagged); accuracy **35.72%** vs majority **34.24%** -> margin **+1.48pp**. On `tr_drussgt_vs_crazy` it is BELOW majority (33.81% vs 37.72%, -3.92pp). - 200 clauses: **27 empty, 45 never fired**, mean length 19.17, max 57. The majority class is starved (class2: 22 non-empty, 18 empty, only 2 positive fired). Class4 fires 11-24-literal clauses -> memorisation signature. - **REPRESENTATION BUG FOUND (reported, not silently fixed):** `tmBuildBits` writes only 38 raw bits into `var bits: array[TM_NBITS=40, uint8]` - bits 38 and 39 are NEVER ASSIGNED, so they are always 0 and their negated literals are always 1. The kit's `constantInputs` confirms 38/39 are constant, and **`UNUSED-38`/`39` rank #6 and #8 in the most-valuable-inputs list** - i.e. the model's highest-usage inputs are information-free. That is a representation bug, not a display artefact, and it is a concrete mechanism for part of the poor learning. DRAFT ENCODING CHECKED: the 49-bit draft is arithmetically consistent (4+4=8 walls, 6+3=9 us, 3+5+3+3+3+3=20 motion, 5+7=12 bullets = 49). No draft inconsistency. Guards: test_tm_diag 48 (new), diag_synthetic 17 (new), test_gun_harness 39, test_vbullet_metric 11, test_power_selection 3, test_adaptive_radar 41, test_tfil_ring_weights 24, test_power_policy 26, test_ram_decision 40, test_rack_membership 48, test_selector_tiebreak 19, test_tm_pattern_registration 20, test_vbullet_admit_gate 12, acceptance_offline_vs_online 12/12; tm_pattern_learning passes. The tm_pattern hook is additive and default-OFF (no behaviour change). NOT YET INCLUDED (the automata metrics discussed for the next step): per-clause automata settledness (distance from the flip point), clause diversity (pairwise overlap), literal-set churn over time, and cross-clause vote disagreement. The kit has clause-level diagnostics but not the automata-state ones. --- common_libs/guns/tm_pattern.nim | 27 + common_libs/tests/diag_synthetic.nim | 201 ++++++ common_libs/tests/diag_tm_pattern_offline.nim | 181 ++++++ common_libs/tests/test_tm_diag.nim | 165 +++++ common_libs/tm_diag/README.md | 124 ++++ common_libs/tm_diag/diagnostics.nim | 589 ++++++++++++++++++ common_libs/tm_diag/feature_spec.nim | 148 +++++ common_libs/tm_diag/tm_core.nim | 173 +++++ 8 files changed, 1608 insertions(+) create mode 100644 common_libs/tests/diag_synthetic.nim create mode 100644 common_libs/tests/diag_tm_pattern_offline.nim create mode 100644 common_libs/tests/test_tm_diag.nim create mode 100644 common_libs/tm_diag/README.md create mode 100644 common_libs/tm_diag/diagnostics.nim create mode 100644 common_libs/tm_diag/feature_spec.nim create mode 100644 common_libs/tm_diag/tm_core.nim diff --git a/common_libs/guns/tm_pattern.nim b/common_libs/guns/tm_pattern.nim index 2581d93..813781b 100644 --- a/common_libs/guns/tm_pattern.nim +++ b/common_libs/guns/tm_pattern.nim @@ -104,6 +104,14 @@ type TmBits* = array[TM_NLITS, uint8] + TmDiagSample* = object + ## DEFAULT-OFF diagnostics capture (common_libs/tm_diag). Populated only + ## when `TmPatternGun.diagCapture` is true; the literal vector is exactly + ## what the TM trained on, plus the eventual label. + lits*: TmBits + label*: int + order*: int ## fire tick (approximate temporal order) + TmPatternTrace = object fireTick: int powerBin: int @@ -200,6 +208,9 @@ type shuffleLabels*: bool ## control: replace the computed GF label with a random class forceBase*: bool ## measurement: ignore the TM, emit the pure LinearGun base debugGraphics*: bool + # ── default-off diagnostics capture (common_libs/tm_diag) ── + diagCapture*: bool ## false by default; no behaviour when false + diagSamples*: seq[TmDiagSample] # ── TM core (Granmo Table 2/3, corrected resource allocation) ──────────────── @@ -303,6 +314,20 @@ proc initTmPatternGun*(): TmPatternGun = result.lastChosen = (TM_CLASSES - 1) div 2 randomize() result.debugGraphics = false + result.diagCapture = false + +proc exportTeams*(g: TmPatternGun): seq[seq[int16]] = + ## READ-ONLY view of the GF clause teams (common_libs/tm_diag introspection). + result = newSeq[seq[int16]](TM_CLASSES) + for c in 0..2 | A&!B->1 | !A->0" + + # ── GROUP 1: data checks ── + var labels = newSeq[int](train.len) + for i, s in train: labels[i] = s.label + let dc = dataChecks(labels, NClasses) + echo "\n## GROUP 1 data checks" + echo &"# n={dc.n} counts={dc.classCounts} shares={dc.classShares} " & + &"majority=class{dc.majorityClass} share={dc.majorityShare*100:.2f}% " & + &"overThreshold={dc.overThreshold}" + for f in dc.flags: echo "# FLAG ", f + # manual majority for the check + var manualMaj = 0 + for c in 0.. dc.classCounts[manualMaj]: manualMaj = c + check "majority class computed from the counts", + dc.majorityClass == manualMaj and + abs(dc.majorityShare - dc.classCounts[manualMaj].float / dc.n.float) < 1e-9 + + # ── train the reference TM ── + let tmpl = newMachine(NBits, NClasses, nClauses = 40, nStates = 64, + sValue = 3.0, seed = 1) + let m = trainModel(tmpl, train, epochs = 25, seed = 777) + let trainAcc = evalAcc(m, train) + let testAcc = evalAcc(m, eval) + echo &"\n## trained reference TM: trainAcc={trainAcc*100:.2f}% testAcc={testAcc*100:.2f}%" + + # ── GROUP 2: clause introspection ── + let infos = clauseInfo(m, eval, spec) + let summ = clauseSummary(infos) + echo "\n## GROUP 2 clause introspection" + echo &"# totalClauses={summ.totalClauses} empty={summ.emptyClauses} " & + &"nonEmpty={summ.nonEmpty} fired>=1={summ.firedAtLeastOnce} " & + &"neverFired={summ.neverFired} posFired={summ.posFired} negFired={summ.negFired} " & + &"meanLen={summ.meanLength:.2f} maxLen={summ.maxLength}" + echo "# lengthHist (idx=len): ", summ.lengthHist + echo "## top firing POSITIVE clauses per class (rule-encoding; polarity annotated)" + for cls in 0..2} {ranked[i].name:<22} appearances={ranked[i].appearances} weighted={ranked[i].weighted:.1f}" + echo "# DEAD-INPUT LIST (weighted < 5% of top, plus never-in-voting-clause): ", dead + echo "# never-used (strict appearances==0): ", neverUsed + echo "# constant inputs: ", consts + check "the planted relevant bits are the top-2 contributors", + ranked[0].bit in {BitA, BitB} and ranked[1].bit in {BitA, BitB} + # the whole US block must be dead + var usDead = true + for b in 8..16: + if b notin dead: usDead = false + check "the whole irrelevant US block is on the dead-input list", usDead + check "the pure-noise MOTION bit is on the dead-input list", NoiseBit in dead + check "no planted relevant bit is called dead", + BitA notin dead and BitB notin dead + + # ── GROUP 4: accuracy diagnostics ── + let ad = accuracyDiagnostics(m, eval) + echo "\n## GROUP 4 accuracy diagnostics" + echo &"# n={ad.n} correct={ad.correct} acc={ad.acc*100:.2f}% " & + &"majority=class{ad.majorityClass} baseline={ad.majorityBaseline*100:.2f}% " & + &"margin={ad.margin*100:+.2f}pp predMajorityShare={ad.predMajorityShare*100:.2f}%" + echo "# confusion [true][pred]:" + for c in 0.. 0.3 + check "uniform-ish predictions (not majority-spam): predMajorityShare < 0.7", + ad.predMajorityShare < 0.7 + + # ── GROUP 5: learning curve ── + let lc = learningCurve(tmpl, train, eval, nPoints = 10, seed = 31) + echo "\n## GROUP 5 learning curve" + echo &"# trend={lc.trend}" + for i in 0.. 0: + echo &"{failures} check(s) FAILED" + quit(1) + echo "All synthetic diagnostic validation checks passed." diff --git a/common_libs/tests/diag_tm_pattern_offline.nim b/common_libs/tests/diag_tm_pattern_offline.nim new file mode 100644 index 0000000..2938bf6 --- /dev/null +++ b/common_libs/tests/diag_tm_pattern_offline.nim @@ -0,0 +1,181 @@ +## TASK 4 — the tm_diag kit against the EXISTING guns/tm_pattern.nim GF head. +## +## Replays the committed DrussGT fixtures through the offline gun range with a +## FRESH, diagCapture-enabled TM per fixture (the live semantics: cold every +## battle, overfit within the battle), then reports: +## * pooled label balance + majority baseline vs the gun's own warm accuracy; +## * the DEAD-INPUT LIST for the actual 40-bit tm_pattern feature set; +## * the top firing positive clauses and the recovered necessary literals. +## +## Offline only. Fixtures are READ-ONLY. +## Run: nim c -r -d:release --path:common_libs common_libs/tests/diag_tm_pattern_offline.nim +## [--fixtures=a,b,c] [--maxsamples=N] + +import std/[os, strformat, strutils, algorithm, random] +import gun_harness/offline_range +import guns/tm_pattern +import tm_diag/diagnostics + +const repoRoot = currentSourcePath().parentDir.parentDir.parentDir +const fixturesDir = repoRoot / "tools" / "fixtures" + +proc toDiag(s: TmDiagSample): DiagSample = + result.lits = newSeq[uint8](TM_NLITS) + for i in 0.. rowSum[maj]: maj = c + let majShare = if total > 0: rowSum[maj].float / total.float else: 0.0 + let acc = if total > 0: correct.float / total.float else: 0.0 + echo &"# pooled warm accuracy = {correct}/{total} = {acc*100:.2f}%" + echo &"# majority class = {maj} share/baseline = {rowSum[maj]}/{total} = {majShare*100:.2f}%" + echo &"# margin = {(acc-majShare)*100:+.2f}pp pred-majority share = " & + &"{colSum[maj].float/max(1,total).float*100:.2f}%" + for c in 0.. 0: cm[c][c].float / rowSum[c].float else: 0.0 + let prec = if colSum[c] > 0: cm[c][c].float / colSum[c].float else: 0.0 + echo &"# class{c}: trueN={rowSum[c]:<6} predN={colSum[c]:<6} TP={cm[c][c]:<6} " & + &"recall={rec*100:5.1f}% precision={prec*100:5.1f}%" + +proc main() = + var names = @["drussgt_vs_crazy", "drussgt_vs_spinbot", "drussgt_vs_drussgt", + "tr_drussgt_vs_crazy", "tr_drussgt_vs_spinbot", + "tr_drussgt_vs_modularbot"] + var maxSamples = 40000 + for i in 1..paramCount(): + let a = paramStr(i) + if a.startsWith("--fixtures="): + names = a[11..^1].split(',') + elif a.startsWith("--maxsamples="): + maxSamples = parseInt(a[13..^1]) + + let spec = tmPatternSpec() + echo &"# tm_pattern GF head over DrussGT fixtures (spec nBits={spec.nBits}, " & + &"TM_CLASSES={TM_CLASSES}, TM_NCLAUSES={TM_NCLAUSES})" + echo "# fixture,samples,captured,classTotal,labelHist" + + var pooledLabels: seq[int] + var pooledConfusion: array[TM_CLASSES, array[TM_CLASSES, int]] + var pooledClassCorrect, pooledClassTotal = 0 + var pooledLabelHist: array[TM_CLASSES, int] + var bestGun: TmPatternGun + var bestSamples: seq[DiagSample] + var bestName = "" + var bestClassTotal = -1 + + for name in names: + let path = fixturesDir / (name & ".jsonl") + if not fileExists(path): + echo &"# SKIP missing fixture {path}" + continue + let fx = loadFixture(path) + let g = new(TmPatternGun) + g[] = initTmPatternGun() + g[].targetMode = tmGF + g[].diagCapture = true + randomize(1234) + let drv = GunDriver( + name: "TMPatGF", + predictCb: proc(state: WorldState, bs: float): GunPrediction = g[].predict(state, bs), + resultCb: proc(e: FeedbackEvent) = g[].onResult(e), + readyCb: proc(): bool = true) + discard replayFixture(fx, @[drv], fx.enemyId) + + var hs: string + for c in 0.. bestClassTotal and thisSamples.len > 0: + bestClassTotal = g[].classTotal + bestGun = g[] + bestSamples = thisSamples + bestName = name + + # ── pooled label balance + majority baseline vs real accuracy ── + for c in 0..=1={summ.firedAtLeastOnce} " & + &"neverFired={summ.neverFired} posFired={summ.posFired} negFired={summ.negFired}" + echo &"# clause length: mean={summ.meanLength:.2f} max={summ.maxLength} hist={summ.lengthHist}" + for cb in clauseBalanceByClass(infos, TM_CLASSES): + echo &"# class{cb.cls}: nonEmpty={cb.nonEmpty} empty={cb.empty} " & + &"posFired={cb.posFired} negFired={cb.negFired} neverFired={cb.neverFired}" + + let contribs = featureContributions(m, bestSamples, spec) + let ranked = rankedInputs(contribs) + let dead = deadInputs(contribs) + let never = neverUsedInputs(contribs) + let consts = constantInputs(bestSamples, TM_NBITS) + echo "\n## INPUT VALUE RANKING (top 15 by weighted vote-share)" + for i in 0..2} {ranked[i].name:<26} appearances={ranked[i].appearances:<8} weighted={ranked[i].weighted:.1f}" + echo "\n## DEAD-INPUT LIST (weighted < 5% of top, or never in a voting clause):" + echo "# ", dead + echo "## never-used (strict appearances==0): ", never + echo "## constant inputs (zero variance in this fixture): ", consts + + echo "\n## TOP FIRING POSITIVE CLAUSES per class" + for cls in 0.. necessary literals: {spec.describeClause(necessaryLiterals(m, bestSamples, cls), cls)}" + +when isMainModule: + main() diff --git a/common_libs/tests/test_tm_diag.nim b/common_libs/tests/test_tm_diag.nim new file mode 100644 index 0000000..02599fd --- /dev/null +++ b/common_libs/tests/test_tm_diag.nim @@ -0,0 +1,165 @@ +## Pure unit tests for the tm_diag diagnostics kit. +## Covers: FeatureSpec rendering, contribution counting, dead-input detection, +## majority baseline, clause introspection and trend classification. +## +## Run: nim c -r --path:common_libs common_libs/tests/test_tm_diag.nim + +import tm_diag/diagnostics + +var checks = 0 +var failures = 0 +proc check(name: string, ok: bool) = + inc checks + if ok: echo "PASS: ", name + else: echo "FAIL: ", name; inc failures + +# ── 1. FeatureSpec rendering ───────────────────────────────────────────────── + +proc testRendering() = + let s = draftTMSpec() + check "draft spec is 49 bits", s.nBits == 49 + check "draft WALLS is 8 bits", s.blocks[0].count + s.blocks[1].count == 8 + check "draft describe(0) = dist-wall<50", s.describe(0) == "dist-wall<50" + check "draft describe(45) = lat DEAD-ON", + s.describe(45) == "lat DEAD-ON -18..+18" + check "draft describe(48) = lat>+72", s.describe(48) == "lat>+72" + check "positive literal", s.describeLiteral(0) == "dist-wall<50" + check "negated literal", s.describeLiteral(49) == "NOT dist-wall<50" + check "clause sentence", + s.describeClause(@[0, 45], 2) == + "IF dist-wall<50 AND lat DEAD-ON -18..+18 THEN class=2" + check "clause sentence with negation", + s.describeClause(@[0, 49 + 45], 1) == + "IF dist-wall<50 AND NOT lat DEAD-ON -18..+18 THEN class=1" + check "empty clause renders TRUE", + s.describeClause(@[], 3) == "IF TRUE (empty clause) THEN class=3" + check "no class suffix when cls<0", + s.describeClause(@[0]) == "IF dist-wall<50" + + let p = tmPatternSpec() + check "tm_pattern spec is 40 bits", p.nBits == 40 + check "tm_pattern UNUSED block is last and 2 bits", + p.blocks[^1].name == "UNUSED" and p.blocks[^1].count == 2 + check "tm_pattern describe(38) = UNUSED-38", p.describe(38) == "UNUSED-38" + +# ── 2. majority baseline / data checks ─────────────────────────────────────── + +proc testDataChecks() = + let dc = dataChecks(@[0, 0, 0, 1, 2], 3) + check "majority count", dc.classCounts == @[3, 1, 1] + check "majority class", dc.majorityClass == 0 + check "majority share 3/5", abs(dc.majorityShare - 0.6) < 1e-9 + check "majority = accuracy baseline", + abs(dc.majorityAccuracy - 0.6) < 1e-9 + check "over-threshold flagged", dc.overThreshold + + let dc2 = dataChecks(@[0, 0, 1, 1, 2, 2, 3, 3], 4) + check "balanced data is not flagged", not dc2.overThreshold + check "balanced majority = 1/4", abs(dc2.majorityShare - 0.25) < 1e-9 + +# ── 3. contribution counting / dead inputs (hand-built model) ──────────────── + +proc tinySpec(): FeatureSpec = + var s = FeatureSpec() + s.addBlock("f", 4, @["f0", "f1", "f2", "f3"]) + s + +proc tinyModel(): TmMachine = + ## nBits=4, 2 classes, 4 clauses/class. + ## class0 clause0 includes +f0 ; class1 clause0 includes +f1. + ## All other clauses are empty. + result = newMachine(4, 2, 4, 64, 3.0, 1) + result.teams[0][0 * result.nLiterals + 0] = 1 + result.teams[1][0 * result.nLiterals + 1] = 1 + +proc testContributions() = + let m = tinyModel() + let spec = tinySpec() + let samples = @[ + makeSample(4, @[1, 0, 0, 0], 0), + makeSample(4, @[0, 1, 0, 0], 1), + ] + let infos = clauseInfo(m, samples, spec) + var class0Info: ClauseInfo + for c in infos: + if c.cls == 0 and c.index == 0: class0Info = c + check "clauseInfo: class0/clause0 fires on sample A", class0Info.votes == 1 + check "clauseInfo: clause length 1", class0Info.length == 1 + check "clauseInfo: rendered", class0Info.text == "IF f0 THEN class=0" + check "clauseInfo: positive polarity", class0Info.polarity == 1 + + let summ = clauseSummary(infos) + check "clauseSummary: 8 clauses total", summ.totalClauses == 8 + check "clauseSummary: 6 empty clauses", summ.emptyClauses == 6 + check "clauseSummary: 2 firing clauses", summ.firedAtLeastOnce == 2 + check "clauseSummary: mean length 1.0", abs(summ.meanLength - 1.0) < 1e-9 + check "clauseSummary: 2 positive fired", summ.posFired == 2 and summ.negFired == 0 + + let contribs = featureContributions(m, samples, spec) + check "contribution f0 = 1.0", abs(contribs[0].weighted - 1.0) < 1e-9 and + contribs[0].appearances == 1 + check "contribution f1 = 1.0", abs(contribs[1].weighted - 1.0) < 1e-9 and + contribs[1].appearances == 1 + check "contribution f2 = 0", contribs[2].weighted == 0.0 + check "contribution f3 = 0", contribs[3].weighted == 0.0 + + let dead = deadInputs(contribs) + check "deadInputs flags f2,f3", dead == @[2, 3] + check "neverUsedInputs flags f2,f3", neverUsedInputs(contribs) == @[2, 3] + let ranked = rankedInputs(contribs) + check "rankedInputs puts f0/f1 first", ranked[0].bit in {0, 1} and ranked[1].bit in {0, 1} + check "rankedInputs puts f2/f3 last", ranked[^1].bit in {2, 3} + + check "constantInputs finds the info-free f2,f3", + constantInputs(samples, 4) == @[2, 3] + +proc testConstantInputs() = + let samples = @[ + makeSample(4, @[1, 0, 0, 0], 0), + makeSample(4, @[1, 1, 0, 0], 1), + ] + check "constantInputs finds f0,f2,f3", constantInputs(samples, 4) == @[0, 2, 3] + +proc testNecessaryLiterals() = + let m = tinyModel() + let spec = tinySpec() + let samples = @[ + makeSample(4, @[1, 0, 0, 0], 0), + makeSample(4, @[0, 1, 0, 0], 1), + ] + check "necessaryLiterals(class0) = {+f0}", necessaryLiterals(m, samples, 0) == @[0] + check "necessaryLiterals(class1) = {+f1}", necessaryLiterals(m, samples, 1) == @[1] + +# ── 4. trend classification ────────────────────────────────────────────────── + +proc testTrend() = + check "rising trend", classifyTrend(@[0.1, 0.3, 0.5, 0.7]) == "rising" + check "flat trend", classifyTrend(@[0.5, 0.5, 0.5, 0.5]) == "flat" + check "falling trend", classifyTrend(@[0.7, 0.5, 0.3, 0.1]) == "falling" + check "rising-then-falling", + classifyTrend(@[0.3, 0.7, 0.75, 0.5]) == + "rising-then-falling (noise fitting)" + +# ── 5. machineFromTeams ────────────────────────────────────────────────────── + +proc testMachineFromTeams() = + let m = tinyModel() + let m2 = machineFromTeams(4, 2, 4, 64, 3.0, m.teams) + check "machineFromTeams copies teams", m2.teams[0] == m.teams[0] + check "machineFromTeams predicts consistently", + m2.predictClass(makeSample(4, @[1, 0, 0, 0], 0).lits) == + m.predictClass(makeSample(4, @[1, 0, 0, 0], 0).lits) + +when isMainModule: + testRendering() + testDataChecks() + testContributions() + testConstantInputs() + testNecessaryLiterals() + testTrend() + testMachineFromTeams() + echo "" + if failures > 0: + echo failures, " / ", checks, " check(s) FAILED" + quit(1) + echo "All ", checks, " tm_diag unit checks passed." diff --git a/common_libs/tm_diag/README.md b/common_libs/tm_diag/README.md new file mode 100644 index 0000000..e8704c4 --- /dev/null +++ b/common_libs/tm_diag/README.md @@ -0,0 +1,124 @@ +# tm_diag — the Tsetlin diagnostics kit + +First-class, offline diagnostics for a Tsetlin-Machine head. Answers the two +questions that clause-reading alone cannot: + +> **Is the TM learning badly, or is the data bad?** + +Built before the new TM gun, so a bad design cannot hide behind the data and a +bad dataset cannot hide behind the model. + +Everything is pure and offline — no battles, no Java, no harness. + +## Files + +| file | contents | +|---|---| +| `feature_spec.nim` | `FeatureSpec`, `describe`, `describeClause`, `draftTMSpec()` (49-bit draft), `tmPatternSpec()` (40-bit shipped encoding) | +| `tm_core.nim` | compact deterministic Granmo Table 2/3 multiclass TM (mirrors the tm_pattern core), introspectable clause layout | +| `diagnostics.nim` | the six groups (re-exports the two above) | + +Import everything with: + +```nim +import tm_diag/diagnostics +``` + +## Task 1 — named features / clause rendering + +```nim +let spec = draftTMSpec() # 49 bits, all one-hot +spec.nBits # 49 +spec.describe(45) # "lat DEAD-ON -18..+18" +spec.describeLiteral(49 + 45) # "NOT lat DEAD-ON -18..+18" +spec.describeClause(@[0, 45], 2) # "IF dist-wall<50 AND lat DEAD-ON -18..+18 THEN class=2" +spec.describeClause(@[], 1) # "IF TRUE (empty clause) THEN class=1" +``` + +A block is added with `addBlock(name, count, bitNames?)`; a bit with no explicit +name renders as `blockName[k]`, and a single-bit block renders as its name. + +`tmPatternSpec()` mirrors `guns/tm_pattern.nim`'s `tmBuildBits` exactly. **Its two +last bits (`UNUSED-38/39`) are a real bug**: `tmBuildBits` writes only 38 raw +bits into an `array[TM_NBITS=40, uint8]`, so bits 38 and 39 are always 0 and +their negations always 1. The kit reports them as constant dead inputs. + +## Task 2 — the six groups + +All functions take a trained `TmMachine` (or an externally supplied clause set) +plus `seq[DiagSample]` where `DiagSample.lits` is the pos-then-neg literal vector +and `DiagSample.label` the true class. + +```nim +# build samples from raw bits +let s = makeSample(nBits, rawBits, label, order) + +# 1. pre-flight DATA checks +let dc = dataChecks(labels, nClasses, threshold = 0.30) +# dc.classCounts, dc.classShares, dc.majorityClass, dc.majorityShare, +# dc.majorityAccuracy, dc.overThreshold, dc.flags +let sc = shuffledLabelControl(tmplMachine, samples) # (acc, majority, ...) + +# 2. clause introspection +let infos = clauseInfo(m, samples, spec) +let summ = clauseSummary(infos) # empty / neverFired / length hist +for c in topClauses(infos, 10): echo c.text, " votes=", c.votes +for cb in clauseBalanceByClass(infos, nClasses): echo cb +for cl in 0..= nBits, + "rawBits.len (" & $rawBits.len & ") < nBits (" & $nBits & ")" + result.lits = newSeq[uint8](2 * nBits) + for i in 0..= 0 and l < nClasses: inc result.classCounts[l] + var maj = 0 + for c in 0.. 0: result.classCounts[c].float / result.n.float else: 0.0 + if result.classCounts[c] > result.classCounts[maj]: maj = c + result.majorityClass = maj + result.majorityShare = result.classShares[maj] + result.majorityAccuracy = result.majorityShare + result.maxShare = result.classShares[maj] + for c in 0.. threshold: + result.overThreshold = true + result.flags.add &"MAJORITY HEAVY: class {c} = " & + &"{result.classCounts[c]}/{result.n} = " & + &"{result.classShares[c]*100.0:.1f}% > {threshold*100.0:.0f}%" + +proc shuffledLabelControl*(tmpl: TmMachine, samples: openArray[DiagSample], + epochs = 1, seed = 4242'u64): + tuple[acc, majority: float, n, nClasses: int, trainedObs: int] = + ## PIPELINE LEAK TEST: retrain a fresh copy on SHUFFLED labels. A correct + ## pipeline cannot beat the majority baseline on shuffled labels; a large + ## margin means information is leaking (features from the future, duplicated + ## samples, ...). + var rng = seedRng(seed xor 0xabcdef12'u64) + var shuf = newSeq[DiagSample](samples.len) + for i, s in samples: shuf[i] = s + var labels = newSeq[int](samples.len) + for i in 0.. maxLen: maxLen = inf.length + if inf.votes > 0: + inc result.firedAtLeastOnce + if inf.polarity > 0: inc result.posFired + else: inc result.negFired + else: + inc result.neverFired + result.maxLength = maxLen + result.lengthHist = newSeq[int](maxLen + 1) + for inf in infos: + if inf.length > 0: inc result.lengthHist[inf.length] + result.meanLength = + if result.nonEmpty > 0: lenSum.float / result.nonEmpty.float else: 0.0 + +type + ClassClauseBalance* = object + cls*: int + empty*: int + nonEmpty*: int + posFired*: int + negFired*: int + neverFired*: int + +proc clauseBalanceByClass*(infos: openArray[ClauseInfo], + nClasses: int): seq[ClassClauseBalance] = + ## Positive/negative firing balance PER CLASS, plus empty and never-fired + ## counts. A healthy class should use both polarities; a class with 0 firing + ## positive clauses has no readable rule. + result = newSeq[ClassClauseBalance](nClasses) + for c in 0..= nClasses: continue + if inf.length == 0: + inc result[inf.cls].empty + continue + inc result[inf.cls].nonEmpty + if inf.votes == 0: + inc result[inf.cls].neverFired + elif inf.polarity > 0: + inc result[inf.cls].posFired + else: + inc result[inf.cls].negFired + +proc topClauses*(infos: openArray[ClauseInfo], k = 10, + minVotes = 1): seq[ClauseInfo] = + ## The clauses that actually fire, ranked by vote count (ties: longer first, + ## then class/index). + for inf in infos: + if inf.votes >= minVotes and inf.length > 0: result.add inf + result.sort(proc(a, b: ClauseInfo): int = + result = cmp(b.votes, a.votes) + if result == 0: result = cmp(b.length, a.length) + if result == 0: result = cmp(a.cls, b.cls) + if result == 0: result = cmp(a.index, b.index)) + if result.len > k: result.setLen(k) + +proc topClausesByPolarity*(infos: openArray[ClauseInfo], polarity: int, + k = 10): seq[ClauseInfo] = + ## Firing clauses of one polarity only (positive = the rule-encoding clauses, + ## negative = the discriminators), ranked by votes. + var filtered: seq[ClauseInfo] + for inf in infos: + if inf.polarity == polarity: filtered.add inf + result = topClauses(filtered, k) + +proc necessaryLiterals*(m: TmMachine, samples: openArray[DiagSample], + cls: int): seq[int] = + ## The literals common to EVERY positive-polarity clause of `cls` that fires + ## at least once. A converged multiclass TM pads its clauses with redundant + ## literals; the intersection strips the padding and recovers the class rule + ## (e.g. the planted `A AND B`). + var first = true + for cl in 0.. mx: mx = c.weighted + let cut = relThreshold * mx + for c in contribs: + if c.appearances == 0 or c.weighted < cut: + result.add c.bit + +proc neverUsedInputs*(contribs: openArray[FeatureContribution]): seq[int] = + ## Strict form: raw bits that never appear in a clause that casts a vote. + for c in contribs: + if c.appearances == 0: result.add c.bit + +proc constantInputs*(samples: openArray[DiagSample], nBits: int): seq[int] = + ## Raw bits with zero variance across the sample set (information-free even if + ## a clause happens to include them). + if samples.len == 0: return + for b in 0..= 0 and s.label < m.nClasses and p >= 0 and p < m.nClasses: + inc result.confusion[s.label][p] + inc result.predCounts[p] + if p == s.label: inc result.correct + result.acc = + if result.n > 0: result.correct.float / result.n.float else: 0.0 + var rowSum = newSeq[int](m.nClasses) + for c in 0.. rowSum[maj]: maj = c + result.majorityClass = maj + result.majorityShare = + if result.n > 0: rowSum[maj].float / result.n.float else: 0.0 + result.majorityBaseline = result.majorityShare + result.margin = result.acc - result.majorityBaseline + for c in 0.. 0: result.confusion[c][c].float / rowSum[c].float else: 0.0 + result.perClassPrecision[c] = + if result.predCounts[c] > 0: + result.confusion[c][c].float / result.predCounts[c].float + else: 0.0 + result.predMajorityShare = + if result.n > 0: result.predCounts[maj].float / result.n.float else: 0.0 + +# ───────────────────────────────────────────────────────────────────────────── +# GROUP 5 — learning curve +# ───────────────────────────────────────────────────────────────────────────── + +type + LearningCurve* = object + points*: seq[int] ## samples seen + accs*: seq[float] ## accuracy on the (held-out) eval set + trend*: string + +proc classifyTrend*(accs: seq[float], eps = 0.02): string = + if accs.len < 3: return "insufficient" + let first = accs[0] + let last = accs[^1] + var mx = first + var mxIdx = 0 + for i, a in accs: + if a > mx: + mx = a + mxIdx = i + if mx - last > eps and mxIdx > 0 and mxIdx < accs.len - 1: + return "rising-then-falling (noise fitting)" + if last - first > eps: return "rising" + if first - last > eps: return "falling" + "flat" + +proc learningCurve*(tmpl: TmMachine, train, eval: openArray[DiagSample], + nPoints = 10, seed = 777'u64): LearningCurve = + ## Train a fresh machine on growing prefixes of `train` and score each prefix + ## on `eval`, so fast vs slow adaptation is directly visible. + var m = tmpl + m.resetMachine(seed) + let np = max(2, nPoints) + let step = max(1, train.len div np) + var seen = 0 + var i = 0 + while i < np: + let upto = min(train.len, (i + 1) * step) + var k = seen + while k < upto: + m.trainSample(train[k].lits, train[k].label) + inc k + seen = upto + var correct = 0 + for s in eval: + if m.predictClass(s.lits) == s.label: inc correct + result.points.add upto + result.accs.add (if eval.len > 0: correct.float / eval.len.float else: 0.0) + inc i + if upto >= train.len: + while i < np: + result.points.add train.len + result.accs.add result.accs[^1] + inc i + break + result.trend = classifyTrend(result.accs) + +# ───────────────────────────────────────────────────────────────────────────── +# GROUP 6 — ablation hooks +# ───────────────────────────────────────────────────────────────────────────── + +type + AblationResult* = object + name*: string + kind*: string ## "baseline" | "drop-block" | "scramble-feature" + baselineAcc*: float + ablatedAcc*: float + delta*: float ## ablated - baseline + +proc ablateAdd*(a, b: AblationResult): AblationResult = + ## Accumulate ablations across folds/fixtures. + result.name = a.name + result.kind = a.kind + let n = 2.0 + result.baselineAcc = (a.baselineAcc + b.baselineAcc) / n + result.ablatedAcc = (a.ablatedAcc + b.ablatedAcc) / n + result.delta = (a.delta + b.delta) / n + +proc dropBits(s: DiagSample, nBits, first, count: int): DiagSample = + result = s + result.lits = s.lits + for b in first.. the block does not earn its bits. + ## Pass `baselineAcc` (from ablateBaseline) to skip recomputing it per block. + doAssert blockIdx >= 0 and blockIdx < spec.blocks.len + let b = spec.blocks[blockIdx] + var tr = newSeq[DiagSample](train.len) + for i, s in train: tr[i] = dropBits(s, spec.nBits, b.first, b.count) + var ev = newSeq[DiagSample](eval.len) + for i, s in eval: ev[i] = dropBits(s, spec.nBits, b.first, b.count) + let base = + if baselineAcc >= 0.0: baselineAcc + else: evalAcc(trainModel(tmpl, train, epochs, seed), eval) + let abl = evalAcc(trainModel(tmpl, tr, epochs, seed), ev) + AblationResult(name: b.name, kind: "drop-block", + baselineAcc: base, ablatedAcc: abl, delta: abl - base) + +proc ablateScrambleFeature*(tmpl: TmMachine, train, eval: openArray[DiagSample], + spec: FeatureSpec, bit: int, epochs = 1, + seed = 777'u64, baselineAcc = -1.0): AblationResult = + ## Retrain from scratch with ONE raw bit permuted across samples in both train + ## and eval. A negative delta => the bit is genuinely informative. + var rng = seedRng(seed xor 0x5bd1e995'u64) + var tr = newSeq[DiagSample](train.len) + for i, s in train: tr[i] = s + var ev = newSeq[DiagSample](eval.len) + for i, s in eval: ev[i] = s + scrambleBit(tr, spec.nBits, bit, rng) + scrambleBit(ev, spec.nBits, bit, rng) + let base = + if baselineAcc >= 0.0: baselineAcc + else: evalAcc(trainModel(tmpl, train, epochs, seed), eval) + let abl = evalAcc(trainModel(tmpl, tr, epochs, seed), ev) + AblationResult(name: spec.describe(bit), kind: "scramble-feature", + baselineAcc: base, ablatedAcc: abl, delta: abl - base) + +proc ablateDropAllBlocks*(tmpl: TmMachine, train, eval: openArray[DiagSample], + spec: FeatureSpec, epochs = 1, seed = 777'u64): + seq[AblationResult] = + for i in 0.. 0, "block '" & name & "' must have at least one bit" + doAssert bitNames.len == 0 or bitNames.len == count, + "block '" & name & "': bitNames.len (" & $bitNames.len & + ") != count (" & $count & ")" + s.blocks.add FeatureBlock(name: name, first: s.nBits, count: count, + bitNames: bitNames) + s.nBits += count + +proc blockOf*(s: FeatureSpec, bit: int): int = + ## Index of the block owning `bit`, or -1. + for i, b in s.blocks: + if bit >= b.first and bit < b.first + b.count: return i + -1 + +proc describe*(s: FeatureSpec, bit: int): string = + ## Human-readable name of a single raw bit. + if bit < 0 or bit >= s.nBits: return "bit" & $bit + let bi = s.blockOf(bit) + if bi < 0: return "bit" & $bit + let b = s.blocks[bi] + let k = bit - b.first + if b.bitNames.len == b.count: + result = b.bitNames[k] + elif b.count == 1: + result = b.name + else: + result = b.name & "[" & $k & "]" + +proc describeLiteral*(s: FeatureSpec, lit: int): string = + ## `lit < nBits` is the positive literal; `lit >= nBits` is its negation + ## (the kit uses the same pos-then-neg literal layout as the shipped TM). + if lit < 0: return "?" + if lit < s.nBits: return s.describe(lit) + "NOT " & s.describe(lit - s.nBits) + +proc describeClause*(s: FeatureSpec, lits: openArray[int], cls = -1): string = + ## Render a conjunction. Example: + ## IF near-wall AND bullet-dead-on AND NOT turn-left(t-2) THEN class=3 + var parts: seq[string] + for l in lits: parts.add s.describeLiteral(l) + let body = if parts.len == 0: "TRUE (empty clause)" else: parts.join(" AND ") + result = "IF " & body + if cls >= 0: result.add " THEN class=" & $cls + +# ── the DRAFT encoding, as the worked example ──────────────────────────────── +# +# From the design session; all one-hot binaries. TOTAL 49 bits. This is NOT +# wired into any gun — it is the example the diagnostics are demonstrated on. +# WALLS (8): dist-to-nearest-wall 4 bins; which-wall-nearest 4 +# US (9): dist-from-us 6 bins; enemy-heading-vs-line-to-us 3 +# MOTION (20): turn-direction 3; ticks-since-reversal 5; +# turn-consistency-10 3; distance-moved-10 3; +# speed-trend-10 3; turn-rate-change-5 3 +# BULLETS(12): time-until-our-bullet-arrives 5; bullet-lateral-offset 7 +proc draftTMSpec*(): FeatureSpec = + result = FeatureSpec() + result.addBlock("dist-to-nearest-wall", 4, + @["dist-wall<50", "dist-wall 50-100", "dist-wall 100-200", "dist-wall>200"]) + result.addBlock("which-wall-nearest", 4, + @["wall-left", "wall-right", "wall-top", "wall-bottom"]) + + result.addBlock("dist-from-us", 6, + @["dist-us<100", "dist-us 100-200", "dist-us 200-300", + "dist-us 300-400", "dist-us 400-600", "dist-us>600"]) + result.addBlock("enemy-heading-vs-line-to-us", 3, + @["hdg-vs-us perpendicular", "hdg-vs-us angled", "hdg-vs-us along-line"]) + + result.addBlock("turn-direction", 3, + @["turn-left(t-0)", "turn-left(t-1)", "turn-left(t-2)"]) + result.addBlock("ticks-since-reversal", 5, + @["since-rev<5", "since-rev 5-10", "since-rev 10-20", + "since-rev 20-40", "since-rev>40"]) + result.addBlock("turn-consistency-10", 3, + @["turn-consistency low", "turn-consistency med", "turn-consistency high"]) + result.addBlock("distance-moved-10", 3, + @["dist-moved-10 low", "dist-moved-10 med", "dist-moved-10 high"]) + result.addBlock("speed-trend-10", 3, + @["speed-trend falling", "speed-trend flat", "speed-trend rising"]) + result.addBlock("turn-rate-change-5", 3, + @["turnrate-change down", "turnrate-change flat", "turnrate-change up"]) + + result.addBlock("time-until-bullet", 5, + @["tta none", "tta<5", "tta 5-10", "tta 10-20", "tta>20"]) + result.addBlock("bullet-lateral-offset", 7, + @["lat<-72", "lat -72..-36", "lat -36..-18", "lat DEAD-ON -18..+18", + "lat +18..+36", "lat +36..+72", "lat>+72"]) + +# ── the SHIPPED tm_pattern encoding (40 raw bits) ──────────────────────────── +# +# Exact mirror of `tmBuildBits` in common_libs/guns/tm_pattern.nim, in the order +# it writes them. NOTE: tmBuildBits writes only 38 bits into an +# `array[TM_NBITS=40, uint8]`; bits 38 and 39 are never assigned and stay 0 for +# the whole life of the gun. They are named UNUSED-* here on purpose so the +# dead-input detector can flag them. +proc tmPatternSpec*(): FeatureSpec = + result = FeatureSpec() + for i in 0..2: + result.addBlock("lat-sign(t-" & $i & ")", 2, + @["lat-pos(t-" & $i & ")", "lat-neg(t-" & $i & ")"]) + for i in 0..2: + result.addBlock("turn-sign(t-" & $i & ")", 2, + @["turn-left(t-" & $i & ")", "turn-right(t-" & $i & ")"]) + result.addBlock("since-reversal", 3, + @["since-rev<=3", "since-rev 3-10", "since-rev>10"]) + result.addBlock("lat-persist", 1, @["lat-persist"]) + result.addBlock("speed-band", 3, + @["speed<1", "speed 1-4", "speed>=4"]) + result.addBlock("distance-band", 3, + @["dist<150", "dist 150-350", "dist>=350"]) + result.addBlock("flight-band", 3, + @["flight<10", "flight 10-25", "flight>=25"]) + result.addBlock("wall-near", 4, + @["wall-near-bottom", "wall-near-top", "wall-near-right", "wall-near-left"]) + result.addBlock("radial-frac", 3, + @["radial<0.35", "radial 0.35-0.7", "radial>=0.7"]) + result.addBlock("enemy-energy", 2, @["enemyE<20", "enemyE>=20"]) + result.addBlock("heading-vs-los", 2, + @["heading-toward-us", "heading-away-us"]) + result.addBlock("closing", 2, @["closing<0.3", "closing>0.3"]) + result.addBlock("UNUSED", 2, @["UNUSED-38", "UNUSED-39"]) diff --git a/common_libs/tm_diag/tm_core.nim b/common_libs/tm_diag/tm_core.nim new file mode 100644 index 0000000..f139a13 --- /dev/null +++ b/common_libs/tm_diag/tm_core.nim @@ -0,0 +1,173 @@ +## tm_diag/tm_core.nim — a compact, deterministic Granmo Table 2/3 Tsetlin +## Machine (multiclass), with an introspection-friendly clause layout. +## +## This mirrors the corrected core in common_libs/guns/tm_pattern.nim +## (`tmEval` / `tmForward` / `tmLearnDir` and the Eq. 6 empty-clause bootstrap) +## so the diagnostics get demonstrated on the SAME algorithm the real gun uses. +## It is a standalone copy because the gun's core is private and pulls in the +## gun harness; here everything is pure and offline. +## +## Layout: `teams[c][cl * nLiterals + lit]`, literal `i` is the positive literal +## and `i + nBits` its negation — identical to the gun. + +import std/math + +type + TmRng* = object + s*: uint64 + + TmMachine* = object + nBits*: int + nLiterals*: int + nClauses*: int + nClasses*: int + half*: int + nStates*: int + sValue*: float + teams*: seq[seq[int16]] + rng*: TmRng + +proc seedRng*(seed: uint64): TmRng = + result.s = seed + if result.s == 0: result.s = 0x9e3779b97f4a7c15'u64 + +proc nextU64*(r: var TmRng): uint64 = + r.s = r.s xor (r.s shl 13) + r.s = r.s xor (r.s shr 7) + r.s = r.s xor (r.s shl 17) + r.s + +proc rand01*(r: var TmRng): float = + ## Uniform [0,1). + (r.nextU64() shr 11).float / 9007199254740992.0 + +proc newMachine*(nBits, nClasses: int, nClauses = 40, nStates = 64, + sValue = 3.0, seed = 12345'u64): TmMachine = + result.nBits = nBits + result.nLiterals = 2 * nBits + result.nClauses = nClauses + result.nClasses = nClasses + result.half = nClauses div 2 + result.nStates = nStates + result.sValue = sValue + result.rng = seedRng(seed) + result.teams = newSeq[seq[int16]](nClasses) + for c in 0.. 0: + hasInc = true + if lits[lit] == 0'u8: return 0'u8 + if hasInc: return 1'u8 + # Eq. 6: the empty conjunction is vacuously true while learning, false when + # classifying. Without this the all-Exclude init deadlocks. + return if learning: 1'u8 else: 0'u8 + +proc tmForward*(m: TmMachine, team: seq[int16], lits: openArray[uint8], + cache: var seq[uint8]): float = + var v = 0.0 + for cl in 0..= pFeedback: continue + let pol = m.tmPolarity(cl) + let cOut = cache[cl] + let base = cl * m.nLiterals + if pol * d > 0.0: + # Type I (Table 2) collapsed to the resulting state move. + for lit in 0.. bestV: + bestV = v + best = c + best + +proc trainSample*(m: var TmMachine, lits: openArray[uint8], label: int) = + var votes = newSeq[float](m.nClasses) + var caches = newSeq[seq[uint8]](m.nClasses) + for c in 0.. 0: + hasInc = true + if lits[lit] == 0'u8: return false + hasInc + +proc clauseLits*(m: TmMachine, cls, cl: int): seq[int] = + ## The literal indices included by one clause (>0 state). + let base = cl * m.nLiterals + for lit in 0.. 0: result.add lit