Files
SirRoboGarage/common_libs/tm_diag/diagnostics.nim
T
SirStone f9f8d84671 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.
2026-09-22 21:24:27 +02:00

590 lines
24 KiB
Nim

## 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)