summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config.py3
-rw-r--r--src/data.py31
-rw-r--r--src/eval.py30
-rw-r--r--src/train.py5
4 files changed, 56 insertions, 13 deletions
diff --git a/src/config.py b/src/config.py
index 0b161e7..60984db 100644
--- a/src/config.py
+++ b/src/config.py
@@ -8,9 +8,11 @@ from dataclasses import dataclass, fields
class Config:
# data
vocab_mode: str = "digits" # "digits" | "integers"
+ task_mode: str = "next_prime" # "next_prime" | "is_prime"
range_start: int = 2
range_end: int = 100 # inclusive
holdout_frac: float = 0.30
+ train_frac: float = 1.0 # 1.0 = use full train split; 0.4/0.5 for E5
seed: int = 0
# model
model: str = "rnn" # "rnn" | "transformer"
@@ -25,6 +27,7 @@ class Config:
n_heads: int = 4
# training
lr: float = 1e-3
+ lr_decay: bool = False # cosine 1e-3 -> 1e-4 over the run (E3)
weight_decay: float = 1.0
max_train_steps: int = 200_000
eval_every: int = 200
diff --git a/src/data.py b/src/data.py
index e6ea57b..fab4c21 100644
--- a/src/data.py
+++ b/src/data.py
@@ -44,21 +44,44 @@ def decode_tokens(ts, cfg: Config) -> int:
return int("".join(map(str, digits))) if digits else -1
+def is_prime_n(n: int) -> bool:
+ """Exact primality for n >= 2."""
+ if n < 2:
+ return False
+ d = 2
+ while d * d <= n:
+ if n % d == 0:
+ return False
+ d += 1
+ return True
+
+
def get_splits(cfg: Config) -> tuple[list[int], list[int]]:
- """(train, val) input lists, seeded shuffle, no overlap."""
+ """(train, val) input lists, seeded shuffle, no overlap. train_frac subsamples the
+ TRAIN split only (E5); the val split is untouched (its size is locked by prereg)."""
rng = random.Random(cfg.seed)
inputs = list(range(cfg.range_start, cfg.range_end + 1))
rng.shuffle(inputs)
n_val = max(1, round(len(inputs) * cfg.holdout_frac))
- return sorted(inputs[n_val:]), sorted(inputs[:n_val])
+ train, val = sorted(inputs[n_val:]), sorted(inputs[:n_val])
+ if cfg.train_frac < 1.0:
+ n_tr = max(1, round(len(train) * cfg.train_frac))
+ # deterministic subsample: seeded shuffle, take first n_tr
+ sub = random.Random(cfg.seed + 1000) # distinct stream from split shuffle
+ sub.shuffle(train)
+ train = sorted(train[:n_tr])
+ return train, val
def build_examples(inputs: list[int], cfg: Config) -> list[tuple[list[int], list[int]]]:
- """[(input_tokens, target_tokens+EOS), ...]"""
+ """[(input_tokens, target_tokens+EOS), ...]. task_mode selects the target function."""
primes = sieve_primes(cfg.range_end + 100)
out = []
for n in inputs:
- p = next_prime(n, primes)
+ if cfg.task_mode == "is_prime":
+ p = 1 if is_prime_n(n) else 0
+ else:
+ p = next_prime(n, primes)
out.append((encode_int(n, cfg), encode_int(p, cfg) + [cfg.eos_id]))
return out
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 {
diff --git a/src/train.py b/src/train.py
index aa0fb4b..406defd 100644
--- a/src/train.py
+++ b/src/train.py
@@ -166,6 +166,11 @@ def main() -> None:
halt = out["halt_steps"]
penalty = halt.float().mean() * lam if halt is not None and lam > 0 else 0.0
loss = loss_tokens + penalty
+ if cfg.lr_decay:
+ # cosine 1e-3 -> 1e-4 over the full run (E3, locked in Addendum 4)
+ frac = min(1.0, step / cfg.max_train_steps)
+ lr_t = cfg.lr * 0.1 + 0.5 * (cfg.lr - cfg.lr * 0.1) * (1 + math.cos(math.pi * frac))
+ opt.param_groups[0]["lr"] = lr_t
opt.zero_grad()
loss.backward()
opt.step()