"""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 # '

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