summaryrefslogtreecommitdiff
path: root/src/eval.py
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-08-16 14:41:54 +0100
committerVoid Agent <void@jayrup.hermes>2026-08-16 14:41:54 +0100
commit751cbe0d93af1fcc69147d27ca8babd63105e4dc (patch)
tree5a026310f5969db18e5284217ff18ad6bc7488ce /src/eval.py
parentababeb8ec06d7f10f6e6ef556702bd55304bb190 (diff)
phase3: task_mode (is_prime), train_frac, lr_decay flags + tests (41 green); Addendum 5 locked; 11-job batch
Diffstat (limited to 'src/eval.py')
-rw-r--r--src/eval.py30
1 files changed, 21 insertions, 9 deletions
diff --git a/src/eval.py b/src/eval.py
index 50ec40e..c964d38 100644
--- a/src/eval.py
+++ b/src/eval.py
@@ -14,7 +14,7 @@ import numpy as np
import torch
from src.config import Config
-from src.data import build_examples, decode_tokens, encode_int, get_splits, next_prime, sieve_primes
+from src.data import build_examples, decode_tokens, encode_int, get_splits, is_prime_n, next_prime, sieve_primes
from src.model_api import build_model, greedy_decode
from src.train import evaluate
@@ -31,15 +31,21 @@ def probe_report(model, cfg: Config, lo: int = 101, hi: int = 200) -> dict:
errors = []
easy_total = 0
easy_wrong = 0
+ is_prime_task = cfg.task_mode == "is_prime"
for n in range(lo, hi + 1):
x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0)
gen = greedy_decode(model, x, cfg)[0].tolist()
pred = decode_tokens(gen, cfg)
- target = next_prime(n, primes)
- is_easy = (n % 2 == 0) or (n % 5 == 0) # trivial composites (skip-evens / skip-5s)
+ if is_prime_task:
+ target = 1 if is_prime_n(n) else 0
+ ok = (pred == 1) == (target == 1) # any non-"1" output reads as "composite"
+ else:
+ target = next_prime(n, primes)
+ ok = pred == target
+ is_easy = (n % 2 == 0) or (n % 5 == 0) # trivial composites (skip-evens / skip-5s)
if is_easy:
easy_total += 1
- if pred == target:
+ if ok:
correct += 1
else:
errors.append({"n": n, "target": target, "pred": pred})
@@ -47,16 +53,22 @@ def probe_report(model, cfg: Config, lo: int = 101, hi: int = 200) -> dict:
easy_wrong += 1
total = hi - lo + 1
acc = correct / total
- flagged = [e for e in errors if e["pred"] in FLAGGED_SIEVE_PREDS]
- distinct_flagged = len({e["pred"] for e in flagged})
+ if is_prime_task:
+ flagged = [e for e in errors if e["n"] in FLAGGED_SIEVE_PREDS] # input IS the classified number
+ distinct_flagged = len({e["n"] for e in flagged})
+ else:
+ flagged = [e for e in errors if e["pred"] in FLAGGED_SIEVE_PREDS]
+ distinct_flagged = len({e["pred"] for e in flagged})
flagged_frac = len(flagged) / len(errors) if errors else 0.0
- # classification per prereg + Addendum 3 operationalization; P4 checked first
+ # classification per prereg + Addendum 3/5 operationalization
if easy_total and easy_wrong / easy_total > 0.5:
code = "P4" # fails trivial evens/5-multiples -> pure memorization
+ elif is_prime_task and len(errors) >= 3 and distinct_flagged >= 3 and flagged_frac >= 0.8:
+ code = "P1" # is-prime: errors concentrated on no-small-factor composites -> learned sieve
elif acc >= 0.85:
- code = "P3" # surprising success beyond expectation
+ code = "P3" # surprising success beyond expectation (next_prime semantics)
elif len(errors) >= 3 and distinct_flagged >= 3 and flagged_frac >= 0.8:
- code = "P1" # errors = sieves predicting no-small-factor composites -> learned {2,3,5,7} sieve
+ code = "P1" # next_prime: errors = sieves predicting no-small-factor composites
else:
code = "P2" # scattered errors -> memorization / non-transferable heuristics
return {