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.
This commit is contained in:
@@ -0,0 +1,201 @@
|
||||
## Task 3 — VALIDATE THE DIAGNOSTICS AGAINST A KNOWN GROUND TRUTH.
|
||||
##
|
||||
## A synthetic dataset is generated from a planted decision list over the DRAFT
|
||||
## 49-bit encoding:
|
||||
## A = bit 0 ("dist-wall<50", WALLS block)
|
||||
## B = bit 45 ("lat DEAD-ON -18..+18", BULLETS block)
|
||||
## class 2 = A AND B ; class 1 = A AND NOT B ; class 0 = NOT A
|
||||
## Deliberately irrelevant: the whole US block (bits 8..16) plus one pure-noise
|
||||
## MOTION bit (bit 17).
|
||||
##
|
||||
## The kit must recover the rule, flag the dead inputs, show ~0 ablation delta
|
||||
## for the irrelevant block, and sit at the majority baseline under shuffled
|
||||
## labels. Every number printed here is MEASURED.
|
||||
##
|
||||
## Run: nim c -r --path:common_libs common_libs/tests/diag_synthetic.nim
|
||||
|
||||
import std/[random, strformat, strutils, algorithm]
|
||||
import tm_diag/diagnostics
|
||||
|
||||
const
|
||||
NBits = 49
|
||||
BitA = 0 ## dist-wall<50
|
||||
BitB = 45 ## bullet lateral DEAD-ON
|
||||
NoiseBit = 17
|
||||
NClasses = 3
|
||||
|
||||
var failures = 0
|
||||
proc check(name: string, ok: bool) =
|
||||
if ok: echo "PASS: ", name
|
||||
else: echo "FAIL: ", name; inc failures
|
||||
|
||||
proc genDataset(n, seed: int): seq[DiagSample] =
|
||||
var rng = initRand(seed)
|
||||
for i in 0..<n:
|
||||
var raw = newSeq[int](NBits)
|
||||
for b in 0..<NBits:
|
||||
raw[b] = (if rng.rand(1.0) < 0.5: 1 else: 0)
|
||||
let a = if rng.rand(1.0) < 0.4: 1 else: 0
|
||||
let bb = if rng.rand(1.0) < 0.5: 1 else: 0
|
||||
raw[BitA] = a
|
||||
raw[BitB] = bb
|
||||
let label =
|
||||
if a == 1 and bb == 1: 2
|
||||
elif a == 1: 1
|
||||
else: 0
|
||||
result.add makeSample(NBits, raw, label, i)
|
||||
|
||||
when isMainModule:
|
||||
let spec = draftTMSpec()
|
||||
let train = genDataset(3000, 1)
|
||||
let eval = genDataset(1500, 2)
|
||||
echo &"# dataset: train={train.len} eval={eval.len} nBits={spec.nBits} rule=A&B->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..<NClasses:
|
||||
if dc.classCounts[c] > 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..<NClasses:
|
||||
var clsInfos: seq[ClauseInfo]
|
||||
for c in infos:
|
||||
if c.cls == cls: clsInfos.add c
|
||||
for c in topClausesByPolarity(clsInfos, 1, 3):
|
||||
echo &"# class{c.cls} votes={c.votes:<5} len={c.length} pol={c.polarity:+d} {c.text}"
|
||||
|
||||
# planted rule recovery: the literals common to every firing positive clause
|
||||
# of a class are exactly the planted rule (the TM pads clauses with junk).
|
||||
proc hasRule(cls: int, wantPositive: set[uint8], wantNeg: set[uint8]): bool =
|
||||
var pos, neg: set[uint8]
|
||||
for l in necessaryLiterals(m, eval, cls):
|
||||
if l < NBits: pos.incl uint8(l)
|
||||
else: neg.incl uint8(l - NBits)
|
||||
pos == wantPositive and neg == wantNeg
|
||||
echo "# recovered necessary literals per class (intersection of firing positive clauses):"
|
||||
for c in 0..<NClasses:
|
||||
echo &"# class{c}: {spec.describeClause(necessaryLiterals(m, eval, c), c)}"
|
||||
check "clause introspection recovers class2 = (A AND B)",
|
||||
hasRule(2, {uint8(BitA), uint8(BitB)}, {})
|
||||
check "clause introspection recovers class1 = (A AND NOT B)",
|
||||
hasRule(1, {uint8(BitA)}, {uint8(BitB)})
|
||||
check "clause introspection recovers class0 = (NOT A)",
|
||||
hasRule(0, {}, {uint8(BitA)})
|
||||
|
||||
# ── GROUP 3: per-feature contribution + dead inputs ──
|
||||
let contribs = featureContributions(m, eval, spec)
|
||||
let dead = deadInputs(contribs)
|
||||
let neverUsed = neverUsedInputs(contribs)
|
||||
let consts = constantInputs(eval, NBits)
|
||||
let ranked = rankedInputs(contribs)
|
||||
echo "\n## GROUP 3 per-feature contribution"
|
||||
echo "# top-8 ranked inputs:"
|
||||
for i in 0..<min(8, ranked.len):
|
||||
echo &"# {ranked[i].bit:>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..<NClasses:
|
||||
echo &"# true {c}: {ad.confusion[c]}"
|
||||
for c in 0..<NClasses:
|
||||
echo &"# class{c}: recall={ad.perClassRecall[c]*100:.1f}% precision={ad.perClassPrecision[c]*100:.1f}%"
|
||||
check "accuracy beats the majority baseline by a large margin", ad.margin > 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..<lc.points.len:
|
||||
echo &"# samples={lc.points[i]:<5} acc={lc.accs[i]*100:.2f}%"
|
||||
check "the learning curve is rising (not flat)",
|
||||
lc.trend.startsWith("rising") and not lc.trend.startsWith("rising-then")
|
||||
|
||||
# ── GROUP 6: ablation ──
|
||||
echo "\n## GROUP 6 ablation (retrain from scratch per variant)"
|
||||
let base = ablateBaseline(tmpl, train, eval, epochs = 15, seed = 777)
|
||||
echo &"# baseline acc={base.baselineAcc*100:.2f}%"
|
||||
# drop every block
|
||||
var usDelta = 0.0
|
||||
var wallsDelta = 0.0
|
||||
var bulletsDelta = 0.0
|
||||
for i, b in spec.blocks:
|
||||
let r = ablateDropBlock(tmpl, train, eval, spec, i, epochs = 15, seed = 777,
|
||||
baselineAcc = base.baselineAcc)
|
||||
echo &"# drop-block {b.name:<28} acc={r.ablatedAcc*100:.2f}% delta={r.delta*100:+.2f}pp"
|
||||
if b.name == "dist-from-us": usDelta = r.delta
|
||||
if b.name == "dist-to-nearest-wall": wallsDelta = r.delta
|
||||
if b.name == "bullet-lateral-offset": bulletsDelta = r.delta
|
||||
check "dropping the irrelevant US block costs ~0 accuracy", abs(usDelta) < 0.03
|
||||
check "dropping the WALLS block (holds A) costs real accuracy", wallsDelta < -0.10
|
||||
check "dropping the BULLETS block (holds B) costs real accuracy", bulletsDelta < -0.05
|
||||
|
||||
# scramble ONE relevant bit vs ONE noise bit
|
||||
let scA = ablateScrambleFeature(tmpl, train, eval, spec, BitA, epochs = 15, seed = 777, baselineAcc = base.baselineAcc)
|
||||
let scNoise = ablateScrambleFeature(tmpl, train, eval, spec, NoiseBit, epochs = 15, seed = 777, baselineAcc = base.baselineAcc)
|
||||
echo &"# scramble {scA.name:<20} delta={scA.delta*100:+.2f}pp"
|
||||
echo &"# scramble {scNoise.name:<20} delta={scNoise.delta*100:+.2f}pp"
|
||||
check "scrambling relevant bit A costs real accuracy", scA.delta < -0.10
|
||||
check "scrambling the noise bit costs ~0 accuracy", abs(scNoise.delta) < 0.03
|
||||
|
||||
# ── GROUP 1b: shuffled-label control ──
|
||||
let sc = shuffledLabelControl(tmpl, train, epochs = 10, seed = 4242)
|
||||
echo "\n## GROUP 1 shuffled-label control"
|
||||
echo &"# shuffled-label accuracy={sc.acc*100:.2f}% majority={sc.majority*100:.2f}% " &
|
||||
&"n={sc.n} gap={(sc.acc-sc.majority)*100:+.2f}pp"
|
||||
check "shuffled-label control sits at the majority baseline (|gap| < 4pp)",
|
||||
abs(sc.acc - sc.majority) < 0.04
|
||||
|
||||
echo ""
|
||||
if failures > 0:
|
||||
echo &"{failures} check(s) FAILED"
|
||||
quit(1)
|
||||
echo "All synthetic diagnostic validation checks passed."
|
||||
@@ -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..<TM_NLITS: result.lits[i] = s.lits[i]
|
||||
result.label = s.label
|
||||
result.order = s.order
|
||||
|
||||
proc printPooledConfusion(cm: array[TM_CLASSES, array[TM_CLASSES, int]]) =
|
||||
var total = 0
|
||||
var correct = 0
|
||||
var rowSum, colSum: array[TM_CLASSES, int]
|
||||
for c in 0..<TM_CLASSES:
|
||||
for p in 0..<TM_CLASSES:
|
||||
total += cm[c][p]
|
||||
if c == p: correct += cm[c][p]
|
||||
rowSum[c] += cm[c][p]
|
||||
colSum[p] += cm[c][p]
|
||||
var maj = 0
|
||||
for c in 1..<TM_CLASSES:
|
||||
if rowSum[c] > 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..<TM_CLASSES:
|
||||
let rec = if rowSum[c] > 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..<TM_CLASSES:
|
||||
inc pooledLabelHist[c], g[].labelHist[c]
|
||||
hs.add $g[].labelHist[c] & ","
|
||||
for p in 0..<TM_CLASSES:
|
||||
pooledConfusion[c][p] += g[].confusion[c][p]
|
||||
pooledClassCorrect += g[].classCorrect
|
||||
pooledClassTotal += g[].classTotal
|
||||
|
||||
echo &"# {name},{fx.states.len},{g[].diagSamples.len},{g[].classTotal},[{hs}]"
|
||||
|
||||
# Keep the fixture with the most warm labelled samples for clause
|
||||
# introspection (the model is fresh per fixture, so we cannot pool teams).
|
||||
var thisSamples: seq[DiagSample]
|
||||
let cap = min(g[].diagSamples.len, maxSamples)
|
||||
thisSamples = newSeq[DiagSample](cap)
|
||||
for i in 0..<cap: thisSamples[i] = toDiag(g[].diagSamples[i])
|
||||
if g[].classTotal > 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..<TM_CLASSES:
|
||||
for _ in 0..<pooledLabelHist[c]: pooledLabels.add c
|
||||
let dc = dataChecks(pooledLabels, TM_CLASSES)
|
||||
echo "\n## POOLED label balance (all resolved bullets, warm+cold)"
|
||||
echo &"# n={dc.n} counts={dc.classCounts}"
|
||||
var sh = ""
|
||||
for c in 0..<TM_CLASSES: sh.add &"c{c}={dc.classShares[c]*100:.1f}% "
|
||||
echo "# shares: ", sh
|
||||
echo &"# majority = class{dc.majorityClass} @ {dc.majorityShare*100:.2f}% overThreshold={dc.overThreshold}"
|
||||
for f in dc.flags: echo "# ", f
|
||||
|
||||
echo "\n## POOLED accuracy (gun's own warm confusion)"
|
||||
echo &"# cross-check: gun classCorrect/classTotal = {pooledClassCorrect}/{pooledClassTotal}"
|
||||
printPooledConfusion(pooledConfusion)
|
||||
|
||||
if bestSamples.len == 0:
|
||||
echo "\n# no captured samples; aborting introspection"
|
||||
return
|
||||
|
||||
echo &"\n## CLAUSE / INPUT INTROSPECTION on fixture '{bestName}' " &
|
||||
&"(bestGun: {bestSamples.len} captured samples, {bestClassTotal} warm)"
|
||||
|
||||
let m = machineFromTeams(TM_NBITS, TM_CLASSES, TM_NCLAUSES, TM_NSTATES, TM_S,
|
||||
bestGun.exportTeams())
|
||||
let infos = clauseInfo(m, bestSamples, spec)
|
||||
let summ = clauseSummary(infos)
|
||||
echo &"# clauses: total={summ.totalClauses} empty={summ.emptyClauses} " &
|
||||
&"nonEmpty={summ.nonEmpty} fired>=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..<min(15, ranked.len):
|
||||
echo &"# {ranked[i].bit:>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..<TM_CLASSES:
|
||||
var clsInfos: seq[ClauseInfo]
|
||||
for c in infos:
|
||||
if c.cls == cls: clsInfos.add c
|
||||
let tc = topClausesByPolarity(clsInfos, 1, 4)
|
||||
if tc.len == 0:
|
||||
echo &"# class{cls}: (no firing positive clauses)"
|
||||
for c in tc:
|
||||
echo &"# class{cls} votes={c.votes:<7} len={c.length} {c.text}"
|
||||
echo &"# -> necessary literals: {spec.describeClause(necessaryLiterals(m, bestSamples, cls), cls)}"
|
||||
|
||||
when isMainModule:
|
||||
main()
|
||||
@@ -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."
|
||||
Reference in New Issue
Block a user