File size: 2,589 Bytes
c91c9ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Odzysk tokenow edu z blendu fork-B (RB4 runda 1) + naprawa separatora.

Blend fork-B = przetasowane bloki: arcmix[0:~1.15B] (separator dokumentu 12285, zero 12287)
oraz dokumenty edu zakonczone 12287 (<|im_end|>, blad buildu: eot=vocab-1).

Dzielimy strumien na kawalki konczace sie 12287. Kawalek bez 12285 i bez 12286 = czysty
dokument edu -> zapis z separatorem podmienionym na 12285. Kawalek zawierajacy 12285/12286
niesie blok arcmixu -> odrzucony w calosci (traci sie co najwyzej jeden dokument edu na blok).

Wyjscie: edu.bin (uint16) + edu.json (liczby do weryfikacji wobec logu buildu).
"""
import argparse
import hashlib
import json

import numpy as np

EOS, IM_START, IM_END = 12285, 12286, 12287
CHUNK = 200_000_000


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--blend", required=True)
    ap.add_argument("--out", required=True)
    a = ap.parse_args()

    b = np.memmap(a.blend, dtype=np.uint16, mode="r")
    n = len(b)
    out = open(a.out, "wb")
    sha = hashlib.sha256()
    kept_docs = kept_tok = drop_pieces = drop_tok = 0
    carry = np.empty(0, dtype=np.uint16)
    for lo in range(0, n, CHUNK):
        x = np.concatenate([carry, np.asarray(b[lo:lo + CHUNK])])
        ends = np.flatnonzero(x == IM_END)
        if lo + CHUNK >= n and (len(ends) == 0 or ends[-1] != len(x) - 1):
            tail = len(x) - (ends[-1] + 1 if len(ends) else 0)
            drop_pieces += 1
            drop_tok += tail
        if len(ends) == 0:
            carry = x
            continue
        starts = np.concatenate([[0], ends[:-1] + 1])
        bad = np.zeros(len(ends), dtype=bool)
        for code in (EOS, IM_START):
            pos = np.flatnonzero(x[:ends[-1] + 1] == code)
            bad[np.searchsorted(ends, pos)] = True
        seg = x[:ends[-1] + 1].copy()
        seg[ends] = EOS
        keep_mask = np.repeat(~bad, ends - starts + 1)
        kept = seg[keep_mask]
        kept.tofile(out)
        sha.update(kept.tobytes())
        kept_docs += int((~bad).sum())
        kept_tok += int(len(kept))
        drop_pieces += int(bad.sum())
        drop_tok += int((ends - starts + 1)[bad].sum())
        carry = x[ends[-1] + 1:]
    out.close()
    stats = {"blend_tokens": int(n), "edu_docs": kept_docs, "edu_tokens": kept_tok,
             "dropped_pieces": drop_pieces, "dropped_tokens": drop_tok,
             "separator": EOS, "sha256": sha.hexdigest()}
    json.dump(stats, open(a.out.rsplit(".", 1)[0] + ".json", "w"), indent=1)
    print(json.dumps(stats), flush=True)


if __name__ == "__main__":
    main()