File size: 5,770 Bytes
31dc8dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | """Compare two lm-eval code-task run directories."""
from __future__ import annotations
import argparse
import ast
import json
import statistics
from pathlib import Path
from diffulex_bench.tasks.humaneval.llada2_utils import extract_code as extract_humaneval_code
from diffulex_bench.tasks.mbpp.llada2_utils import (
_code_text_for_prediction as mbpp_code_text_for_prediction,
)
from diffulex_bench.tasks.mbpp.llada2_utils import extract_code as extract_mbpp_code
def _first_file(run_dir: Path, pattern: str) -> Path:
matches = sorted(run_dir.glob(pattern))
if not matches:
raise FileNotFoundError(f"No file matching {pattern!r} under {run_dir}")
return matches[0]
def _sample_path(run_dir: Path) -> Path:
return _first_file(run_dir, "**/samples_*.jsonl")
def _result_path(run_dir: Path) -> Path:
return _first_file(run_dir, "**/results_*.json")
def _load_samples(run_dir: Path) -> dict[int, dict]:
path = _sample_path(run_dir)
rows = [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines()]
return {int(row["doc_id"]): row for row in rows}
def _extract_code(row: dict) -> str:
pred = row["filtered_resps"][0]
doc = row["doc"]
task_id = str(doc.get("task_id", "")).lower()
if "humaneval" in task_id:
text = pred if "```" in pred else doc.get("prefix", "") + pred
return extract_humaneval_code(text)
return extract_mbpp_code(mbpp_code_text_for_prediction(doc, pred))
def _syntax_error(row: dict) -> bool:
try:
ast.parse(_extract_code(row))
return False
except SyntaxError:
return True
def _score(row: dict) -> int:
return int(row.get("exact_match", 0))
def _short(text: str, limit: int) -> str:
text = text.replace("\r", "")
if len(text) <= limit:
return text
return text[:limit] + "\n...<truncated>"
def _print_run_summary(label: str, run_dir: Path, rows: dict[int, dict]) -> None:
result = json.loads(_result_path(run_dir).read_text(encoding="utf-8"))
task_name = next(iter(result["results"]))
scores = [_score(row) for row in rows.values()]
preds = [row["filtered_resps"][0] for row in rows.values()]
lengths = [len(pred) for pred in preds]
syntax = sum(_syntax_error(row) for row in rows.values())
masks = sum("<|mask|>" in pred for pred in preds)
no_fence = sum("```" not in pred for pred in preds)
long = sum(len(pred) > 3000 for pred in preds)
metadata = result["configs"][task_name].get("metadata", {})
print(f"\n### {label}: {run_dir.name}")
print("result:", result["results"][task_name])
print(f"pass: {sum(scores)}/{len(scores)} = {sum(scores) / len(scores):.6f}")
print(
"metadata:",
{
k: metadata.get(k)
for k in (
"buffer_size",
"block_size",
"max_nfe",
"max_new_tokens",
"decoding_strategy",
"sampling_mode",
"accept_threshold",
"remask_threshold",
"token_merge_mode",
"token_merge_top_k",
)
},
)
print(
"response_chars mean/p50/max:",
round(statistics.mean(lengths), 1),
statistics.median(lengths),
max(lengths),
)
print("syntax_errors:", syntax, "mask_outputs:", masks, "no_fence:", no_fence, "long>3000:", long)
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("base_run", type=Path)
parser.add_argument("candidate_run", type=Path)
parser.add_argument("--examples", type=int, default=5)
parser.add_argument("--chars", type=int, default=900)
args = parser.parse_args()
base = _load_samples(args.base_run)
cand = _load_samples(args.candidate_run)
ids = sorted(set(base) & set(cand))
_print_run_summary("base", args.base_run, base)
_print_run_summary("candidate", args.candidate_run, cand)
both_right = [i for i in ids if _score(base[i]) == 1 and _score(cand[i]) == 1]
regress = [i for i in ids if _score(base[i]) == 1 and _score(cand[i]) == 0]
improve = [i for i in ids if _score(base[i]) == 0 and _score(cand[i]) == 1]
both_wrong = [i for i in ids if _score(base[i]) == 0 and _score(cand[i]) == 0]
print("\n=== migration ===")
print("both_right:", len(both_right), "regress:", len(regress), "improve:", len(improve), "both_wrong:", len(both_wrong))
print("regress first:", regress[:20])
print("improve first:", improve[:20])
print("regress syntax in base/candidate:", sum(_syntax_error(base[i]) for i in regress), sum(_syntax_error(cand[i]) for i in regress))
for title, selected in (("REGRESS", regress[: args.examples]), ("IMPROVE", improve[: args.examples])):
print(f"\n=== {title} examples ===")
for doc_id in selected:
before = base[doc_id]
after = cand[doc_id]
doc = before["doc"]
prompt = doc.get("prompt") or doc.get("text") or doc.get("question") or ""
tests = "\n".join(doc.get("test_list", []))
print(f"\n--- doc_id {doc_id} task_id {doc.get('task_id')} ---")
print("prompt:", _short(prompt, 240))
print("tests:", _short(tests, 360))
print("[base]", _score(before), "syntax_error", _syntax_error(before), "chars", len(before["filtered_resps"][0]))
print(_short(before["filtered_resps"][0], args.chars))
print("[candidate]", _score(after), "syntax_error", _syntax_error(after), "chars", len(after["filtered_resps"][0]))
print(_short(after["filtered_resps"][0], args.chars))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|