Text Generation
Transformers
Safetensors
Arabic
llama
arabic
reasoning
chain-of-thought
math
gsm8k
small-language-model
slm
sft
conversational
text-generation-inference
oddadmix's picture
training code: data generation, SFT, eval, GRPO
867d0f3 verified
|
Raw
History Blame Contribute Delete
10.3 kB

Nawah-Math-Reasoning β€” training code

Everything that produced the model in this repo: two dataset pipelines, the SFT trainer, the eval harness, a pass@k diagnostic, and a GRPO implementation. Run top to bottom and you get the release; run any single stage and it resumes from what is already on disk.

Paths below assume this directory is the working directory and $P is a Python with torch, transformers>=5.15, pyarrow and huggingface_hub.

0. Two environments, and why

env for pins
training / eval / Hub train_reasoning.py, eval_reasoning.py, everything else torch 2.9 + cu126, transformers 5.15
generation translate_gsm.py, synth_generate.py vLLM 0.8.5.post1, torch 2.6 + cu124, transformers 4.51.3

They are not interchangeable. The model configs are transformers-v5 format (rope_parameters, dtype, tokenizer_class: TokenizersBackend) and transformers 4.x cannot read them β€” but vLLM 0.8.5 breaks on transformers 5.x (TokenizersBackend has no attribute all_special_tokens_extended), so the generation env must stay at 4.51.3. Do not pip install -U in it.

vLLM 0.8.5 is itself a pin: every release from 0.20.2 up requires torch 2.11, which is a CUDA 13 build needing driver β‰₯ 580. On a CUDA 12.4 host that is a hard stop, and cu12 wheels do not rescue it (ImportError: libcudart.so.13). 0.8.5 is the last cu124 release.

In transformers v5 TrainingArguments it is eval_strategy (not evaluation_strategy) and warmup_steps (there is no warmup_ratio).

1. gsm8k-reasoning-ar β€” 142,969 machine-translated rows

$V translate_gsm.py 150000     # vLLM env; ~95 min on 1x A6000; resumable
$P build_dataset.py            # -> out_gsm/*.parquet + rejects.jsonl
$P push_dataset.py --repo <user>/gsm8k-reasoning-ar

Ajhesh7/gsm8k-reasoning-SFT-datas translated with ByteDance-Seed/Seed-X-PPO-7B. The trailing language tag in the prompt is mandatory: Translate the following English sentence into Arabic:\n{text} <ar>.

Only 150k of the source's 600,000 rows are translated, and the sample is stratified by question pattern with a floor of 13 rows/pattern: the corpus expands from just 2,814 question patterns, so translating all 600k is ~99% redundant.

The reason build_dataset.py has a validator at all: Seed-X silently corrupts arithmetic in ~0.6% of segments β€” $13751 β†’ 13571, 4 Γ— 44 = 176 β†’ 4 Γ— 46 = 176. These read as fluent Arabic and pass any fluency check. So every row is numeral-audited, asymmetrically: reasoning chains strictly (every numeral must survive exactly), questions leniently (a value ≀ 12 may be verbalised β€” 6 friends β†’ Ψ£Ψ΅Ψ―Ω‚Ψ§Ψ¦Ω‡Ψ§ Ψ§Ω„Ψ³ΨͺΨ© β€” but any number the English never contained is rejected). Naive exact-multiset matching rejects ~6% of correct translations; do not "simplify" it back to that.

translate_gsm.py appends every chunk to out_gsm/translations.jsonl and skips cached segments on restart, so an interruption costs at most one 20k chunk. Do not pass enable_prefix_caching=True β€” it hangs at startup on this vLLM/V1 engine.

2. arabic-math-reasoning-synth β€” 120,462 verified rows

# general pool: 100,323 rows, ~22 h
GEN_MODEL=<path>/gemma-3-12b-it BACKEND=vllm OUT_DIR=out_synth TARGET=100000 \
BATCH=256 MAX_NEW=1200 MAX_LEN=3072 GPU_UTIL=0.90 \
nohup $V -u synth_generate.py > synth_run.log 2>&1 &

# relational pool: 20,139 rows, ~3.5 h. Disjoint task-id range, on purpose.
GEN_MODEL=<path>/gemma-3-12b-it BACKEND=vllm OUT_DIR=out_synth_rel POOL=relational \
START_TASK=1000000 TARGET=20000 BATCH=256 MAX_NEW=1200 MAX_LEN=3072 GPU_UTIL=0.90 SEED=1234 \
nohup $V -u synth_generate.py > synth_rel.log 2>&1 &

./finish_merge_push.sh          # merge primary cache + any node shards -> out_merged_v6/
$P split_synth_v6.py            # -> data_synth_v6_sft/{train,eval,eval_rel}.jsonl
$P push_synth_dataset.py --repo <user>/arabic-math-reasoning-synth

Resumability is the design, not a feature. A task is one generation call asking for 4 problems, and task N's prompt is a pure function of N (synth_common.build_task) β€” nothing about the plan is persisted, so a restart redraws identical prompts. Finished tasks are appended to generations.jsonl and fsynced; a torn final line is dropped with a warning. The stop condition is accepted rows, not tasks.

Raw completions are stored, never just the parsed rows. Every validator change can be re-scored over the whole cache with no GPU β€” which is how the accept rate went 48.4% β†’ 64.9% without regenerating anything.

