summaryrefslogtreecommitdiff
path: root/src/data.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/data.py')
-rw-r--r--src/data.py31
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