Text Generation
Transformers
Safetensors
Arabic
llama
arabic
reasoning
chain-of-thought
math
gsm8k
small-language-model
slm
sft
conversational
text-generation-inference
File size: 7,184 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
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
"""
Reasoning SFT for oddadmix/50M-2048-Emhotob on Arabic_Reasoning_Dataset.

ChatML format with the derivation wrapped in <think>...</think>. Loss is computed on the
assistant turn only — the user prompt is masked out, same as the earlier Emhotob SFT runs.
No TRL; plain HF Trainer.
"""
import json
import os
from dataclasses import dataclass
from pathlib import Path

import torch
from torch.utils.data import Dataset
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    Trainer,
    TrainingArguments,
)

os.environ.setdefault("CUDA_VISIBLE_DEVICES", "0")

# Defaults reproduce v1 (Arabic_Reasoning_Dataset); every value can be overridden by env var
# so the same recipe can be pointed at a different corpus.
def _env(name, default, cast=str):
    return cast(os.environ.get(name, default))

BASE_MODEL   = _env("BASE_MODEL", "/notebooks/50M/50M-2048-Emhotob")
OUTPUT_DIR   = _env("OUTPUT_DIR", "./Nawah-Reasoning-v1")
TRAIN_FILE   = _env("TRAIN_FILE", "data/train.jsonl")
EVAL_FILE    = _env("EVAL_FILE", "data/eval.jsonl")

MAX_LENGTH   = _env("MAX_LENGTH", 768, int)   # v1: p100 of that corpus is 708 tokens
IGNORE_INDEX = -100

LEARNING_RATE = _env("LEARNING_RATE", 3e-4, float)   # same as the Emhotob translation SFT ladder
EPOCHS        = _env("EPOCHS", 8, int)   # v1 is tiny (~840k tok/epoch); best checkpoint wins
BATCH_SIZE    = _env("BATCH_SIZE", 16, int)
GRAD_ACCUM    = _env("GRAD_ACCUM", 2, int)
WARMUP_STEPS  = _env("WARMUP_STEPS", 100, int)
EVAL_STEPS    = _env("EVAL_STEPS", 100, int)
# On the v3 mix, eval loss is a bad model selector: the repeated Arabic_Reasoning rows start
# memorising around epoch 1.4 and drag the loss up while generation quality on *both* halves is
# still improving. Set LOAD_BEST=0 there and keep the final checkpoint.
LOAD_BEST     = _env("LOAD_BEST", 1, int) == 1
# Point at a checkpoint dir to continue an interrupted run (optimizer/scheduler/RNG/step are
# restored from it). Empty = fresh run, so v1-v5 still reproduce exactly.
RESUME        = _env("RESUME", "") or None
WEIGHT_DECAY  = 0.0
MAX_GRAD_NORM = 1.0
SEED          = 42

SPECIAL_TOKENS = ["<|im_start|>", "<|im_end|>", "<think>", "</think>"]

CHAT_TEMPLATE = (
    "{% for message in messages %}"
    "{{ '<|im_start|>' + message['role'] + '\n' + message['content'] + '<|im_end|>' + '\n' }}"
    "{% endfor %}"
    "{% if add_generation_prompt %}{{ '<|im_start|>assistant\n' }}{% endif %}"
)

PROMPT_TMPL   = "<|im_start|>user\n{instruction}<|im_end|>\n<|im_start|>assistant\n"
RESPONSE_TMPL = "<think>\n{reasoning}\n</think>\n{answer}<|im_end|>"


def load_jsonl(path):
    with open(path, encoding="utf-8") as fh:
        return [json.loads(line) for line in fh]


class ReasoningDataset(Dataset):
    """Prompt tokens are masked so loss falls only on <think>…</think> + answer."""

    def __init__(self, rows, tokenizer, max_length):
        self.rows = rows
        self.tok = tokenizer
        self.max_length = max_length

    def __len__(self):
        return len(self.rows)

    def __getitem__(self, idx):
        row = self.rows[idx]
        prompt = PROMPT_TMPL.format(instruction=row["instruction"])
        response = RESPONSE_TMPL.format(reasoning=row["reasoning"], answer=row["answer"])

        prompt_ids = [self.tok.bos_token_id] + self.tok.encode(prompt, add_special_tokens=False)
        response_ids = self.tok.encode(response, add_special_tokens=False)

        input_ids = (prompt_ids + response_ids)[: self.max_length]
        prompt_len = min(len(prompt_ids), len(input_ids))
        labels = [IGNORE_INDEX] * prompt_len + input_ids[prompt_len:]

        return {
            "input_ids": torch.tensor(input_ids, dtype=torch.long),
            "labels": torch.tensor(labels, dtype=torch.long),
        }


