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}
|