File size: 4,013 Bytes
1b929be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Text pool for embedding distillation: FineWeb-2 / FineWeb (ODC-BY) snippets of 1-8 consecutive sentences (20-2000 chars;
a third are single sentences or short spans), ~90 languages.   python emb/build_texts.py OUT.parquet [PER_LANG]   (CPU job)
Columns: pid, text, lang, url. English gets 4x PER_LANG. Different random seed and documents than rerank/build_passages.py.
"""
import json, os, random, re, sys, time
from multiprocessing import Pool

OUT = sys.argv[1]
PER = int(sys.argv[2]) if len(sys.argv) > 2 else 8000
FW = json.load(open(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "lid", "fw_files.json")))
LANGS2 = ("afr_Latn als_Latn amh_Ethi hye_Armn azj_Latn eus_Latn bel_Cyrl bos_Latn cym_Latn epo_Latn glg_Latn guj_Gujr hau_Latn "
          "isl_Latn gle_Latn jav_Latn kan_Knda kat_Geor kaz_Cyrl khm_Khmr kir_Cyrl lao_Laoo mal_Mlym mkd_Cyrl mlt_Latn khk_Cyrl mya_Mymr "
          "npi_Deva ory_Orya pan_Guru pbt_Arab sin_Sinh som_Latn sun_Latn tgk_Cyrl tat_Cyrl uig_Arab uzn_Latn xho_Latn zul_Latn "
          "ltz_Latn fao_Latn snd_Arab ckb_Arab kin_Latn").split()
LANGS = ("arb_Arab ben_Beng deu_Latn spa_Latn fas_Arab fin_Latn fra_Latn hin_Deva ind_Latn jpn_Jpan kor_Hang rus_Cyrl swh_Latn "
         "tel_Telu tha_Thai yor_Latn cmn_Hani bul_Cyrl ces_Latn dan_Latn ita_Latn nld_Latn nob_Latn por_Latn ron_Latn srp_Cyrl "
         "swe_Latn pol_Latn tur_Latn ukr_Cyrl vie_Latn heb_Hebr ell_Grek hun_Latn zsm_Latn fil_Latn urd_Arab tam_Taml mar_Deva "
         "slk_Latn hrv_Latn cat_Latn lit_Latn ekk_Latn lvs_Latn slv_Latn").split()
SPLIT = re.compile(r"(?<=[.!?。!?।؟።])\s+|\n+")


def collect(args):
    lang, repo, files, per = args
    import pyarrow.parquet as pq
    from huggingface_hub import HfFileSystem
    rng = random.Random(1007 + sum(map(ord, lang))); fs = HfFileSystem(); rows = []

    def retry(fn):
        for k in range(6):
            try:
                return fn()
            except Exception:
                if k == 5:
                    raise
                time.sleep(20 * (k + 1))
    try:
        for f in files:
            pf = retry(lambda: pq.ParquetFile(fs.open(f"datasets/{repo}/{f}")))
            for g in range(pf.num_row_groups - 1, -1, -1):   # last row groups first: other documents than the passage pool
                tb = retry(lambda: pf.read_row_group(g, columns=["text", "url"]))
                for text, url in zip(tb.column("text").to_pylist(), tb.column("url").to_pylist()):
                    sents = [s.strip() for s in SPLIT.split(text or "") if len(s.strip()) >= 10]
                    if len(sents) < 2:
                        continue
                    k = rng.choice([1, 1, 2, 3, 4, 6, 8]); i = rng.randrange(max(1, len(sents) - k + 1))
                    t = " ".join(sents[i:i + k])
                    if rng.random() < 0.15 and len(t) > 60:   # short span (query-like)
                        st = rng.randrange(0, len(t) - 40); t = t[st:st + rng.randint(20, 60)]
                    if 20 <= len(t) <= 2000:
                        rows.append((t, lang, url))
                    if len(rows) >= per:
                        break
                if len(rows) >= per:
                    break
            if len(rows) >= per:
                break
    except Exception as e:
        print(lang, "failed", e, flush=True)
    print(lang, len(rows), flush=True)
    return rows


if __name__ == "__main__":
    import pandas as pd
    jobs = [("eng_Latn", "HuggingFaceFW/fineweb", FW["fineweb_eng"], PER * 4)] + [(l, "HuggingFaceFW/fineweb-2", FW["fineweb2"][l], PER) for l in LANGS + [x for x in LANGS2 if x in FW["fineweb2"]]]
    with Pool(int(os.environ.get("NPROC", 14))) as p:
        res = p.map(collect, jobs)
    df = pd.DataFrame([r for rs in res for r in rs], columns=["text", "lang", "url"]).drop_duplicates("text").reset_index(drop=True)
    df.insert(0, "pid", range(len(df)))
    df.to_parquet(OUT); print("passages", len(df), df.lang.nunique(), "langs", flush=True)
    os._exit(0)