gonzalolinares commited on
Commit
5fb8489
·
verified ·
1 Parent(s): 927b117

DPO: merge SFT LoRA then train

Browse files
Files changed (1) hide show
  1. train_dpo.py +49 -41
train_dpo.py CHANGED
@@ -1,74 +1,82 @@
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,
 
1
  # /// script
2
+ # dependencies = ["trl>=0.12.0", "peft>=0.7.0", "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, PeftModel
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
+ SFT_ADAPTER = os.environ.get("BASE_MODEL", "gonzalolinares/qwen25-1.5b-cpp-sft")
15
+ BASE_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 load_policy():
21
+ """Load base NL model, merge SFT LoRA if present."""
22
+ tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
23
+ if tokenizer.pad_token is None:
24
+ tokenizer.pad_token = tokenizer.eos_token
25
+ model = AutoModelForCausalLM.from_pretrained(BASE_MODEL, torch_dtype="auto")
26
  try:
27
+ model = PeftModel.from_pretrained(model, SFT_ADAPTER)
28
+ model = model.merge_and_unload()
29
+ print(f"Merged SFT adapter from {SFT_ADAPTER}")
30
+ except Exception as e:
31
+ print(f"No SFT adapter merge ({e}); training from base {BASE_MODEL}")
32
+ return model, tokenizer
33
 
34
 
35
  def main() -> None:
 
36
  ds = load_dataset(DATASET_ID, split="train")
 
 
 
 
 
 
 
 
 
 
 
 
 
37
 
38
+ def to_dpo(example):
39
  prompt = example.get("prompt")
40
  chosen = example.get("chosen")
41
  rejected = example.get("rejected")
42
+
43
+ def last_assistant(messages):
44
+ if isinstance(messages, list) and messages:
45
+ last = messages[-1]
46
+ if isinstance(last, dict):
47
+ return last.get("content", str(last))
48
+ return str(messages)
49
+
50
+ def user_text(messages):
51
+ if isinstance(messages, list):
52
+ parts = []
53
+ for m in messages:
54
+ if isinstance(m, dict) and m.get("role") in {"system", "user"}:
55
+ parts.append(m.get("content", ""))
56
+ return "\n\n".join(parts)
57
+ return str(messages)
58
+
59
+ # TRL DPO conversational: prompt=list[messages], chosen/rejected=list with assistant
60
  if isinstance(prompt, list) and prompt and isinstance(prompt[0], dict):
61
+ ch = chosen if isinstance(chosen, list) else [{"role": "assistant", "content": str(chosen)}]
62
+ rj = rejected if isinstance(rejected, list) else [{"role": "assistant", "content": str(rejected)}]
63
+ return {"prompt": prompt, "chosen": ch, "rejected": rj}
64
+
 
 
 
 
65
  return {
66
+ "prompt": user_text(prompt),
67
+ "chosen": last_assistant(chosen),
68
+ "rejected": last_assistant(rejected),
69
  }
70
 
71
+ ds = ds.map(to_dpo, remove_columns=[c for c in ds.column_names if c not in {"prompt", "chosen", "rejected"}])
72
+ # keep only needed columns after map - re-add by selecting
73
+ keep = {"prompt", "chosen", "rejected"}
74
+ drop = [c for c in ds.column_names if c not in keep]
75
+ if drop:
76
+ ds = ds.remove_columns(drop)
77
  split = ds.train_test_split(test_size=0.1, seed=42)
78
 
79
+ model, tokenizer = load_policy()
 
 
 
80
 
81
  trainer = DPOTrainer(
82
  model=model,