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
|
"""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"
range_start: int = 2
range_end: int = 100 # inclusive
holdout_frac: float = 0.30
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_eps: float = 0.05 # ACT cumulative threshold (documented; K small so no early break)
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 before this step
n_layers: int = 2
n_heads: int = 4
# training
lr: float = 1e-3
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
max_out_len: int = 6
log_n_examples: int = 10
out_dir: str = "runs"
@property
def vocab(self) -> int:
"""Token count. Integers mode: 0..range_end+1 (next prime can exceed range_end)."""
if self.vocab_mode == "integers":
return self.range_end + 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
|