TM diagnostics kit: VALIDATED (finds a known dead input), and it found a real bug
Built `common_libs/tm_diag/` as a first-class offline diagnostics kit for Tsetlin
work, BEFORE writing the new gun - because we hit two data problems tonight that no
amount of reading the TM's clauses would have revealed (a 38.8% majority answer,
and 36-58% mislabelled training samples).
WHAT IT PROVIDES
- `feature_spec.nim`: a NAMED feature container, so a learned clause prints as a
sentence (`IF near-wall AND bullet-dead-on AND turn-left(t-2) THEN class=3`)
instead of "feature 17". Includes the 49-bit draft spec from the design session
and the shipped 40-bit encoding.
- `tm_core.nim`: a compact deterministic Granmo multiclass TM with an
INTROSPECTABLE clause layout (mirrors the tm_pattern core).
- `diagnostics.nim`, six groups: (1) pre-flight DATA checks + shuffled-label
control, (2) clause introspection (readable dump, per-clause vote counts, empty
and never-fired clauses, length distribution, per-class balance), (3)
per-feature contribution with an explicit DEAD-INPUT LIST and a ranked
most-valuable list, (4) accuracy vs the majority baseline with per-class
precision/recall and pred-majority share, (5) learning curve, (6) ablation hooks
(drop a block / scramble a bit).
=== TASK 3: THE VALIDATION THAT GATES EVERYTHING - PASSED WITH NUMBERS ===
A diagnostic we never checked is worthless, so the kit was tested on a synthetic
set with a PLANTED RULE (class2 = A and B, class1 = A and not B, class0 = not A),
a deliberately IRRELEVANT block (US, 9 bits) and a PURE-NOISE bit (17).
- majority baseline 60.63% (class0); over-30% correctly flagged
- **the planted rule is recovered EXACTLY** via `necessaryLiterals`:
class0 IF NOT dist-wall<50 | class1 IF dist-wall<50 AND NOT lat DEAD-ON |
class2 IF dist-wall<50 AND lat DEAD-ON
- **DEAD-INPUT LIST = all 9 US bits AND the noise bit 17**, while the planted bits
0 and 45 are correctly NOT listed
- top contributors: bit0 w=1241.7, bit45 w=583.3, then 49.8 - a 12-25x gap, so the
relevant bits are unmistakable
- **ABLATION: drop WALLS -39.47pp, drop BULLETS -19.33pp, drop US 0.00pp**,
scramble A -42.00pp, scramble the noise bit 0.00pp
- shuffled-label control 60.40% vs majority 60.63% = -0.23pp -> no leak
So the kit reliably finds a known dead input and a known relevant one.
=== TASK 4: THE REAL READING, AND A BUG IN THE SHIPPED GUN ===
`tm_pattern` GF head, 6 DrussGT fixtures, pooled 250,745 samples:
- label balance c2 = **34.4%** (majority-heavy, flagged); accuracy **35.72%** vs
majority **34.24%** -> margin **+1.48pp**. On `tr_drussgt_vs_crazy` it is BELOW
majority (33.81% vs 37.72%, -3.92pp).
- 200 clauses: **27 empty, 45 never fired**, mean length 19.17, max 57. The
majority class is starved (class2: 22 non-empty, 18 empty, only 2 positive
fired). Class4 fires 11-24-literal clauses -> memorisation signature.
- **REPRESENTATION BUG FOUND (reported, not silently fixed):** `tmBuildBits` writes
only 38 raw bits into `var bits: array[TM_NBITS=40, uint8]` - bits 38 and 39 are
NEVER ASSIGNED, so they are always 0 and their negated literals are always 1.
The kit's `constantInputs` confirms 38/39 are constant, and **`UNUSED-38`/`39`
rank #6 and #8 in the most-valuable-inputs list** - i.e. the model's
highest-usage inputs are information-free. That is a representation bug, not a
display artefact, and it is a concrete mechanism for part of the poor learning.
DRAFT ENCODING CHECKED: the 49-bit draft is arithmetically consistent
(4+4=8 walls, 6+3=9 us, 3+5+3+3+3+3=20 motion, 5+7=12 bullets = 49). No draft
inconsistency.
Guards: test_tm_diag 48 (new), diag_synthetic 17 (new), test_gun_harness 39,
test_vbullet_metric 11, test_power_selection 3, test_adaptive_radar 41,
test_tfil_ring_weights 24, test_power_policy 26, test_ram_decision 40,
test_rack_membership 48, test_selector_tiebreak 19, test_tm_pattern_registration 20,
test_vbullet_admit_gate 12, acceptance_offline_vs_online 12/12; tm_pattern_learning
passes. The tm_pattern hook is additive and default-OFF (no behaviour change).
NOT YET INCLUDED (the automata metrics discussed for the next step): per-clause
automata settledness (distance from the flip point), clause diversity (pairwise
overlap), literal-set churn over time, and cross-clause vote disagreement. The kit
has clause-level diagnostics but not the automata-state ones.
This commit is contained in:
@@ -0,0 +1,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.
|
||||
@@ -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)
|
||||
@@ -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"])
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user