summaryrefslogtreecommitdiff
path: root/src/config.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/config.py')
-rw-r--r--src/config.py101
1 files changed, 101 insertions, 0 deletions
diff --git a/src/config.py b/src/config.py
new file mode 100644
index 0000000..74cb4d3
--- /dev/null
+++ b/src/config.py
@@ -0,0 +1,101 @@
+"""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 off before this (anti-collapse)
+ 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