cpp-train-scripts / train_dpo.py
gonzalolinares's picture
fix DPOConfig: drop max_prompt_length
2ac14b9 verified
Raw
History Blame Contribute Delete
4.49 kB
# /// script
# dependencies = ["trl>=0.12.0", "peft>=0.7.0", "datasets", "transformers", "accelerate", "torch"]
# ///
"""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)
# TRL DPO conversational: prompt=list[messages], chosen/rejected=list with assistant
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 only needed columns after map - re-add by selecting
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()