summaryrefslogtreecommitdiff
path: root/src/data.py
blob: 411ba7e93a93cc98dea9ce7ffb57ca2dd8a574b9 (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
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
"""Prime dataset: n -> next prime, digit-tokenized; splits and batching."""
import bisect
import random

import torch

from src.config import Config


# Special token constants for digits mode (Addendum 8, E8)
EOS_ID = 10
SEP_ID = 11            # '#' separator between scratchpad and final target
PAUSE_ID = 12          # '<p>' pause/filler token (Arm C)
RAND_START_ID = 13     # 'a' (token IDs 13..28 represent 'a'..'p' for Arms D1/D2)
SYM_C = 29             # 'c'
SYM_EQ = 30            # '='
SYM_D = 31             # 'd'
SYM_COLON = 32         # ':'
NOISE_SLOT_ID = 33     # '<noise>' placeholder for Arm D3 continuous dynamic noise


def sieve_primes(limit: int) -> list[int]:
    """All primes <= limit (inclusive)."""
    if limit < 2:
        return []
    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
    return [i for i in range(2, limit + 1) if is_prime[i]]


def next_prime(n: int, primes: list[int]) -> int:
    idx = bisect.bisect_right(primes, n)
    if idx < len(primes):
        return primes[idx]
    raise ValueError(f"no prime > {n} in supplied list")


def encode_int(n: int, cfg: Config) -> list[int]:
    if cfg.vocab_mode == "integers":
        return [n]
    return [int(d) for d in str(n)]


def make_structured_trace(n: int, target_prime: int, primes_list: list[int]) -> list[int]:
    """Generates trace tokens for candidate search and trial division:
    For each candidate c in [n+1 .. target_prime]:
      c, =, digits(c), d, digits(p), :, 0/1, ...
    """
    tokens = []
    for c in range(n + 1, target_prime + 1):
        tokens.append(SYM_C)
        tokens.append(SYM_EQ)
        tokens.extend([int(d) for d in str(c)])
        for p in primes_list:
            if p * p > c:
                break
            tokens.append(SYM_D)
            tokens.extend([int(d) for d in str(p)])
            tokens.append(SYM_COLON)
            if c % p == 0:
                tokens.append(0)  # divisible -> composite found, halt checks for c
                break
            else:
                tokens.append(1)  # not divisible -> check next prime
    return tokens


def decode_tokens(ts, cfg: Config) -> int:
    """Decode a token sequence, stopping at EOS. If SEP_ID is present, decodes digits AFTER SEP_ID.
    -1 if nothing decodable."""
    if cfg.vocab_mode == "integers":
        return int(ts[0]) if len(ts) else -1

    if cfg.scratch_mode != "none":
        if SEP_ID in ts:
            idx = ts.index(SEP_ID)
            ts = ts[idx + 1:]
        else:
            return -1

    digits = []
    for t in ts:
        if t == cfg.eos_id:
            break
        if 0 <= t <= 9:
            digits.append(int(t))
        else:
            break
    return int("".join(map(str, digits))) if digits else -1


def is_prime_n(n: int) -> bool:
    """Exact primality for n >= 2."""
    if n < 2:
        return False
    d = 2
    while d * d <= n:
        if n % d == 0:
            return False
        d += 1
    return True


def get_splits(cfg: Config) -> tuple[list[int], list[int]]:
    """(train, val) input lists, seeded shuffle, no overlap. train_frac subsamples the
    TRAIN split only (E5); the val split is untouched (its size is locked by prereg)."""
    rng = random.Random(cfg.seed)
    inputs = list(range(cfg.range_start, cfg.range_end + 1))
    rng.shuffle(inputs)
    n_val = max(1, round(len(inputs) * cfg.holdout_frac))
    train, val = sorted(inputs[n_val:]), sorted(inputs[:n_val])
    if cfg.train_frac < 1.0:
        n_tr = max(1, round(len(train) * cfg.train_frac))
        sub = random.Random(cfg.seed + 1000)   # distinct stream from split shuffle
        sub.shuffle(train)
        train = sorted(train[:n_tr])
    return train, val


def build_examples(inputs: list[int], cfg: Config) -> list[tuple[list[int], list[int]]]:
    """[(input_tokens, target_tokens+EOS), ...]. task_mode and scratch_mode select format."""
    margin = max(100, int(cfg.range_end * 0.05) + 50)
    primes = sieve_primes(cfg.range_end + margin) if (cfg.task_mode != "is_prime" or cfg.scratch_mode == "structured") else []
    out = []
    for n in inputs:
        x_toks = encode_int(n, cfg)
        if cfg.task_mode == "is_prime":
            p = 1 if is_prime_n(n) else 0
            ans_toks = encode_int(p, cfg) + [cfg.eos_id]
            if cfg.scratch_mode == "none":
                out.append((x_toks, ans_toks))
            elif cfg.scratch_mode == "filler":
                filler = [PAUSE_ID] * cfg.scratch_len
                y_toks = filler + [SEP_ID] + ans_toks
                out.append((x_toks, y_toks))
            else:
                out.append((x_toks, ans_toks))
        else:
            p = next_prime(n, primes)
            ans_toks = encode_int(p, cfg) + [cfg.eos_id]
            if cfg.scratch_mode == "none":
                out.append((x_toks, ans_toks))
            elif cfg.scratch_mode == "structured":
                trace = make_structured_trace(n, p, primes)
                y_toks = trace + [SEP_ID] + ans_toks
                out.append((x_toks, y_toks))
            elif cfg.scratch_mode == "filler":
                filler = [PAUSE_ID] * cfg.scratch_len
                y_toks = filler + [SEP_ID] + ans_toks
                out.append((x_toks, y_toks))
            elif cfg.scratch_mode in ("random_learned", "random_frozen"):
                rng = random.Random(cfg.seed + n * 37)
                rand_toks = [RAND_START_ID + rng.randint(0, 15) for _ in range(cfg.scratch_len)]
                y_toks = rand_toks + [SEP_ID] + ans_toks
                out.append((x_toks, y_toks))
            elif cfg.scratch_mode == "random_noise":
                noise_toks = [NOISE_SLOT_ID] * cfg.scratch_len
                y_toks = noise_toks + [SEP_ID] + ans_toks
                out.append((x_toks, y_toks))
            else:
                out.append((x_toks, ans_toks))
    return out


def _global_lengths(cfg: Config) -> tuple[int, int]:
    """(in_max, out_max): fixed global lengths so batch layout == singleton layout."""
    if cfg.vocab_mode == "integers":
        return 1, 2
    in_max = len(str(cfg.range_end))
    if cfg.task_mode == "is_prime":
        base_out = 2
    else:
        margin = max(100, int(cfg.range_end * 0.05) + 50)
        primes = sieve_primes(cfg.range_end + margin)
        max_target = next_prime(cfg.range_end, primes)
        base_out = len(str(max_target)) + 1

    if cfg.scratch_mode == "none":
        return in_max, base_out
    elif cfg.scratch_mode in ("filler", "random_learned", "random_frozen", "random_noise"):
        return in_max, cfg.scratch_len + 1 + base_out
    elif cfg.scratch_mode == "structured":
        margin = max(100, int(cfg.range_end * 0.05) + 50)
        primes = sieve_primes(cfg.range_end + margin)
        # sample max trace length across the entire range
        max_trace_len = 0
        for n in range(cfg.range_start, cfg.range_end + 1):
            p = next_prime(n, primes)
            tr = make_structured_trace(n, p, primes)
            if len(tr) > max_trace_len:
                max_trace_len = len(tr)
        return in_max, max_trace_len + 1 + base_out
    return in_max, base_out


def pad_inputs(x: torch.Tensor, cfg: Config) -> torch.Tensor:
    """LEFT-pad inputs to the global in_max so absolute positions are layout-invariant."""
    in_max, _ = _global_lengths(cfg)
    if x.shape[1] < in_max:
        pad = torch.full((x.shape[0], in_max - x.shape[1]), cfg.pad_id, dtype=x.dtype, device=x.device)
        x = torch.cat([pad, x], dim=1)
    return x


def make_batch(examples, cfg: Config) -> dict[str, torch.Tensor]:
    """Fixed global layout: x left-padded to in_max, y right-padded to out_max."""
    in_max, out_max = _global_lengths(cfg)
    B = len(examples)
    x = torch.full((B, in_max), cfg.pad_id, dtype=torch.long)
    y = torch.full((B, out_max), cfg.pad_id, dtype=torch.long)
    loss_mask = torch.zeros((B, out_max), dtype=torch.bool)
    for i, (xi, yi) in enumerate(examples):
        x[i, in_max - len(xi):] = torch.tensor(xi, dtype=torch.long)
        y[i, : len(yi)] = torch.tensor(yi, dtype=torch.long)
        if cfg.scratch_mode in ("filler", "random_learned", "random_frozen", "random_noise"):
            loss_mask[i, cfg.scratch_len : len(yi)] = True
        else:
            loss_mask[i, : len(yi)] = True
    y_in = torch.cat([torch.full((B, 1), cfg.eos_id, dtype=torch.long), y[:, :-1]], dim=1)
    y_mask = y != cfg.pad_id
    return {"x": x, "y": y, "y_in": y_in, "y_mask": y_mask, "loss_mask": loss_mask}