@dataclass
class PaddingCollator:
    pad_token_id: int

    def __call__(self, features):
        longest = max(len(f["input_ids"]) for f in features)
        input_ids, labels, attention = [], [], []
        for f in features:
            pad = longest - len(f["input_ids"])
            input_ids.append(torch.cat([f["input_ids"], torch.full((pad,), self.pad_token_id, dtype=torch.long)]))
            labels.append(torch.cat([f["labels"], torch.full((pad,), IGNORE_INDEX, dtype=torch.long)]))
            attention.append(torch.cat([torch.ones(len(f["input_ids"]), dtype=torch.long), torch.zeros(pad, dtype=torch.long)]))
        return {
            "input_ids": torch.stack(input_ids),
            "labels": torch.stack(labels),
            "attention_mask": torch.stack(attention),
        }


def main():
    print("[*] loading tokenizer + base model")
    tok = AutoTokenizer.from_pretrained(BASE_MODEL)
    added = tok.add_special_tokens({"additional_special_tokens": SPECIAL_TOKENS})
    tok.chat_template = CHAT_TEMPLATE
    print(f"    added {added} special tokens -> vocab {len(tok)}")

    model = AutoModelForCausalLM.from_pretrained(BASE_MODEL, dtype=torch.float32)
    model.resize_token_embeddings(len(tok))
    model.config.use_cache = False
    print(f"    params: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M")

    train_rows, eval_rows = load_jsonl(TRAIN_FILE), load_jsonl(EVAL_FILE)
    print(f"[*] train {len(train_rows)} / eval {len(eval_rows)}")

    args = TrainingArguments(
        output_dir=OUTPUT_DIR,
        num_train_epochs=EPOCHS,
        per_device_train_batch_size=BATCH_SIZE,
        per_device_eval_batch_size=BATCH_SIZE,
        gradient_accumulation_steps=GRAD_ACCUM,
        learning_rate=LEARNING_RATE,
        lr_scheduler_type="cosine",
        warmup_steps=WARMUP_STEPS,
        weight_decay=WEIGHT_DECAY,
        max_grad_norm=MAX_GRAD_NORM,
        bf16=True,
        logging_steps=25,
        eval_strategy="steps",
        eval_steps=EVAL_STEPS,
        save_strategy="steps",
        save_steps=EVAL_STEPS,
        save_total_limit=2,
        load_best_model_at_end=LOAD_BEST,
        metric_for_best_model="eval_loss",
        greater_is_better=False,
        report_to=[],
        seed=SEED,
        dataloader_num_workers=2,
        remove_unused_columns=False,
    )

    trainer = Trainer(
        model=model,
        args=args,
        train_dataset=ReasoningDataset(train_rows, tok, MAX_LENGTH),
        eval_dataset=ReasoningDataset(eval_rows, tok, MAX_LENGTH),
        data_collator=PaddingCollator(pad_token_id=tok.pad_token_id),
    )

    if RESUME:
        print(f"[*] resuming from {RESUME}")
    trainer.train(resume_from_checkpoint=RESUME)

    print("[*] saving best checkpoint")
    im_end_id = tok.convert_tokens_to_ids("<|im_end|>")
    model.config.use_cache = True
    model.generation_config.eos_token_id = [tok.eos_token_id, im_end_id]
    model.generation_config.pad_token_id = tok.pad_token_id
    trainer.save_model(OUTPUT_DIR)
    tok.save_pretrained(OUTPUT_DIR)

    metrics = trainer.evaluate()
    print("[*] final eval:", metrics)
    Path(OUTPUT_DIR, "train_metrics.json").write_text(
        json.dumps({"final_eval": metrics, "log_history": trainer.state.log_history}, ensure_ascii=False, indent=2),
        encoding="utf-8",
    )
    print(f"[+] done -> {OUTPUT_DIR}")


if __name__ == "__main__":
    main()