1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
|
"""Experiment configuration. Defaults = Experiment 1 (design/preregistration.md)."""
import argparse
import json
from dataclasses import dataclass, fields
@dataclass
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"
d_model: int = 128
max_steps: int = 20 # K tied iterations (RNN)
halting: bool = True # False -> fixed-K ablation
halt_penalty: float = 0.01 # lambda on mean steps
halt_warmup_steps: int = 1000 # penalty = 0 before this (anti-collapse)
halt_ramp_end_steps: int = 5000 # penalty ramps linearly 0 -> halt_penalty between warmup and here
min_steps: int = 2 # ACT: halt prob forced to 0 for the first min_steps-1 steps
n_layers: int = 2
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
early_stop_em: float = 1.0
early_stop_patience: int = 5
batch_size: int = 32
eval_batch_size: int = 512 # chunked evaluation batch size to avoid GPU OOM
max_out_len: int = 6
log_n_examples: int = 10
out_dir: str = "runs"
device: str = "auto" # "auto" | "cuda" | "cpu" | "mps"
@property
def vocab(self) -> int:
"""Token count. Integers mode: value tokens 0..next_prime(range_end), plus EOS (+1) and pad (+1)."""
if self.vocab_mode == "integers":
# inline mini-sieve: next prime above range_end bounds the target domain
limit = self.range_end + 100
is_prime = [True] * (limit + 1)
is_prime[0] = is_prime[1] = False
for p in range(2, int(limit ** 0.5) + 1):
if is_prime[p]:
for m in range(p * p, limit + 1, p):
is_prime[m] = False
np_ = next(i for i in range(self.range_end + 1, limit + 1) if is_prime[i])
return np_ + 2 # values 0..np_, EOS=np_+1, pad=np_+2
return 11 # digits 0-9 + EOS
@property
def eos_id(self) -> int:
return self.vocab - 1
@property
def pad_id(self) -> int:
return self.vocab # one extra embedding row reserved for pad
def to_json(self) -> dict:
d = {f.name: getattr(self, f.name) for f in fields(self)}
d["vocab"] = self.vocab
d["eos_id"] = self.eos_id
d["pad_id"] = self.pad_id
return d
@classmethod
def from_json(cls, d: dict) -> "Config":
cfg = cls()
for f in fields(cls):
if f.name in d:
setattr(cfg, f.name, d[f.name])
return cfg
def save(self, path: str) -> None:
with open(path, "w") as fh:
json.dump(self.to_json(), fh, indent=2)
@classmethod
def load(cls, path: str) -> "Config":
with open(path) as fh:
return cls.from_json(json.load(fh))
def _bool_arg(s: str) -> bool:
return s.lower() in ("1", "true", "yes", "on")
def parse_args(argv=None) -> Config:
cfg = Config()
p = argparse.ArgumentParser(description="prime-grokking train")
p.add_argument("model_pos", nargs="?", default=None, help="model: rnn | transformer")
p.add_argument("seed_pos", nargs="?", default=None, help="seed (int)")
for f in fields(Config):
if f.type is bool:
p.add_argument(f"--{f.name}", default=None, type=_bool_arg)
elif f.type in (int, float, str):
p.add_argument(f"--{f.name}", default=None, type=f.type)
a = p.parse_args(argv)
if a.model_pos is not None:
cfg.model = a.model_pos
if a.seed_pos is not None:
cfg.seed = int(a.seed_pos)
for f in fields(Config):
v = getattr(a, f.name, None)
if v is not None:
setattr(cfg, f.name, v)
return cfg
|