Upload blend_edu.py with huggingface_hub
Browse files- blend_edu.py +80 -0
blend_edu.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""fork-B blend: doc-level 50:50 (parametrizable) mix of arcmix + edu -> train.bin (uint16).
|
| 3 |
+
|
| 4 |
+
Single-factor ablation (RB4/S10): fork-A=arcmix, fork-B=arcmix+edu. This builds fork-B's
|
| 5 |
+
train.bin by taking ~equal token budgets from each source, splitting on eot-doc-boundaries,
|
| 6 |
+
shuffling the tagged doc list (seeded) and concatenating. The trainer samples random windows
|
| 7 |
+
uniformly over the flat stream -> overall token exposure = arcmix:edu at the chosen ratio.
|
| 8 |
+
|
| 9 |
+
Val: fork-B should reuse arcmix_val.bin (same val as fork-A) so val-loss is comparable;
|
| 10 |
+
copy it into the out-dir separately. Only train.bin is produced here.
|
| 11 |
+
"""
|
| 12 |
+
import argparse
|
| 13 |
+
from pathlib import Path
|
| 14 |
+
|
| 15 |
+
import numpy as np
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def doc_spans(arr, eot):
|
| 19 |
+
"""(start,end) spans split on eot (end exclusive, includes the eot token)."""
|
| 20 |
+
idx = np.where(arr == eot)[0]
|
| 21 |
+
prev = 0
|
| 22 |
+
for e in idx:
|
| 23 |
+
yield prev, int(e) + 1
|
| 24 |
+
prev = int(e) + 1
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def take_docs(arr, budget, eot):
|
| 28 |
+
spans, tot = [], 0
|
| 29 |
+
for s, e in doc_spans(arr, eot):
|
| 30 |
+
spans.append((s, e))
|
| 31 |
+
tot += e - s
|
| 32 |
+
if tot >= budget:
|
| 33 |
+
break
|
| 34 |
+
return spans, tot
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def main():
|
| 38 |
+
ap = argparse.ArgumentParser()
|
| 39 |
+
ap.add_argument("--arcmix", required=True, help="path to arcmix train.bin (uint16)")
|
| 40 |
+
ap.add_argument("--edu", required=True, help="path to edu train.bin (uint16)")
|
| 41 |
+
ap.add_argument("--out", required=True, help="output blend train.bin")
|
| 42 |
+
ap.add_argument("--arcmix-tokens", type=int, default=2_000_000_000)
|
| 43 |
+
ap.add_argument("--edu-tokens", type=int, default=2_000_000_000)
|
| 44 |
+
ap.add_argument("--eot", type=int, default=12287)
|
| 45 |
+
ap.add_argument("--seed", type=int, default=1337)
|
| 46 |
+
a = ap.parse_args()
|
| 47 |
+
|
| 48 |
+
arc = np.memmap(a.arcmix, dtype=np.uint16, mode="r")
|
| 49 |
+
edu = np.memmap(a.edu, dtype=np.uint16, mode="r")
|
| 50 |
+
print(f"arcmix.bin={len(arc):,} tok | edu.bin={len(edu):,} tok", flush=True)
|
| 51 |
+
|
| 52 |
+
arc_spans, arc_tot = take_docs(arc, a.arcmix_tokens, a.eot)
|
| 53 |
+
edu_spans, edu_tot = take_docs(edu, a.edu_tokens, a.eot)
|
| 54 |
+
print(f"take arcmix docs={len(arc_spans):,} tok={arc_tot:,} | "
|
| 55 |
+
f"edu docs={len(edu_spans):,} tok={edu_tot:,}", flush=True)
|
| 56 |
+
|
| 57 |
+
tagged = [(0, s, e) for s, e in arc_spans] + [(1, s, e) for s, e in edu_spans]
|
| 58 |
+
rng = np.random.default_rng(a.seed)
|
| 59 |
+
rng.shuffle(tagged)
|
| 60 |
+
|
| 61 |
+
out = open(a.out, "wb")
|
| 62 |
+
buf, buflen, written = [], 0, 0
|
| 63 |
+
for src, s, e in tagged:
|
| 64 |
+
buf.append(np.asarray((arc if src == 0 else edu)[s:e]))
|
| 65 |
+
buflen += e - s
|
| 66 |
+
if buflen >= 20_000_000:
|
| 67 |
+
np.concatenate(buf).tofile(out)
|
| 68 |
+
written += buflen
|
| 69 |
+
buf, buflen = [], 0
|
| 70 |
+
if buf:
|
| 71 |
+
np.concatenate(buf).tofile(out)
|
| 72 |
+
written += buflen
|
| 73 |
+
out.close()
|
| 74 |
+
frac = arc_tot / max(written, 1)
|
| 75 |
+
print(f"WROTE {a.out} tok={written:,} | arcmix={arc_tot:,} ({frac:.1%}) edu={edu_tot:,} ({1-frac:.1%})",
|
| 76 |
+
flush=True)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
if __name__ == "__main__":
|
| 80 |
+
main()
|