From 54f2b3c6cca7be8d78e0764d7dd4f9d0d66615b0 Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 09:45:46 +0200 Subject: [PATCH] MultiPL-E Rust harness: compile-run verifier, STaR labeling, transfer eval Co-Authored-By: Claude Fable 5 --- scripts/eval_rust.py | 86 ++++++++++++++++++++++++++++++++++++++++++ scripts/prep_rust.py | 66 ++++++++++++++++++++++++++++++++ scripts/rust_common.py | 66 ++++++++++++++++++++++++++++++++ 3 files changed, 218 insertions(+) create mode 100644 scripts/eval_rust.py create mode 100644 scripts/prep_rust.py create mode 100644 scripts/rust_common.py diff --git a/scripts/eval_rust.py b/scripts/eval_rust.py new file mode 100644 index 0000000..2c58310 --- /dev/null +++ b/scripts/eval_rust.py @@ -0,0 +1,86 @@ +"""Rust eval: pass@1 vs k with the frozen-prompt loop (transfer or trained).""" + +import argparse +import json +import os +import sys +import time +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import torch + +from loop_common import BandLooper, MergeAdapter +from rust_common import DIRECT_SUFFIX, extract_rust, run_rust_tests, rust_prompt + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) +from jlens.core import load_model # noqa: E402 + +OUT = Path(os.environ.get("LOOP_OUT", + Path(__file__).resolve().parent.parent / "results-loop")) + + +@torch.no_grad() +def pass1(looper, adapter, tok, items, k, batch=8, max_new=300): + codes = [] + for i in range(0, len(items), batch): + torch.cuda.empty_cache() + chunk = items[i : i + batch] + enc = tok([rust_prompt(tok, it, DIRECT_SUFFIX) for it in chunk], + return_tensors="pt", padding=True, + add_special_tokens=False).to("cuda") + gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k, + max_new_tokens=max_new, + attention_mask=enc["attention_mask"]) + for j in range(len(chunk)): + codes.append(extract_rust( + tok.decode(gen[j, enc["input_ids"].shape[1]:], + skip_special_tokens=True))) + with ThreadPoolExecutor(16) as ex: + oks = list(ex.map(lambda ci: run_rust_tests(ci[0], ci[1]), + zip(codes, items))) + by = {} + for it, ok in zip(items, oks): + d = by.setdefault(it["label"], [0, 0]) + d[0] += ok + d[1] += 1 + per_item = [{"task_id": it["task_id"], "ok": bool(ok)} + for it, ok in zip(items, oks)] + return sum(oks) / len(items), {l: c / n for l, (c, n) in by.items()}, per_item + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--adapter", default=None) + ap.add_argument("--tag", default="rust") + ap.add_argument("--ks", default="0,2,4") + ap.add_argument("--split", default="test") + args = ap.parse_args() + ks = [int(x) for x in args.ks.split(",")] + + model, tok = load_model(dtype=torch.bfloat16) + tok.padding_side = "left" + looper = BandLooper(model) + adapter = MergeAdapter( + d=model.config.get_text_config().hidden_size).cuda() + if args.adapter: + adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) + adapter.eval() + + items = [it for it in json.load(open(OUT / "rust_data.json")) + if it["split"] == args.split] + print(f"[{args.tag}] Rust eval on {len(items)} items, ks={ks}", flush=True) + res = {"tag": args.tag, "ks": {}, "n": len(items)} + for k in ks: + t0 = time.time() + acc, by, per_item = pass1(looper, adapter, tok, items, k) + res["ks"][k] = {"acc": acc, "by_label": by, "per_item": per_item} + print(f"k={k}: pass@1={acc:.3f} " + f"by_label={ {l: round(v,3) for l,v in by.items()} }" + f" ({time.time()-t0:.0f}s)", flush=True) + json.dump(res, open(OUT / f"eval_rust_{args.tag}.json", "w"), indent=1) + print("wrote", OUT / f"eval_rust_{args.tag}.json") + + +if __name__ == "__main__": + main() diff --git a/scripts/prep_rust.py b/scripts/prep_rust.py new file mode 100644 index 0000000..0a9322c --- /dev/null +++ b/scripts/prep_rust.py @@ -0,0 +1,66 @@ +"""STaR labeling of MultiPL-E Rust (mbpp-rs) with the frozen model.""" + +import json +import os +import sys +import time +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import torch +from datasets import load_dataset + +from prep_mbpp import batch_generate +from rust_common import (DIRECT_SUFFIX, PLAN_SUFFIX, extract_rust, + run_rust_tests, rust_prompt) + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) +from jlens.core import load_model # noqa: E402 + +OUT = Path(os.environ.get("LOOP_OUT", + Path(__file__).resolve().parent.parent / "results-loop")) + + +def main(): + model, tok = load_model(dtype=torch.bfloat16) + tok.padding_side = "left" + items = [dict(r) for r in load_dataset("nuprl/MultiPL-E", "mbpp-rs", + split="test")] + # split: first 200 train-pool, rest test (354 total) + for i, it in enumerate(items): + it["split"] = "train" if i < 200 else "test" + it["task_id"] = it["name"] + print(f"{len(items)} rust items", flush=True) + + for tag, suffix, mx in (("direct", DIRECT_SUFFIX, 300), + ("plan", PLAN_SUFFIX, 700)): + t0 = time.time() + gens = batch_generate(model, tok, + [rust_prompt(tok, it, suffix) for it in items], + max_new_tokens=mx, batch_size=24) + codes = [extract_rust(g) for g in gens] + with ThreadPoolExecutor(16) as ex: + oks = list(ex.map(lambda ci: run_rust_tests(ci[0], ci[1]), + zip(codes, items))) + for it, c, ok in zip(items, codes, oks): + it[f"{tag}_ok"] = bool(ok) + it[f"{tag}_code"] = c if ok else None + print(f"{tag}: pass@1={sum(oks)/len(items):.3f} " + f"({time.time()-t0:.0f}s)", flush=True) + + for it in items: + it["label"] = ("easy" if it["direct_ok"] + else "hard" if it["plan_ok"] else "drop") + it["sol_code"] = (it["direct_code"] if it["direct_ok"] + else it["plan_code"]) + for split in ("train", "test"): + sub = [it for it in items if it["split"] == split] + print(f"{split}: easy={sum(i['label']=='easy' for i in sub)} " + f"hard={sum(i['label']=='hard' for i in sub)} " + f"drop={sum(i['label']=='drop' for i in sub)}", flush=True) + json.dump(items, open(OUT / "rust_data.json", "w"), indent=1) + print("wrote", OUT / "rust_data.json") + + +if __name__ == "__main__": + main() diff --git a/scripts/rust_common.py b/scripts/rust_common.py new file mode 100644 index 0000000..e44aab5 --- /dev/null +++ b/scripts/rust_common.py @@ -0,0 +1,66 @@ +"""MultiPL-E Rust (mbpp-rs / humaneval-rs): prompts + compile-run verifier. + +Cross-LANGUAGE transfer testbed: same problems as our Python MBPP/HumanEval, +in Rust. Verifier: rustc compile of [function + tests-main], run, exit code. +""" + +import os +import re +import subprocess +import tempfile +from pathlib import Path + +RUSTC = os.path.expanduser("~/.cargo/bin/rustc") +CODE_RE = re.compile(r"```(?:rust|rs)?\s*\n(.*?)```", re.S) + +DIRECT_SUFFIX = ("\n\nWrite ONLY the complete Rust function (signature " + "included, no main) in a ```rust code block. No explanation.") +PLAN_SUFFIX = ("\n\nFirst write a very brief plan: at most 4 short bullet " + "lines. Then write the complete Rust function (signature " + "included, no main) in a ```rust code block.") + + +def rust_prompt(tok, item, suffix): + msg = ("Complete this Rust function:\n\n```rust\n" + item["prompt"] + + "\n```" + suffix) + return tok.apply_chat_template([{"role": "user", "content": msg}], + tokenize=False, add_generation_prompt=True) + + +def extract_rust(text): + m = CODE_RE.findall(text) + return m[-1].strip() if m else None + + +def run_rust_tests(code, item, timeout=30): + if not code: + return False + tests = item["tests"] + if tests.lstrip().startswith("}"): # designed to close an open fn body + tests = tests.lstrip()[1:] + program = code + "\n" + tests + try: + with tempfile.TemporaryDirectory() as td: + src = Path(td) / "main.rs" + src.write_text(program) + c = subprocess.run([RUSTC, "-O", "--edition", "2021", "-o", + str(Path(td) / "prog"), str(src)], + capture_output=True, timeout=timeout) + if c.returncode != 0: + return False + r = subprocess.run([str(Path(td) / "prog")], capture_output=True, + timeout=10) + return r.returncode == 0 + except (subprocess.TimeoutExpired, OSError): + return False + + +if __name__ == "__main__": + item = {"prompt": "/// doubles x\nfn double(x: isize) -> isize {\n", + "tests": "}\n\nfn main() {\n assert_eq!(double(2), 4);\n}"} + good = "fn double(x: isize) -> isize {\n x * 2\n}" + bad = "fn double(x: isize) -> isize {\n x + 1\n}" + assert run_rust_tests(good, item), "good should pass" + assert not run_rust_tests(bad, item), "bad should fail" + assert not run_rust_tests("garbage", item), "non-compiling should fail" + print("rust verifier self-test ok")