diff options
Diffstat (limited to 'src/data.py')
| -rw-r--r-- | src/data.py | 31 |
1 files changed, 27 insertions, 4 deletions
diff --git a/src/data.py b/src/data.py index e6ea57b..fab4c21 100644 --- a/src/data.py +++ b/src/data.py @@ -44,21 +44,44 @@ def decode_tokens(ts, cfg: Config) -> int: 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, 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)) - return sorted(inputs[n_val:]), sorted(inputs[:n_val]) + 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)) + # deterministic subsample: seeded shuffle, take first n_tr + 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), ...]""" + """[(input_tokens, target_tokens+EOS), ...]. task_mode selects the target function.""" primes = sieve_primes(cfg.range_end + 100) out = [] for n in inputs: - p = next_prime(n, primes) + if cfg.task_mode == "is_prime": + p = 1 if is_prime_n(n) else 0 + else: + p = next_prime(n, primes) out.append((encode_int(n, cfg), encode_int(p, cfg) + [cfg.eos_id])) return out |
