summaryrefslogtreecommitdiff
path: root/src/data.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/data.py')
-rw-r--r--src/data.py17
1 files changed, 11 insertions, 6 deletions
diff --git a/src/data.py b/src/data.py
index fab4c21..26fb6b9 100644
--- a/src/data.py
+++ b/src/data.py
@@ -1,4 +1,5 @@
"""Prime dataset: n -> next prime, digit-tokenized; splits and batching."""
+import bisect
import random
import torch
@@ -20,9 +21,9 @@ def sieve_primes(limit: int) -> list[int]:
def next_prime(n: int, primes: list[int]) -> int:
- for p in primes:
- if p > n:
- return p
+ idx = bisect.bisect_right(primes, n)
+ if idx < len(primes):
+ return primes[idx]
raise ValueError(f"no prime > {n} in supplied list")
@@ -75,7 +76,8 @@ def get_splits(cfg: Config) -> tuple[list[int], list[int]]:
def build_examples(inputs: list[int], cfg: Config) -> list[tuple[list[int], list[int]]]:
"""[(input_tokens, target_tokens+EOS), ...]. task_mode selects the target function."""
- primes = sieve_primes(cfg.range_end + 100)
+ margin = max(100, int(cfg.range_end * 0.05) + 50)
+ primes = sieve_primes(cfg.range_end + margin) if cfg.task_mode != "is_prime" else []
out = []
for n in inputs:
if cfg.task_mode == "is_prime":
@@ -88,10 +90,13 @@ def build_examples(inputs: list[int], cfg: Config) -> list[tuple[list[int], list
def _global_lengths(cfg: Config) -> tuple[int, int]:
"""(in_max, out_max): fixed global lengths so batch layout == singleton layout (codex BLOCKER fix)."""
- primes = sieve_primes(cfg.range_end + 100)
- max_target = next_prime(cfg.range_end, primes)
if cfg.vocab_mode == "integers":
return 1, 2 # [value], [value, EOS]
+ if cfg.task_mode == "is_prime":
+ return len(str(cfg.range_end)), 2 # digit token "1"/"0" + EOS
+ 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)
return len(str(cfg.range_end)), len(str(max_target)) + 1 # digits + EOS