summaryrefslogtreecommitdiff
path: root/src/config.py
blob: 0b161e72af88441444790e1d175bc2f06745cc16 (plain)
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
"""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_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
    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: 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