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,834 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
"""
Turn out_synth/generations.jsonl into the finished synthetic corpus.

Re-parses and re-validates every cached generation with synth_common (the generator's inline
accounting is only a progress estimate; this is the authoritative pass), drops duplicates, and
writes a parquet plus the SFT splits train_reasoning.py reads directly.

Duplicates are keyed on the question with its numbers masked out, so "3 apples at 5 riyal" and
"7 apples at 9 riyal" collapse to one template — this is the exact failure the translated GSM8K
set has (142,969 rows from 2,814 templates), and the whole point of the variation grid is to not
repeat it. The template count is reported so it can be checked rather than assumed.

Writes:
  out_synth/arabic_math_reasoning_synth.parquet   the corpus
  out_synth/rejects.jsonl                         every rejected item with its reason
  out_synth/build_stats.json                      counts, reject histogram, axis coverage
  data_synth_sft/{train,eval}.jsonl               ready for train_reasoning.py
"""
import collections
import json
import os
import random
from pathlib import Path

import pyarrow as pa
import pyarrow.parquet as pq

import synth_common as sc

OUT_DIR = Path(os.environ.get("OUT_DIR", "out_synth"))
CACHE = OUT_DIR / "generations.jsonl"
SFT_DIR = Path(os.environ.get("SFT_DIR", "data_synth_sft"))
EVAL_N = int(os.environ.get("EVAL_N", 2000))
LIMIT = int(os.environ.get("LIMIT", 0))          # 0 = keep everything that validates
SEED = 42


def main():
    stats = collections.Counter()
    rejects_by_reason = collections.Counter()
    axes_seen = collections.defaultdict(collections.Counter)
    rows, rejects = [], []
    seen = {}

    with open(CACHE, encoding="utf-8") as fh:
        for line in fh:
            try:
                rec = json.loads(line)
            except json.JSONDecodeError:
                stats["truncated_lines"] += 1        # a crash mid-write; the rest is still good
                continue
            stats["tasks"] += 1
            items = sc.parse_items(rec["raw"])
            stats["parsed_items"] += len(items)
            if not items:
                stats["tasks_with_no_parsable_item"] += 1
            for item in items:
                ok, reason = sc.validate(item)
                if not ok:
                    rejects_by_reason[reason] += 1
                    rejects.append({**item, "reason": reason, "task_id": rec["task_id"]})
                    continue
                key = sc.dedup_key(item["instruction"])
                if key in seen:
                    stats["dropped_duplicate_template"] += 1
                    continue
                seen[key] = True
                for axis, value in rec["axes"].items():
                    axes_seen[axis][value] += 1
                rows.append({**item, **{f"axis_{k}": v for k, v in rec["axes"].items()},
                             "task_id": rec["task_id"],
                             # rows cached before multi-node generation carry no model field
                             "gen_model": rec.get("model", "gemma-3-12b-it")})

    for r in rows:
        stats[f"model_{r['gen_model']}"] += 1

    stats["kept"] = len(rows)
    stats["rejected"] = sum(rejects_by_reason.values())

    random.Random(SEED).shuffle(rows)
    if LIMIT:
        rows = rows[:LIMIT]

    OUT_DIR.mkdir(exist_ok=True)
    pq.write_table(pa.Table.from_pylist(rows), OUT_DIR / "arabic_math_reasoning_synth.parquet")
    with open(OUT_DIR / "rejects.jsonl", "w", encoding="utf-8") as fh:
        for r in rejects:
            fh.write(json.dumps(r, ensure_ascii=False) + "\n")

    SFT_DIR.mkdir(exist_ok=True)
    eval_rows, train_rows = rows[:EVAL_N], rows[EVAL_N:]
    for name, split in (("train", train_rows), ("eval", eval_rows)):
        with open(SFT_DIR / f"{name}.jsonl", "w", encoding="utf-8") as fh:
            for r in split:
                fh.write(json.dumps({"instruction": r["instruction"], "reasoning": r["reasoning"],
                                     "answer": r["answer"], "source": "synth_math_ar"},
                                    ensure_ascii=False) + "\n")
        print(f"[+] {name}: {len(split):,} -> {SFT_DIR / f'{name}.jsonl'}")

    summary = {
        "stats": dict(stats),
        "reject_reasons": dict(rejects_by_reason.most_common()),
        "accept_rate": stats["kept"] / max(stats["parsed_items"], 1),
        "unique_templates": len(seen),
        "axis_coverage": {k: len(v) for k, v in axes_seen.items()},
    }
    (OUT_DIR / "build_stats.json").write_text(json.dumps(summary, ensure_ascii=False, indent=2),
                                              encoding="utf-8")
    print(json.dumps(summary, ensure_ascii=False, indent=2))
    print(f"[+] {OUT_DIR / 'arabic_math_reasoning_synth.parquet'}")


if __name__ == "__main__":
    main()