""" Reasoning SFT for oddadmix/50M-2048-Emhotob on Arabic_Reasoning_Dataset. ChatML format with the derivation wrapped in .... 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|>", "", ""] 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 = "\n{reasoning}\n\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 + 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()