diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config.py | 3 | ||||
| -rw-r--r-- | src/data.py | 31 | ||||
| -rw-r--r-- | src/eval.py | 30 | ||||
| -rw-r--r-- | src/train.py | 5 |
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() |
