File size: 6,040 Bytes
625c0f0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6831b86
625c0f0
 
 
 
 
 
 
 
 
 
 
 
 
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
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
# /// script
# dependencies = ["trl>=0.12.0", "peft>=0.7.0", "datasets", "transformers", "accelerate", "torch"]
# ///
"""GRPO with g++ compiler reward (online RL). For Hugging Face Jobs (uv)."""

from __future__ import annotations

import os
import re
import shutil
import subprocess
import tempfile
from pathlib import Path

from datasets import load_dataset
from peft import LoraConfig, PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import GRPOConfig, GRPOTrainer

DATASET_ID = os.environ.get("DATASET_ID", "gonzalolinares/cpp-compiler-grpo")
SFT_ADAPTER = os.environ.get("BASE_MODEL", "gonzalolinares/qwen25-1.5b-cpp-sft")
DPO_ADAPTER = os.environ.get("DPO_MODEL", "gonzalolinares/qwen25-1.5b-cpp-dpo")
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-grpo")
OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "qwen25-1.5b-cpp-grpo")

CODE_FENCE_RE = re.compile(r"```(?:cpp|c\+\+)?\s*([\s\S]*?)```", re.IGNORECASE)


def ensure_gpp() -> None:
    if shutil.which("g++"):
        return
    print("Installing build-essential for g++...")
    subprocess.run(
        ["bash", "-lc", "apt-get update -qq && apt-get install -y -qq build-essential"],
        check=True,
    )
    if not shutil.which("g++"):
        raise RuntimeError("g++ not available after apt install")


def extract_code(text: str) -> str:
    m = CODE_FENCE_RE.search(text)
    if m:
        return m.group(1).strip() + "\n"
    lines = text.splitlines()
    start = 0
    for i, line in enumerate(lines):
        if line.lstrip().startswith("#include") or re.match(r"\s*int\s+main\b", line):
            start = i
            break
    return "\n".join(lines[start:]).strip() + "\n"


def judge_code(code: str, expected_stdout: str | None = None) -> float:
    code = extract_code(code)
    if not code.strip():
        return 0.0
    with tempfile.TemporaryDirectory(prefix="grpo_judge_") as tmp:
        root = Path(tmp)
        src = root / "prog.cpp"
        bin_path = root / "prog"
        src.write_text(code, encoding="utf-8")
        try:
            cp = subprocess.run(
                ["g++", "-std=c++20", "-O0", "-Wall", "-o", str(bin_path), str(src)],
                capture_output=True,
                text=True,
                timeout=15.0,
            )
        except subprocess.TimeoutExpired:
            return 0.0
        if cp.returncode != 0:
            return 0.0
        reward = 1.0
        if expected_stdout:
            try:
                rp = subprocess.run(
                    [str(bin_path)],
                    capture_output=True,
                    text=True,
                    timeout=5.0,
                )
                if rp.returncode == 0 and (rp.stdout or "") == expected_stdout:
                    reward += 0.5
                else:
                    reward = max(reward - 0.25, 0.5)
            except subprocess.TimeoutExpired:
                reward = max(reward - 0.25, 0.5)
        return round(reward, 3)


def completion_text(completion) -> str:
    if isinstance(completion, list):
        if completion and isinstance(completion[-1], dict):
            return str(completion[-1].get("content", ""))
        return str(completion)
    return str(completion)


def compile_reward(
    prompts,
    completions,
    expected_stdout=None,
    **kwargs,
) -> list[float]:
    rewards: list[float] = []
    for i, completion in enumerate(completions):
        text = completion_text(completion)
        exp = None
        if expected_stdout is not None:
            exp = expected_stdout[i] if expected_stdout[i] else None
        rewards.append(judge_code(text, expected_stdout=exp))
    return rewards


def load_policy():
    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"SFT merge skipped ({e})")
    try:
        model = PeftModel.from_pretrained(model, DPO_ADAPTER)
        model = model.merge_and_unload()
        print(f"Merged DPO adapter from {DPO_ADAPTER}")
    except Exception as e:
        print(f"DPO merge skipped ({e})")
    return model, tokenizer


def main() -> None:
    ensure_gpp()
    ds = load_dataset(DATASET_ID, split="train")
    if "prompt" not in ds.column_names:
        raise SystemExit(f"Dataset needs 'prompt' column; got {ds.column_names}")

    model, tokenizer = load_policy()

    trainer = GRPOTrainer(
        model=model,
        processing_class=tokenizer,
        reward_funcs=[compile_reward],
        train_dataset=ds,
        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=GRPOConfig(
            output_dir=OUTPUT_DIR,
            num_train_epochs=1,
            per_device_train_batch_size=1,
            gradient_accumulation_steps=4,
            num_generations=4,
            max_completion_length=512,
            learning_rate=5e-6,
            logging_steps=5,
            save_strategy="steps",
            save_steps=50,
            save_total_limit=1,
            temperature=0.7,
            bf16=True,
            remove_unused_columns=False,
            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 GRPO model to {HUB_MODEL_ID}")


if __name__ == "__main__":
    main()