diff options
Diffstat (limited to 'src/config.py')
| -rw-r--r-- | src/config.py | 101 |
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 |