The arithmetic audit (synth_common.validate) re-evaluates every a op b = c in the reasoning and one wrong equation rejects the row. Two things it must keep doing, both of which were bugs that rejected correct rows:

  1. Equation chains. Models write 75 + 15 Γ— 5 = 75 + 75 = 150. Reading only to the first = compares 150 against 75 and rejects a correct chain. Split the whole chain, require every segment to agree.
  2. Rounding. 3200 / 60 = 53.33 is arithmetic as people write it. _close() forgives rounding at the precision the model displayed (and floor/ceil for integers), so a genuinely wrong number still fails.

Also rejected: noop_step (63 + 0 = 63, padding to hit a step count), meta_commentary, Latin residue, missing conclusion marker, and answers that disagree with the last computed value.

The relational pool is a separate pool (RELATIONAL_OPS), deliberately not appended to OPERATIONS. Appending would change rng.choice() for every task id and silently break reproducibility of the first 100,323 rows. build_task(task_id, seed, pool) takes pool="default" or "relational"; keep it that way.

Check unique_templates in build_stats.json rather than assuming the variation grid worked.

3. SFT

$P prepare_data.py              # Arabic_Reasoning_Dataset -> data/{train,eval}.jsonl
$P prepare_gsm_sft.py           # -> data_gsm_sft/{train,eval}.jsonl
$P prepare_v6_sft.py            # three-way mix -> data_v6_sft/{train,eval}.jsonl

BASE_MODEL=<base> OUTPUT_DIR=./Nawah-Math-Reasoning \
TRAIN_FILE=data_v6_sft/train.jsonl EVAL_FILE=data_v6_sft/eval.jsonl \
MAX_LENGTH=768 EPOCHS=5 BATCH_SIZE=64 GRAD_ACCUM=1 WARMUP_STEPS=200 EVAL_STEPS=1000 LOAD_BEST=0 \
$P -u train_reasoning.py 2>&1 | tee train.log      # 21,535 steps, ~85 min on 1x A6000

LOAD_BEST=0 is load-bearing. Eval loss selects the worse checkpoint on this ladder, and that was measured rather than assumed: on a corpus with no repeated rows, the minimum-loss checkpoint scored 30.9% where the final scored 35.6%. It held for four consecutive runs. Ship the final checkpoint.

Answer styles are not normalised. GSM8K rows end in a bare numeral, the other two corpora in an Ψ₯Ψ°Ω†ΨŒ … sentence. Rewriting them into one style deletes what the mix adds, so a source tag rides on every row and eval scores each half on its own terms.

The eval splits are pinned across versions β€” the same 400 Arabic_Reasoning rows and the same 600 GSM8K rows since the first model, and split_synth_v6.py copies the previous synth eval rows through verbatim rather than re-shuffling. Re-drawing would have moved 1,955 of 2,000 held-out items into train.

4. Eval

EVAL_FILE=data_v6_sft/eval.jsonl        $P eval_reasoning.py ./Nawah-Math-Reasoning 1800
EVAL_FILE=data_synth_sft/eval.jsonl     $P eval_reasoning.py ./Nawah-Math-Reasoning 1000

Greedy. Reports well-formedness, number agreement, exact match and reasoning length, broken down per source when the rows carry the tag. EVAL_FILE comes from the env β€” pointing it at the wrong split silently scores against the wrong data.

Score on number agreement, not exact match: with two answer styles in the mix, exact match measures style compliance, not arithmetic.

5. RL β€” run the diagnostic first

$P -u passk_diag.py ./Nawah-Math-Reasoning     # 240 problems x k=8 @ T=1.0, ~25 min

RLVR reweights samples the model already produces. A group where every sample is wrong scores all-zero, the advantage is zero, and there is no gradient β€” so exploitable headroom is bounded by pass@k βˆ’ pass@1. On the previous version that was +27.1 points, with 40.4% of problems never solved in 8 tries. Run this before spending a GPU-hour.

That diagnostic is also what produced this release: the dead 40% turned out to be a data gap (relational comparisons were 1.34% of the corpus), so the fix was Β§2's relational pool through SFT, not RL. Data first, then RL β€” running RL first leaves the dead tail untouched and caps the gain.

MODEL=./Nawah-Math-Reasoning OUTPUT_DIR=./Nawah-Math-Reasoning-grpo \
GROUP=8 PROMPTS_PER=8 STEPS=500 LR=1e-6 BETA=0.02 \
$P -u grpo_train.py 2>&1 | tee grpo.log

grpo_train.py is hand-rolled β€” TRL is not installable against these pins. Group-normalised advantage with no value network, binary final-answer reward (not partial credit β€” the failure being fixed is fluent reasoning landing on a wrong number), KL to a frozen reference via Schulman's k3 estimator, and no importance ratio or PPO clipping because sampling is on-policy with one step per batch. Zero-spread groups are skipped and counted; the dead-group percentage in the log is the number to watch.

6. Release

$P build_release_code.py                    # stage this directory
$P push_release.py --dry-run                # render the card only
$P push_release.py

Every number in the model card is read from an eval JSON on disk. Nothing is typed by hand, so the card cannot drift from the measurements.

Demo

space/ is the Gradio demo, deployable as-is to a Space. Two things that break it if "cleaned up": the streamer must use skip_special_tokens=False (<think>/</think> are real special tokens and stripping them destroys the reasoning/answer split), and render_prompt() must stay byte-identical to the trainer's rendering β€” a 52M model is very sensitive to format drift.

On ZeroGPU, import spaces must come before torch, and model.to("cuda") at startup must stay unguarded: no GPU is attached at startup and the call is replayed inside the forked GPU process.