-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdata.py
More file actions
91 lines (64 loc) · 2.49 KB
/
Copy pathdata.py
File metadata and controls
91 lines (64 loc) · 2.49 KB
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
80
81
82
83
84
85
86
87
88
89
90
91
"""Sort dataset loader.
Each line of `data_files/{train,test}.txt` is: <input_digits>><sorted_digits>
e.g. 3729>2379
"""
from pathlib import Path
import numpy as np
import jax.numpy as jnp
DIGITS = "0123456789"
SEP_CHAR = ">"
END_CHAR = "$"
PAD_CHAR = "."
VOCAB = DIGITS + SEP_CHAR + END_CHAR + PAD_CHAR
VOCAB_SIZE = len(VOCAB)
C2I = {c: i for i, c in enumerate(VOCAB)}
I2C = {i: c for c, i in C2I.items()}
PAD_ID = C2I[PAD_CHAR]
SEP_ID = C2I[SEP_CHAR]
END_ID = C2I[END_CHAR]
DATA_DIR = Path(__file__).parent / "data_files"
SEQ_LEN = 4
N_TRAIN = 10000
def encode(s: str) -> list[int]:
return [C2I[c] for c in s]
def decode(ids) -> str:
return "".join(I2C[int(i)] for i in ids)
def build_tensors(inputs: np.ndarray, outputs: np.ndarray):
"""Concatenate input>output$ and produce (x, y, loss_mask) for next-token training."""
N, seq_len = inputs.shape
sep = np.full((N, 1), SEP_ID, dtype=np.int32)
end = np.full((N, 1), END_ID, dtype=np.int32)
full = np.concatenate([inputs, sep, outputs, end], axis=1) # (N, 2*seq_len + 2)
x = full[:, :-1] # (N, 2*seq_len + 1)
y = full[:, 1:] # (N, 2*seq_len + 1)
T = x.shape[1]
mask = np.zeros((N, T), dtype=np.float32)
mask[:, seq_len + 1:] = 1.0 # skip the separator; count loss on the output tokens
return x, y, mask
def causal_attention_mask(T):
"""Boolean (T, T) causal mask: row i may attend to columns up to position i."""
return jnp.tril(jnp.ones((T, T), dtype=bool), k=1)
def _load(split: str):
path = DATA_DIR / f"{split}.txt"
inputs, outputs = [], []
with path.open() as f:
for line in f:
inp_s, out_s = line.strip().split(SEP_CHAR)
inputs.append([C2I[c] for c in inp_s])
outputs.append([C2I[c] for c in out_s])
return np.array(inputs, dtype=np.int32), np.array(outputs, dtype=np.int32)
def load_train():
"""Generate training data: each example is SEQ_LEN numbers drawn from 0-9."""
rng = np.random.default_rng(0)
inputs = np.stack([
rng.choice(len(DIGITS), size=SEQ_LEN, replace=False) for _ in range(N_TRAIN)
]).astype(np.int32)
outputs = np.sort(inputs, axis=1)
return build_tensors(inputs, outputs)
def load_test():
"""Load test data from data_files/test.txt as raw (inputs, outputs)."""
return _load("test")
def make_train_batch(rng: np.random.Generator, train_data, batch_size: int):
x, y, mask = train_data
idx = rng.integers(0, len(x), size=batch_size)
return x[idx], y[idx], mask[idx]