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