MultiPL-E Rust harness: compile-run verifier, STaR labeling, transfer eval

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-14 09:45:46 +02:00
co-authored by Claude Fable 5
parent 6c39c38eb9
commit 54f2b3c6cc
3 changed files with 218 additions and 0 deletions
+86
View File
@@ -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()
+66
View File
@@ -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()
+66
View File
@@ -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")