| |
| |
| |
| """DPO on offline compile-ok vs compile-fail preferences (compiler-as-judge).""" |
|
|
| import os |
|
|
| from datasets import load_dataset |
| from peft import LoraConfig, PeftModel |
| from trl import DPOConfig, DPOTrainer |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
| DATASET_ID = os.environ.get("DATASET_ID", "gonzalolinares/cpp-compiler-prefs") |
| SFT_ADAPTER = os.environ.get("BASE_MODEL", "gonzalolinares/qwen25-1.5b-cpp-sft") |
| BASE_MODEL = os.environ.get("FALLBACK_MODEL", "Qwen/Qwen2.5-1.5B-Instruct") |
| HUB_MODEL_ID = os.environ.get("HUB_MODEL_ID", "gonzalolinares/qwen25-1.5b-cpp-dpo") |
| OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "qwen25-1.5b-cpp-dpo") |
|
|
|
|
| def load_policy(): |
| """Load base NL model, merge SFT LoRA if present.""" |
| tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL) |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
| model = AutoModelForCausalLM.from_pretrained(BASE_MODEL, torch_dtype="auto") |
| try: |
| model = PeftModel.from_pretrained(model, SFT_ADAPTER) |
| model = model.merge_and_unload() |
| print(f"Merged SFT adapter from {SFT_ADAPTER}") |
| except Exception as e: |
| print(f"No SFT adapter merge ({e}); training from base {BASE_MODEL}") |
| return model, tokenizer |
|
|
|
|
| def main() -> None: |
| ds = load_dataset(DATASET_ID, split="train") |
|
|
| def to_dpo(example): |
| prompt = example.get("prompt") |
| chosen = example.get("chosen") |
| rejected = example.get("rejected") |
|
|
| def last_assistant(messages): |
| if isinstance(messages, list) and messages: |
| last = messages[-1] |
| if isinstance(last, dict): |
| return last.get("content", str(last)) |
| return str(messages) |
|
|
| def user_text(messages): |
| if isinstance(messages, list): |
| parts = [] |
| for m in messages: |
| if isinstance(m, dict) and m.get("role") in {"system", "user"}: |
| parts.append(m.get("content", "")) |
| return "\n\n".join(parts) |
| return str(messages) |
|
|
| |
| if isinstance(prompt, list) and prompt and isinstance(prompt[0], dict): |
| ch = chosen if isinstance(chosen, list) else [{"role": "assistant", "content": str(chosen)}] |
| rj = rejected if isinstance(rejected, list) else [{"role": "assistant", "content": str(rejected)}] |
| return {"prompt": prompt, "chosen": ch, "rejected": rj} |
|
|
| return { |
| "prompt": user_text(prompt), |
| "chosen": last_assistant(chosen), |
| "rejected": last_assistant(rejected), |
| } |
|
|
| ds = ds.map(to_dpo, remove_columns=[c for c in ds.column_names if c not in {"prompt", "chosen", "rejected"}]) |
| |
| keep = {"prompt", "chosen", "rejected"} |
| drop = [c for c in ds.column_names if c not in keep] |
| if drop: |
| ds = ds.remove_columns(drop) |
| split = ds.train_test_split(test_size=0.1, seed=42) |
|
|
| model, tokenizer = load_policy() |
|
|
| trainer = DPOTrainer( |
| model=model, |
| processing_class=tokenizer, |
| train_dataset=split["train"], |
| eval_dataset=split["test"], |
| peft_config=LoraConfig( |
| r=16, |
| lora_alpha=32, |
| lora_dropout=0.05, |
| bias="none", |
| task_type="CAUSAL_LM", |
| target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], |
| ), |
| args=DPOConfig( |
| output_dir=OUTPUT_DIR, |
| num_train_epochs=2, |
| per_device_train_batch_size=1, |
| per_device_eval_batch_size=1, |
| gradient_accumulation_steps=8, |
| learning_rate=5e-5, |
| logging_steps=5, |
| eval_strategy="steps", |
| eval_steps=20, |
| save_strategy="epoch", |
| save_total_limit=1, |
| max_length=1024, |
| bf16=True, |
| push_to_hub=False, |
| hub_model_id=HUB_MODEL_ID, |
| report_to="none", |
| ), |
| ) |
| trainer.train() |
| trainer.model.push_to_hub(HUB_MODEL_ID, private=False) |
| tokenizer.push_to_hub(HUB_MODEL_ID, private=False) |
| print(f"Pushed to {HUB_MODEL_ID}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|