File size: 4,493 Bytes
332a987
5fb8489
332a987
 
 
 
 
 
5fb8489
332a987
 
 
 
5fb8489
 
332a987
 
 
 
5fb8489
 
 
 
 
 
332a987
5fb8489
 
 
 
 
 
332a987
 
 
 
 
5fb8489
332a987
 
 
5fb8489
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
332a987
5fb8489
 
 
 
332a987
5fb8489
 
 
332a987
 
5fb8489
 
 
 
 
 
332a987
 
5fb8489
332a987
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
927b117
 
332a987
 
927b117
332a987
927b117
332a987
 
 
927b117
 
332a987
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
# /// 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()