# /// 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()