gonzalolinares commited on
Commit
332a987
·
verified ·
1 Parent(s): 17a6842

DPO script

Browse files
Files changed (1) hide show
  1. train_dpo.py +112 -0
train_dpo.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # dependencies = ["trl>=0.12.0", "peft>=0.7.0", "trackio", "datasets", "transformers", "accelerate", "torch"]
3
+ # ///
4
+ """DPO on offline compile-ok vs compile-fail preferences (compiler-as-judge)."""
5
+
6
+ import os
7
+
8
+ from datasets import load_dataset
9
+ from peft import LoraConfig
10
+ from trl import DPOConfig, DPOTrainer
11
+ from transformers import AutoModelForCausalLM, AutoTokenizer
12
+
13
+ DATASET_ID = os.environ.get("DATASET_ID", "gonzalolinares/cpp-compiler-prefs")
14
+ BASE_MODEL = os.environ.get("BASE_MODEL", "gonzalolinares/qwen25-1.5b-cpp-sft")
15
+ FALLBACK_MODEL = os.environ.get("FALLBACK_MODEL", "Qwen/Qwen2.5-1.5B-Instruct")
16
+ HUB_MODEL_ID = os.environ.get("HUB_MODEL_ID", "gonzalolinares/qwen25-1.5b-cpp-dpo")
17
+ OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "qwen25-1.5b-cpp-dpo")
18
+
19
+
20
+ def resolve_model(model_id: str) -> str:
21
+ try:
22
+ AutoTokenizer.from_pretrained(model_id)
23
+ return model_id
24
+ except Exception:
25
+ print(f"Could not load {model_id}, falling back to {FALLBACK_MODEL}")
26
+ return FALLBACK_MODEL
27
+
28
+
29
+ def main() -> None:
30
+ model_id = resolve_model(BASE_MODEL)
31
+ ds = load_dataset(DATASET_ID, split="train")
32
+ # Normalize to prompt/chosen/rejected text if needed
33
+ def to_dpo(example):
34
+ def as_text(x):
35
+ if isinstance(x, list):
36
+ # list of chat messages -> last assistant or join
37
+ parts = []
38
+ for m in x:
39
+ if isinstance(m, dict):
40
+ parts.append(f"{m.get('role', '')}: {m.get('content', '')}")
41
+ else:
42
+ parts.append(str(m))
43
+ return "\n".join(parts)
44
+ return str(x)
45
+
46
+ prompt = example.get("prompt")
47
+ chosen = example.get("chosen")
48
+ rejected = example.get("rejected")
49
+ # Prefer conversational: if prompt is message list without assistant
50
+ if isinstance(prompt, list) and prompt and isinstance(prompt[0], dict):
51
+ # TRL DPO can take conversational format
52
+ return {
53
+ "prompt": prompt,
54
+ "chosen": chosen if isinstance(chosen, list) else [{"role": "assistant", "content": str(chosen)}],
55
+ "rejected": rejected
56
+ if isinstance(rejected, list)
57
+ else [{"role": "assistant", "content": str(rejected)}],
58
+ }
59
+ return {
60
+ "prompt": as_text(prompt),
61
+ "chosen": as_text(chosen),
62
+ "rejected": as_text(rejected),
63
+ }
64
+
65
+ ds = ds.map(to_dpo)
66
+ split = ds.train_test_split(test_size=0.1, seed=42)
67
+
68
+ model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype="auto")
69
+ tokenizer = AutoTokenizer.from_pretrained(model_id)
70
+ if tokenizer.pad_token is None:
71
+ tokenizer.pad_token = tokenizer.eos_token
72
+
73
+ trainer = DPOTrainer(
74
+ model=model,
75
+ processing_class=tokenizer,
76
+ train_dataset=split["train"],
77
+ eval_dataset=split["test"],
78
+ peft_config=LoraConfig(
79
+ r=16,
80
+ lora_alpha=32,
81
+ lora_dropout=0.05,
82
+ bias="none",
83
+ task_type="CAUSAL_LM",
84
+ target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
85
+ ),
86
+ args=DPOConfig(
87
+ output_dir=OUTPUT_DIR,
88
+ num_train_epochs=2,
89
+ per_device_train_batch_size=1,
90
+ per_device_eval_batch_size=1,
91
+ gradient_accumulation_steps=8,
92
+ learning_rate=5e-5,
93
+ logging_steps=5,
94
+ eval_strategy="steps",
95
+ eval_steps=20,
96
+ max_length=1024,
97
+ max_prompt_length=512,
98
+ bf16=True,
99
+ push_to_hub=True,
100
+ hub_model_id=HUB_MODEL_ID,
101
+ report_to="trackio",
102
+ project="cpp-compiler-rl",
103
+ run_name="dpo-compiler-prefs",
104
+ ),
105
+ )
106
+ trainer.train()
107
+ trainer.push_to_hub()
108
+ print(f"Pushed to {HUB_MODEL_ID}")
109
+
110
+
111
+ if __name__ == "__main__":
112
+ main()