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:
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -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")
|
||||||
Reference in New Issue
Block a user