File size: 2,945 Bytes
96282ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""fork-B blend: doc-level 50:50 (parametrizable) mix of arcmix + edu -> train.bin (uint16).

Single-factor ablation (RB4/S10): fork-A=arcmix, fork-B=arcmix+edu. This builds fork-B's
train.bin by taking ~equal token budgets from each source, splitting on eot-doc-boundaries,
shuffling the tagged doc list (seeded) and concatenating. The trainer samples random windows
uniformly over the flat stream -> overall token exposure = arcmix:edu at the chosen ratio.

Val: fork-B should reuse arcmix_val.bin (same val as fork-A) so val-loss is comparable;
copy it into the out-dir separately. Only train.bin is produced here.
"""
import argparse
from pathlib import Path

import numpy as np


def doc_spans(arr, eot):
    """(start,end) spans split on eot (end exclusive, includes the eot token)."""
    idx = np.where(arr == eot)[0]
    prev = 0
    for e in idx:
        yield prev, int(e) + 1
        prev = int(e) + 1


def take_docs(arr, budget, eot):
    spans, tot = [], 0
    for s, e in doc_spans(arr, eot):
        spans.append((s, e))
        tot += e - s
        if tot >= budget:
            break
    return spans, tot


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--arcmix", required=True, help="path to arcmix train.bin (uint16)")
    ap.add_argument("--edu", required=True, help="path to edu train.bin (uint16)")
    ap.add_argument("--out", required=True, help="output blend train.bin")
    ap.add_argument("--arcmix-tokens", type=int, default=2_000_000_000)
    ap.add_argument("--edu-tokens", type=int, default=2_000_000_000)
    ap.add_argument("--eot", type=int, default=12287)
    ap.add_argument("--seed", type=int, default=1337)
    a = ap.parse_args()

    arc = np.memmap(a.arcmix, dtype=np.uint16, mode="r")
    edu = np.memmap(a.edu, dtype=np.uint16, mode="r")
    print(f"arcmix.bin={len(arc):,} tok | edu.bin={len(edu):,} tok", flush=True)

    arc_spans, arc_tot = take_docs(arc, a.arcmix_tokens, a.eot)
    edu_spans, edu_tot = take_docs(edu, a.edu_tokens, a.eot)
    print(f"take arcmix docs={len(arc_spans):,} tok={arc_tot:,} | "
          f"edu docs={len(edu_spans):,} tok={edu_tot:,}", flush=True)

    tagged = [(0, s, e) for s, e in arc_spans] + [(1, s, e) for s, e in edu_spans]
    rng = np.random.default_rng(a.seed)
    rng.shuffle(tagged)

    out = open(a.out, "wb")
    buf, buflen, written = [], 0, 0
    for src, s, e in tagged:
        buf.append(np.asarray((arc if src == 0 else edu)[s:e]))
        buflen += e - s
        if buflen >= 20_000_000:
            np.concatenate(buf).tofile(out)
            written += buflen
            buf, buflen = [], 0
    if buf:
        np.concatenate(buf).tofile(out)
        written += buflen
    out.close()
    frac = arc_tot / max(written, 1)
    print(f"WROTE {a.out} tok={written:,} | arcmix={arc_tot:,} ({frac:.1%}) edu={edu_tot:,} ({1-frac:.1%})",
          flush=True)


if __name__ == "__main__":
    main()