gonzalolinares commited on
Commit
625c0f0
·
verified ·
1 Parent(s): d5c7535

GRPO compiler reward script

Browse files
Files changed (1) hide show
  1. train_grpo.py +185 -0
train_grpo.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # dependencies = ["trl>=0.12.0", "peft>=0.7.0", "datasets", "transformers", "accelerate", "torch"]
3
+ # ///
4
+ """GRPO with g++ compiler reward (online RL). For Hugging Face Jobs (uv)."""
5
+
6
+ from __future__ import annotations
7
+
8
+ import os
9
+ import re
10
+ import shutil
11
+ import subprocess
12
+ import tempfile
13
+ from pathlib import Path
14
+
15
+ from datasets import load_dataset
16
+ from peft import LoraConfig, PeftModel
17
+ from transformers import AutoModelForCausalLM, AutoTokenizer
18
+ from trl import GRPOConfig, GRPOTrainer
19
+
20
+ DATASET_ID = os.environ.get("DATASET_ID", "gonzalolinares/cpp-compiler-grpo")
21
+ SFT_ADAPTER = os.environ.get("BASE_MODEL", "gonzalolinares/qwen25-1.5b-cpp-sft")
22
+ DPO_ADAPTER = os.environ.get("DPO_MODEL", "gonzalolinares/qwen25-1.5b-cpp-dpo")
23
+ BASE_MODEL = os.environ.get("FALLBACK_MODEL", "Qwen/Qwen2.5-1.5B-Instruct")
24
+ HUB_MODEL_ID = os.environ.get("HUB_MODEL_ID", "gonzalolinares/qwen25-1.5b-cpp-grpo")
25
+ OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "qwen25-1.5b-cpp-grpo")
26
+
27
+ CODE_FENCE_RE = re.compile(r"```(?:cpp|c\+\+)?\s*([\s\S]*?)```", re.IGNORECASE)
28
+
29
+
30
+ def ensure_gpp() -> None:
31
+ if shutil.which("g++"):
32
+ return
33
+ print("Installing build-essential for g++...")
34
+ subprocess.run(
35
+ ["bash", "-lc", "apt-get update -qq && apt-get install -y -qq build-essential"],
36
+ check=True,
37
+ )
38
+ if not shutil.which("g++"):
39
+ raise RuntimeError("g++ not available after apt install")
40
+
41
+
42
+ def extract_code(text: str) -> str:
43
+ m = CODE_FENCE_RE.search(text)
44
+ if m:
45
+ return m.group(1).strip() + "\n"
46
+ lines = text.splitlines()
47
+ start = 0
48
+ for i, line in enumerate(lines):
49
+ if line.lstrip().startswith("#include") or re.match(r"\s*int\s+main\b", line):
50
+ start = i
51
+ break
52
+ return "\n".join(lines[start:]).strip() + "\n"
53
+
54
+
55
+ def judge_code(code: str, expected_stdout: str | None = None) -> float:
56
+ code = extract_code(code)
57
+ if not code.strip():
58
+ return 0.0
59
+ with tempfile.TemporaryDirectory(prefix="grpo_judge_") as tmp:
60
+ root = Path(tmp)
61
+ src = root / "prog.cpp"
62
+ bin_path = root / "prog"
63
+ src.write_text(code, encoding="utf-8")
64
+ try:
65
+ cp = subprocess.run(
66
+ ["g++", "-std=c++20", "-O0", "-Wall", "-o", str(bin_path), str(src)],
67
+ capture_output=True,
68
+ text=True,
69
+ timeout=15.0,
70
+ )
71
+ except subprocess.TimeoutExpired:
72
+ return 0.0
73
+ if cp.returncode != 0:
74
+ return 0.0
75
+ reward = 1.0
76
+ if expected_stdout:
77
+ try:
78
+ rp = subprocess.run(
79
+ [str(bin_path)],
80
+ capture_output=True,
81
+ text=True,
82
+ timeout=5.0,
83
+ )
84
+ if rp.returncode == 0 and (rp.stdout or "") == expected_stdout:
85
+ reward += 0.5
86
+ else:
87
+ reward = max(reward - 0.25, 0.5)
88
+ except subprocess.TimeoutExpired:
89
+ reward = max(reward - 0.25, 0.5)
90
+ return round(reward, 3)
91
+
92
+
93
+ def completion_text(completion) -> str:
94
+ if isinstance(completion, list):
95
+ if completion and isinstance(completion[-1], dict):
96
+ return str(completion[-1].get("content", ""))
97
+ return str(completion)
98
+ return str(completion)
99
+
100
+
101
+ def compile_reward(
102
+ prompts,
103
+ completions,
104
+ expected_stdout=None,
105
+ **kwargs,
106
+ ) -> list[float]:
107
+ rewards: list[float] = []
108
+ for i, completion in enumerate(completions):
109
+ text = completion_text(completion)
110
+ exp = None
111
+ if expected_stdout is not None:
112
+ exp = expected_stdout[i] if expected_stdout[i] else None
113
+ rewards.append(judge_code(text, expected_stdout=exp))
114
+ return rewards
115
+
116
+
117
+ def load_policy():
118
+ tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
119
+ if tokenizer.pad_token is None:
120
+ tokenizer.pad_token = tokenizer.eos_token
121
+ model = AutoModelForCausalLM.from_pretrained(BASE_MODEL, torch_dtype="auto")
122
+ try:
123
+ model = PeftModel.from_pretrained(model, SFT_ADAPTER)
124
+ model = model.merge_and_unload()
125
+ print(f"Merged SFT adapter from {SFT_ADAPTER}")
126
+ except Exception as e:
127
+ print(f"SFT merge skipped ({e})")
128
+ try:
129
+ model = PeftModel.from_pretrained(model, DPO_ADAPTER)
130
+ model = model.merge_and_unload()
131
+ print(f"Merged DPO adapter from {DPO_ADAPTER}")
132
+ except Exception as e:
133
+ print(f"DPO merge skipped ({e})")
134
+ return model, tokenizer
135
+
136
+
137
+ def main() -> None:
138
+ ensure_gpp()
139
+ ds = load_dataset(DATASET_ID, split="train")
140
+ if "prompt" not in ds.column_names:
141
+ raise SystemExit(f"Dataset needs 'prompt' column; got {ds.column_names}")
142
+
143
+ model, tokenizer = load_policy()
144
+
145
+ trainer = GRPOTrainer(
146
+ model=model,
147
+ processing_class=tokenizer,
148
+ reward_funcs=[compile_reward],
149
+ train_dataset=ds,
150
+ peft_config=LoraConfig(
151
+ r=16,
152
+ lora_alpha=32,
153
+ lora_dropout=0.05,
154
+ bias="none",
155
+ task_type="CAUSAL_LM",
156
+ target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
157
+ ),
158
+ args=GRPOConfig(
159
+ output_dir=OUTPUT_DIR,
160
+ num_train_epochs=1,
161
+ per_device_train_batch_size=1,
162
+ gradient_accumulation_steps=4,
163
+ num_generations=4,
164
+ max_completion_length=512,
165
+ learning_rate=5e-6,
166
+ logging_steps=5,
167
+ save_strategy="steps",
168
+ save_steps=50,
169
+ save_total_limit=1,
170
+ max_length=1024,
171
+ temperature=0.7,
172
+ bf16=True,
173
+ push_to_hub=False,
174
+ hub_model_id=HUB_MODEL_ID,
175
+ report_to="none",
176
+ ),
177
+ )
178
+ trainer.train()
179
+ trainer.model.push_to_hub(HUB_MODEL_ID, private=False)
180
+ tokenizer.push_to_hub(HUB_MODEL_ID, private=False)
181
+ print(f"Pushed GRPO model to {HUB_MODEL_ID}")
182
+
183
+
184
+ if __name__ == "__main__":
185
+ main()