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
+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