summaryrefslogtreecommitdiff
path: root/src/data.py
blob: 35682d3de8f9741b763d8c73de63f3845d3dc96d (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""Prime dataset: n -> next prime, digit-tokenized; splits and batching."""
import random

import torch

from src.config import Config


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:
    for p in primes:
        if p > n:
            return p
    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 decode_tokens(ts, cfg: Config) -> int:
    """Decode a token sequence, stopping at EOS. -1 if nothing decodable."""
    if cfg.vocab_mode == "integers":
        return int(ts[0]) if len(ts) else -1
    digits = []
    for t in ts:
        if t == cfg.eos_id:
            break
        digits.append(int(t))
    return int("".join(map(str, digits))) if digits else -1


def get_splits(cfg: Config) -> tuple[list[int], list[int]]:
    """(train, val) input lists, seeded shuffle, no overlap."""
    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])


def build_examples(inputs: list[int], cfg: Config) -> list[tuple[list[int], list[int]]]:
    """[(input_tokens, target_tokens+EOS), ...]"""
    primes = sieve_primes(cfg.range_end + 100)
    out = []
    for n in inputs:
        p = next_prime(n, primes)
        out.append((encode_int(n, cfg), encode_int(p, cfg) + [cfg.eos_id]))
    return out


def make_batch(examples, cfg: Config) -> dict[str, torch.Tensor]:
    xs, ys = zip(*examples)
    T_in = max(len(x) for x in xs)
    T_out = max(len(y) for y in ys)
    B = len(examples)
    x = torch.full((B, T_in), cfg.pad_id, dtype=torch.long)
    y = torch.full((B, T_out), cfg.pad_id, dtype=torch.long)
    for i, (xi, yi) in enumerate(examples):
        x[i, : len(xi)] = torch.tensor(xi, dtype=torch.long)
        y[i, : len(yi)] = torch.tensor(yi, dtype=torch.long)
    # teacher-forced decoder input: BOS(=EOS reuse) then shifted y
    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}