Text Generation
Transformers
Safetensors
Arabic
llama
arabic
reasoning
chain-of-thought
math
gsm8k
small-language-model
slm
sft
conversational
text-generation-inference
File size: 4,970 Bytes
867d0f3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""
Assemble the Arabic dataset from the translation cache and validate it.

Validation per row (a translation is only kept if it passes):
  * both segments translated and non-empty
  * the numbers in the Arabic text match the English exactly (order-insensitive multiset)
  * no degenerate repetition loop from the translator
  * not left largely untranslated (Latin-script residue)

Writes out_gsm/gsm8k_reasoning_ar.parquet plus a rejects file for inspection/retry.
"""
import collections
import json
import re
import sys
from pathlib import Path

import pyarrow as pa
import pyarrow.parquet as pq

sys.path.insert(0, ".")
from gsm_common import build_record

OUT = Path("out_gsm")
AR_DIGITS = str.maketrans("٠١٢٣٤٥٦٧٨٩", "0123456789")
NUM_RE = re.compile(r"\d+(?:\.\d+)?")
LATIN_RE = re.compile(r"[A-Za-z]")
ARABIC_RE = re.compile(r"[؀-ۿ]")


def numbers(text):
    return NUM_RE.findall(text.translate(AR_DIGITS).replace(",", ""))


def has_repetition_loop(text, n=6, times=3):
    """Detect the translator getting stuck repeating an n-word window."""
    words = text.split()
    if len(words) < n * times:
        return False
    counts = collections.Counter(
        " ".join(words[i : i + n]) for i in range(len(words) - n + 1)
    )
    return counts.most_common(1)[0][1] >= times


VERBALISABLE_MAX = 12  # Arabic writes small quantities as words ("ضعف" for 2, "الستة" for 6)


def check(en, ar, strict=True):
    """
    strict=True (reasoning chains): every numeral must survive exactly — the arithmetic depends
    on it. strict=False (questions): a small number may be verbalised, but a number the source
    never contained is a translation error (observed: '$13751' -> '13571', '4 × 44' -> '4 × 46'),
    which silently corrupts the math and is always rejected.
    """
    if not ar or not ar.strip():
        return "empty"
    en_n, ar_n = collections.Counter(numbers(en)), collections.Counter(numbers(ar))
    if ar_n - en_n:
        return "invented_number"
    missing = en_n - ar_n
    if missing:
        if strict:
            return "number_dropped"
        if any(float(v) > VERBALISABLE_MAX for v in missing):
            return "number_dropped"
    if has_repetition_loop(ar):
        return "repetition"
    if not ARABIC_RE.search(ar):
        return "not_arabic"
    latin = len(LATIN_RE.findall(ar))
    if latin > 0.25 * len(ar.replace(" ", "")):
        return "latin_residue"
    return None


def main():
    trans = {}
    with open(OUT / "translations.jsonl", encoding="utf-8") as fh:
        for line in fh:
            try:
                r = json.loads(line)
            except json.JSONDecodeError:
                continue
            trans[r["src"]] = r["tgt"]
    print(f"[*] {len(trans)} cached translations")

    rows = [json.loads(l) for l in open(OUT / "selected_rows.jsonl", encoding="utf-8")]
    print(f"[*] {len(rows)} selected rows")

    kept, rejects = [], []
    reasons = collections.Counter()
    for r in rows:
        q_ar, t_ar = trans.get(r["question"]), trans.get(r["thinking"])
        if q_ar is None or t_ar is None:
            reasons["missing"] += 1
            rejects.append({**r, "reason": "missing"})
            continue
        why = check(r["question"], q_ar, strict=False) or check(r["thinking"], t_ar, strict=True)
        if why:
            reasons[why] += 1
            rejects.append({**r, "question_ar": q_ar, "thinking_ar": t_ar, "reason": why})
            continue
        kept.append(
            {
                "text": build_record(q_ar, t_ar, r["answer"]),
                "question": q_ar,
                "thinking": t_ar,
                "answer": r["answer"],
                "question_en": r["question"],
                "thinking_en": r["thinking"],
                "source_index": r["idx"],
            }
        )

    print(f"[*] kept {len(kept)}/{len(rows)} ({len(kept)/len(rows):.2%})")
    print(f"[*] rejects: {dict(reasons)}")

    table = pa.table({k: [row[k] for row in kept] for k in kept[0]})
    pq.write_table(table, OUT / "gsm8k_reasoning_ar.parquet", compression="zstd")
    print(f"[+] wrote {OUT / 'gsm8k_reasoning_ar.parquet'} ({table.num_rows} rows)")

    with open(OUT / "rejects.jsonl", "w", encoding="utf-8") as fh:
        for r in rejects:
            fh.write(json.dumps(r, ensure_ascii=False) + "\n")
    print(f"[+] wrote {OUT / 'rejects.jsonl'} ({len(rejects)} rows)")

    stats = {
        "selected": len(rows),
        "kept": len(kept),
        "kept_pct": 100 * len(kept) / len(rows),
        "rejects": dict(reasons),
        "unique_translations": len(trans),
    }
    (OUT / "build_stats.json").write_text(json.dumps(stats, ensure_ascii=False, indent=2), encoding="utf-8")
    for row in kept[:3]:
        print("-" * 70)
        print("EN:", row["question_en"][:120])
        print("AR:", row["question"][:120])
        print("AR think:", row["thinking"][:160])


if __name__ == "__main__":
    main()