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:
2026-09-22 21:24:27 +02:00
parent b0654d18eb
commit f9f8d84671
8 changed files with 1608 additions and 0 deletions
+27
View File
@@ -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..<TM_CLASSES: result[c] = g.teams[c]
proc exportRadTeams*(g: TmPatternGun): seq[seq[int16]] =
result = newSeq[seq[int16]](TM_CLASSES)
for c in 0..<TM_CLASSES: result[c] = g.radTeams[c]
proc exportRevTeams*(g: TmPatternGun): seq[seq[int16]] =
result = newSeq[seq[int16]](2)
for c in 0..<2: result[c] = g.revTeams[c]
proc initTmRadialGun*(): TmPatternGun =
## The RACK-REGISTERED instance: the RADIAL target mode, which is the
@@ -509,6 +534,8 @@ proc tmResolveTrace(g: var TmPatternGun, t: TmPatternTrace, power: float) =
let winner = if shuffleGF: rand(TM_CLASSES - 1) else: gfToBucket(gf)
inc g.labelHist[winner]
if g.diagCapture:
g.diagSamples.add TmDiagSample(lits: t.lits, label: winner, order: t.fireTick)
if t.warm:
inc g.classTotal
if winner == t.chosen: inc g.classCorrect
+201
View File
@@ -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()
+165
View File
@@ -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."
+124
View File
@@ -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..<nClasses: # recover the class rule
echo spec.describeClause(necessaryLiterals(m, samples, cl), cl)
# 3. per-feature contribution + DEAD-INPUT LIST
let contribs = featureContributions(m, samples, spec)
let dead = deadInputs(contribs) # weighted < 5% of top, or never used
let strict = neverUsedInputs(contribs) # appearances == 0
let consts = constantInputs(samples, nBits)
for c in rankedInputs(contribs)[0..<10]: echo c.name, " ", c.weighted
# 4. accuracy diagnostics
let ad = accuracyDiagnostics(m, samples)
# ad.acc, ad.majorityBaseline, ad.margin, ad.confusion,
# ad.perClassRecall/Precision, ad.predMajorityShare
# 5. learning curve
let lc = learningCurve(tmplMachine, train, eval, nPoints = 10)
# lc.points, lc.accs, lc.trend ("flat" | "rising" | "rising-then-falling ...")
# 6. ablation hooks
let base = ablateBaseline(tmplMachine, train, eval)
for r in ablateDropAllBlocks(tmplMachine, train, eval, spec): echo r.name, r.delta
let rs = ablateScrambleFeature(tmplMachine, train, eval, spec, bit)
```
`ablation` retrains a fresh machine per variant (the strongest form of "does this
input earn its bits"). Pass `baselineAcc` from `ablateBaseline` to avoid
recomputing it per block.
## The default-off real-gun hook
`guns/tm_pattern.nim` gained only additive, default-off instrumentation:
```nim
g.diagCapture = true # default false; no behaviour change when false
# ... replay ...
g.diagSamples # seq[TmDiagSample] (literal vector + label)
g.exportTeams() # read-only GF clause teams
g.exportRadTeams(); g.exportRevTeams()
```
To introspect an externally trained clause set:
```nim
let m = machineFromTeams(TM_NBITS, TM_CLASSES, TM_NCLAUSES, TM_NSTATES, TM_S,
g.exportTeams())
```
## Running the demos / tests
```sh
nim c -r -d:release --path:common_libs common_libs/tests/test_tm_diag.nim # 48 pure unit checks
nim c -r -d:release --path:common_libs common_libs/tests/diag_synthetic.nim # Task 3 proof
nim c -r -d:release --path:common_libs common_libs/tests/diag_tm_pattern_offline.nim # Task 4 real reading
```
See `common_libs/tests/diag_synthetic.nim` for the ground-truth validation and
`common_libs/tests/diag_tm_pattern_offline.nim` for the real reading.
+589
View File
@@ -0,0 +1,589 @@
## tm_diag/diagnostics.nim — THE DIAGNOSTICS KIT (Task 2, six groups).
##
## Everything runs OFFLINE: a trained `TmMachine` plus a labelled sample set.
## No battles, no harness, no Java. The six groups are:
## 1. pre-flight DATA checks (dataChecks, shuffledLabelControl)
## 2. clause introspection (clauseInfo, clauseSummary, topClauses)
## 3. per-feature contribution (featureContributions, deadInputs, ...)
## 4. accuracy diagnostics (accuracyDiagnostics)
## 5. learning curve (learningCurve)
## 6. ablation hooks (ablateDropBlock, ablateScrambleFeature, ...)
##
## A `DiagSample` carries the LITERAL vector (positive literals [0..nBits),
## negations [nBits..2*nBits)) exactly as the shipped TM stores it, plus the
## true label. This keeps the kit usable both on the reference core and on the
## literal vectors captured from the real gun.
import std/[math, strformat, algorithm]
import feature_spec
import tm_core
export feature_spec, tm_core
type
DiagSample* = object
lits*: seq[uint8]
label*: int
order*: int ## optional capture/time index (for post-hoc curves)
# ── sample construction ──────────────────────────────────────────────────────
proc makeSample*(nBits: int, rawBits: openArray[int], label: int,
order = 0): DiagSample =
## Expand `nBits` raw bits into the pos/neg literal vector.
doAssert rawBits.len >= nBits,
"rawBits.len (" & $rawBits.len & ") < nBits (" & $nBits & ")"
result.lits = newSeq[uint8](2 * nBits)
for i in 0..<nBits:
let v = uint8(if rawBits[i] != 0: 1 else: 0)
result.lits[i] = v
result.lits[i + nBits] = 1'u8 - v
result.label = label
result.order = order
proc makeSample*(nBits: int, rawBits: openArray[uint8], label: int,
order = 0): DiagSample =
result.lits = newSeq[uint8](2 * nBits)
for i in 0..<nBits:
let v = uint8(if rawBits[i] != 0'u8: 1 else: 0)
result.lits[i] = v
result.lits[i + nBits] = 1'u8 - v
result.label = label
result.order = order
proc machineFromTeams*(nBits, nClasses, nClauses, nStates: int, sValue: float,
teams: openArray[seq[int16]]): TmMachine =
## Build an introspectable machine around an externally-trained clause set
## (e.g. the shipped tm_pattern gun's `teams`).
result = newMachine(nBits, nClasses, nClauses, nStates, sValue)
doAssert teams.len == nClasses
for c in 0..<nClasses:
result.teams[c] = teams[c]
# ── shared training / scoring helpers (used by every group below) ────────────
proc shuffledIndices(n: int, rng: var TmRng): seq[int] =
result = newSeq[int](n)
for i in 0..<n: result[i] = i
for i in countdown(n - 1, 1):
let j = int(rng.nextU64() mod uint64(i + 1))
swap(result[i], result[j])
proc trainModel*(tmpl: TmMachine, samples: openArray[DiagSample],
epochs = 1, seed = 777'u64): TmMachine =
## Fresh machine, `epochs` shuffled passes over `samples`.
result = tmpl
result.resetMachine(seed)
for _ in 0..<epochs:
let idx = shuffledIndices(samples.len, result.rng)
for k in idx:
result.trainSample(samples[k].lits, samples[k].label)
proc evalAcc*(m: TmMachine, samples: openArray[DiagSample]): float =
var c = 0
for s in samples:
if m.predictClass(s.lits) == s.label: inc c
if samples.len == 0: 0.0 else: c.float / samples.len.float
# ─────────────────────────────────────────────────────────────────────────────
# GROUP 1 — pre-flight DATA checks
# ─────────────────────────────────────────────────────────────────────────────
type
DataCheck* = object
n*: int
nClasses*: int
classCounts*: seq[int]
classShares*: seq[float]
majorityClass*: int
majorityShare*: float
majorityAccuracy*: float
maxShare*: float
overThreshold*: bool
threshold*: float
flags*: seq[string]
proc dataChecks*(labels: openArray[int], nClasses: int,
threshold = 0.30): DataCheck =
## Per-class label counts/shares, the majority-class share (which is also the
## majority-class accuracy baseline), and a loud flag when any class exceeds
## `threshold` of the samples — a too-large class compromises the accuracy
## test.
result.threshold = threshold
result.nClasses = nClasses
result.n = labels.len
result.classCounts = newSeq[int](nClasses)
result.classShares = newSeq[float](nClasses)
for l in labels:
if l >= 0 and l < nClasses: inc result.classCounts[l]
var maj = 0
for c in 0..<nClasses:
result.classShares[c] =
if result.n > 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..<nClasses:
if result.classShares[c] > 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..<samples.len: labels[i] = samples[i].label
for i in countdown(samples.len - 1, 1):
let j = int(rng.nextU64() mod uint64(i + 1))
swap(labels[i], labels[j])
for i in 0..<samples.len: shuf[i].label = labels[i]
let m = trainModel(tmpl, shuf, epochs, seed)
result.acc = evalAcc(m, shuf)
result.majority = dataChecks(labels, tmpl.nClasses).majorityShare
result.n = samples.len
result.nClasses = tmpl.nClasses
result.trainedObs = epochs * samples.len
# ─────────────────────────────────────────────────────────────────────────────
# GROUP 2 — clause introspection
# ─────────────────────────────────────────────────────────────────────────────
type
ClauseInfo* = object
cls*: int
index*: int
polarity*: int
length*: int
lits*: seq[int]
votes*: int ## samples in which the clause fires (casts a vote)
text*: string ## describeClause rendering
ClauseSummary* = object
totalClauses*: int
emptyClauses*: int
nonEmpty*: int
firedAtLeastOnce*: int
neverFired*: int
posFired*: int
negFired*: int
lengthHist*: seq[int] ## index = clause length, value = count (non-empty)
meanLength*: float
maxLength*: int
proc clauseInfo*(m: TmMachine, samples: openArray[DiagSample],
spec: FeatureSpec): seq[ClauseInfo] =
for c in 0..<m.nClasses:
for cl in 0..<m.nClauses:
let ls = m.clauseLits(c, cl)
var votes = 0
for s in samples:
if m.clauseFires(c, cl, s.lits): inc votes
result.add ClauseInfo(cls: c, index: cl,
polarity: (if cl < m.half: 1 else: -1), length: ls.len, lits: ls,
votes: votes, text: spec.describeClause(ls, c))
proc clauseSummary*(infos: openArray[ClauseInfo]): ClauseSummary =
var maxLen = 0
var lenSum = 0
for inf in infos:
inc result.totalClauses
if inf.length == 0:
inc result.emptyClauses
else:
inc result.nonEmpty
lenSum += inf.length
if inf.length > 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: result[c].cls = c
for inf in infos:
if inf.cls < 0 or inf.cls >= 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..<m.half: # positive polarity only
let ls = m.clauseLits(cls, cl)
if ls.len == 0: continue
var fires = false
for s in samples:
if m.clauseFires(cls, cl, s.lits):
fires = true
break
if not fires: continue
if first:
result = ls
first = false
else:
var keep: seq[int]
for l in result:
if l in ls: keep.add l
result = keep
result.sort()
# ─────────────────────────────────────────────────────────────────────────────
# GROUP 3 — per-feature contribution + DEAD-INPUT LIST
# ─────────────────────────────────────────────────────────────────────────────
type
FeatureContribution* = object
bit*: int
name*: string
appearances*: int ## times the bit appears in a clause that casts a vote
weighted*: float ## sum of clause vote-shares (1/nCasting per fire)
proc featureContributions*(m: TmMachine, samples: openArray[DiagSample],
spec: FeatureSpec): seq[FeatureContribution] =
## For every raw bit: how often its positive or negated literal appears in a
## clause that ACTUALLY CASTS A VOTE (non-empty + fires on a sample), weighted
## by that clause's share of the sample's casting votes.
result = newSeq[FeatureContribution](m.nBits)
for b in 0..<m.nBits:
result[b].bit = b
result[b].name = spec.describe(b)
var allLits = newSeq[seq[int]](m.nClasses * m.nClauses)
for c in 0..<m.nClasses:
for cl in 0..<m.nClauses:
allLits[c * m.nClauses + cl] = m.clauseLits(c, cl)
for s in samples:
var casting: seq[int]
for k in 0..<allLits.len:
let ls = allLits[k]
if ls.len == 0: continue
var fires = true
for lit in ls:
if s.lits[lit] == 0'u8:
fires = false
break
if fires: casting.add k
if casting.len == 0: continue
let w = 1.0 / float(casting.len)
for k in casting:
for lit in allLits[k]:
let b = if lit < m.nBits: lit else: lit - m.nBits
inc result[b].appearances
result[b].weighted += w
proc deadInputs*(contribs: openArray[FeatureContribution],
relThreshold = 0.05): seq[int] =
## THE DEAD-INPUT LIST. A bit is dead if it never appears in a voting clause
## (`appearances == 0`) OR its weighted vote-share contribution is below
## `relThreshold` x the largest contribution.
##
## The relative form is the practical one: a CONVERGED multiclass TM keeps
## redundant literals inside otherwise-correct clauses (each fires rarely, so
## its vote share is tiny). Strict `appearances == 0` therefore misses bits
## that are effectively dead; pass `relThreshold = 0.0` for the strict list.
var mx = 0.0
for c in contribs:
if c.weighted > 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..<nBits:
let v0 = samples[0].lits[b]
var constant = true
for i in 1..<samples.len:
if samples[i].lits[b] != v0:
constant = false
break
if constant: result.add b
proc rankedInputs*(contribs: seq[FeatureContribution]): seq[FeatureContribution] =
## Most valuable first.
result = contribs
result.sort(proc(a, b: FeatureContribution): int =
result = cmp(b.weighted, a.weighted)
if result == 0: result = cmp(b.appearances, a.appearances)
if result == 0: result = cmp(a.bit, b.bit))
# ─────────────────────────────────────────────────────────────────────────────
# GROUP 4 — accuracy diagnostics
# ─────────────────────────────────────────────────────────────────────────────
type
AccDiag* = object
n*: int
correct*: int
acc*: float
confusion*: seq[seq[int]] ## [true][pred]
majorityClass*: int
majorityShare*: float
majorityBaseline*: float
margin*: float ## acc - majority baseline (pp as fraction)
perClassRecall*: seq[float]
perClassPrecision*: seq[float]
predCounts*: seq[int]
predMajorityShare*: float ## how often the model predicts the majority class
proc accuracyDiagnostics*(m: TmMachine,
samples: openArray[DiagSample]): AccDiag =
result.confusion = newSeq[seq[int]](m.nClasses)
for c in 0..<m.nClasses: result.confusion[c] = newSeq[int](m.nClasses)
result.predCounts = newSeq[int](m.nClasses)
result.perClassRecall = newSeq[float](m.nClasses)
result.perClassPrecision = newSeq[float](m.nClasses)
result.n = samples.len
for s in samples:
let p = m.predictClass(s.lits)
if s.label >= 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..<m.nClasses:
for p in 0..<m.nClasses: rowSum[c] += result.confusion[c][p]
var maj = 0
for c in 0..<m.nClasses:
if rowSum[c] > 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..<m.nClasses:
result.perClassRecall[c] =
if rowSum[c] > 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..<first + count:
result.lits[b] = 0'u8
result.lits[b + nBits] = 1'u8
proc scrambleBit(samples: var seq[DiagSample], nBits, bit: int,
rng: var TmRng) =
let n = samples.len
if n < 2: return
var vals = newSeq[uint8](n)
for i in 0..<n: vals[i] = samples[i].lits[bit]
for i in countdown(n - 1, 1):
let j = int(rng.nextU64() mod uint64(i + 1))
swap(vals[i], vals[j])
for i in 0..<n:
samples[i].lits[bit] = vals[i]
samples[i].lits[bit + nBits] = 1'u8 - vals[i]
proc ablateDropBlock*(tmpl: TmMachine, train, eval: openArray[DiagSample],
spec: FeatureSpec, blockIdx: int, epochs = 1,
seed = 777'u64, baselineAcc = -1.0): AblationResult =
## Retrain from scratch with one whole block held at 0 (feature removed) and
## report the accuracy delta. ~zero => 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..<spec.blocks.len:
result.add ablateDropBlock(tmpl, train, eval, spec, i, epochs, seed)
proc ablateBaseline*(tmpl: TmMachine, train, eval: openArray[DiagSample],
epochs = 1, seed = 777'u64): AblationResult =
AblationResult(name: "baseline", kind: "baseline",
baselineAcc: evalAcc(trainModel(tmpl, train, epochs, seed), eval),
ablatedAcc: evalAcc(trainModel(tmpl, train, epochs, seed), eval),
delta: 0.0)
+148
View File
@@ -0,0 +1,148 @@
## tm_diag/feature_spec.nim — NAMED FEATURE CONTAINER (Task 1).
##
## The diagnostic kit is worthless if it prints "feature 17". A `FeatureSpec`
## is an ORDERED list of feature blocks, each with a human name, a bit range and
## (optionally) a name per bit. `describe` turns a single raw bit into its
## readable name; `describeClause` renders a Tsetlin conjunction as a sentence.
##
## Everything here is pure and offline — no battles, no harness.
import std/[strutils]
type
FeatureBlock* = object
## A contiguous run of raw bits forming one logical feature (often one-hot
## bins). `first` is the global raw-bit index of bit 0 of the block.
name*: string
first*: int
count*: int
bitNames*: seq[string] ## optional; len == count for per-bit names
FeatureSpec* = object
## An ordered list of blocks covering `nBits` raw bits.
nBits*: int
blocks*: seq[FeatureBlock]
proc addBlock*(s: var FeatureSpec, name: string, count: int,
bitNames: seq[string] = @[]) =
## Append a block; its first bit is the current end of the spec.
doAssert count > 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"])
+173
View File
@@ -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..<nClasses:
result.teams[c] = newSeq[int16](nClauses * result.nLiterals)
proc cloneMachine*(m: TmMachine): TmMachine =
result = m
result.teams = newSeq[seq[int16]](m.nClasses)
for c in 0..<m.nClasses:
result.teams[c] = m.teams[c]
proc resetMachine*(m: var TmMachine, seed: uint64) =
## Wipe every clause back to the Exclude boundary and reseed the RNG.
for c in 0..<m.nClasses:
for i in 0..<m.teams[c].len: m.teams[c][i] = 0
m.rng = seedRng(seed)
proc tmPolarity*(m: TmMachine, cl: int): float {.inline.} =
if cl < m.half: 1.0 else: -1.0
proc tmEval*(m: TmMachine, team: seq[int16], lits: openArray[uint8],
cl: int, learning: bool): uint8 =
let base = cl * m.nLiterals
var hasInc = false
for lit in 0..<m.nLiterals:
if team[base + lit] > 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..<m.nClauses:
let o = tmEval(m, team, lits, cl, learning = false)
cache[cl] = tmEval(m, team, lits, cl, learning = true)
v += tmPolarity(m, cl) * float(o)
clamp(v, -float(m.half), float(m.half))
proc tmLearnDir*(m: var TmMachine, team: var seq[int16],
lits: openArray[uint8], cache: seq[uint8],
vote, d: float) =
## One Granmo update of one class team with desired vote direction `d`.
let T = float(m.half)
let pFeedback = (T - d * vote) / (2.0 * T)
if pFeedback <= 0.0: return
for cl in 0..<m.nClauses:
if m.rng.rand01() >= 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..<m.nLiterals:
var st = int(team[base + lit])
if lits[lit] == 1'u8:
if cOut == 1'u8:
if m.rng.rand01() < (m.sValue - 1.0) / m.sValue:
st = min(st + 1, m.nStates)
else:
if m.rng.rand01() < 1.0 / m.sValue:
st = max(st - 1, -m.nStates)
else:
if m.rng.rand01() < 1.0 / m.sValue:
st = max(st - 1, -m.nStates)
team[base + lit] = int16(st)
else:
# Type II (Table 3): penalise exclusion of a zero literal when firing.
if cOut == 1'u8:
for lit in 0..<m.nLiterals:
if lits[lit] == 0'u8:
if team[base + lit] <= 0:
team[base + lit] = int16(min(int(team[base + lit]) + 1, m.nStates))
proc votesOf*(m: TmMachine, lits: openArray[uint8]): seq[float] =
result = newSeq[float](m.nClasses)
var cache = newSeq[uint8](m.nClauses)
for c in 0..<m.nClasses:
result[c] = tmForward(m, m.teams[c], lits, cache)
proc predictClass*(m: TmMachine, lits: openArray[uint8]): int =
## Hard argmax over the per-class votes; ties break to the lowest class.
var best = 0
var bestV = -Inf
var cache = newSeq[uint8](m.nClauses)
for c in 0..<m.nClasses:
let v = tmForward(m, m.teams[c], lits, cache)
if v > 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..<m.nClasses:
caches[c] = newSeq[uint8](m.nClauses)
votes[c] = tmForward(m, m.teams[c], lits, caches[c])
for c in 0..<m.nClasses:
let d = if c == label: 1.0 else: -1.0
tmLearnDir(m, m.teams[c], lits, caches[c], votes[c], d)
proc clauseFires*(m: TmMachine, cls, cl: int, lits: openArray[uint8]): bool =
## Classification-mode output of one clause: true iff it is non-empty and
## every included literal is 1.
let base = cl * m.nLiterals
var hasInc = false
for lit in 0..<m.nLiterals:
if m.teams[cls][base + lit] > 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..<m.nLiterals:
if m.teams[cls][base + lit] > 0: result.add lit