diff options
Diffstat (limited to 'src/data.py')
| -rw-r--r-- | src/data.py | 17 |
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 |
