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()
|