diff options
| author | Void Agent <void@jayrup.hermes> | 2026-08-16 14:41:54 +0100 |
|---|---|---|
| committer | Void Agent <void@jayrup.hermes> | 2026-08-16 14:41:54 +0100 |
| commit | 751cbe0d93af1fcc69147d27ca8babd63105e4dc (patch) | |
| tree | 5a026310f5969db18e5284217ff18ad6bc7488ce /src/eval.py | |
| parent | ababeb8ec06d7f10f6e6ef556702bd55304bb190 (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.py | 30 |
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 { |
