Download scripts/audit_multibench_results.py from Orangerl/umm: direct link, hf CLI and curl.
- Browser
- Download file 7.93 kB
-
https://huggingface.co/Orangerl/umm/resolve/main/scripts/audit_multibench_results.py
- Command line
-
hf download hf://Orangerl/umm/scripts/audit_multibench_results.py
-
curl -L -o audit_multibench_results.py https://huggingface.co/Orangerl/umm/resolve/main/scripts/audit_multibench_results.py
7.93 kB
| #!/usr/bin/env python3 | |
| """Independently verify multibench result completeness and recompute scores.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| import re | |
| from collections import Counter, defaultdict | |
| from pathlib import Path | |
| import pyarrow.parquet as pq | |
| ANSWER_RE = re.compile(r"^\s*\(?([A-Z])\)?\s*$", re.IGNORECASE) | |
| def _write_json(path: Path, payload: dict) -> None: | |
| path.write_text( | |
| json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True) + "\n", | |
| encoding="utf-8", | |
| ) | |
| def _expected(value: str) -> str: | |
| match = ANSWER_RE.fullmatch(str(value)) | |
| if not match: | |
| raise ValueError(f"Invalid answer: {value!r}") | |
| return match.group(1).upper() | |
| def _read_jsonl(path: Path) -> list[dict]: | |
| return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] | |
| def _cvbench_official_scores( | |
| manifest: list[dict], result_by_id: dict[str, dict], raw_root: Path | |
| ) -> dict: | |
| source_by_idx: dict[int, str] = {} | |
| for parquet_path in (raw_root / "CV-Bench/test_2d.parquet", raw_root / "CV-Bench/test_3d.parquet"): | |
| table = pq.read_table(parquet_path, columns=["idx", "source"]) | |
| for row in table.to_pylist(): | |
| source_by_idx[int(row["idx"])] = str(row["source"]) | |
| source_scores: dict[str, list[int]] = defaultdict(list) | |
| for source_row in manifest: | |
| source = source_by_idx[int(source_row["source_idx"])] | |
| source_scores[source].append(int(result_by_id[str(source_row["sample_id"])]["correct"])) | |
| required = {"ADE20K", "COCO", "Omni3D"} | |
| if source_scores.keys() != required: | |
| raise RuntimeError(f"Unexpected CV-Bench sources: {sorted(source_scores)}") | |
| components = { | |
| source: { | |
| "correct": sum(values), | |
| "samples": len(values), | |
| "accuracy": sum(values) / len(values), | |
| } | |
| for source, values in sorted(source_scores.items()) | |
| } | |
| accuracy_2d = (components["ADE20K"]["accuracy"] + components["COCO"]["accuracy"]) / 2 | |
| accuracy_3d = components["Omni3D"]["accuracy"] | |
| return { | |
| "source_scores": components, | |
| "official_2d_accuracy": accuracy_2d, | |
| "official_3d_accuracy": accuracy_3d, | |
| "official_cvbench_accuracy": (accuracy_2d + accuracy_3d) / 2, | |
| "formula": "((ADE20K accuracy + COCO accuracy) / 2 + Omni3D accuracy) / 2", | |
| } | |
| def audit_one(manifest_path: Path, result_dir: Path, raw_root: Path) -> dict: | |
| manifest = _read_jsonl(manifest_path) | |
| results = _read_jsonl(result_dir / "results.jsonl") | |
| metrics_file = json.loads((result_dir / "metrics.json").read_text(encoding="utf-8")) | |
| benchmark = str(metrics_file["metadata"]["benchmark"]) | |
| by_id = {str(row["sample_id"]): row for row in manifest} | |
| result_by_id = {str(row["sample_id"]): row for row in results} | |
| if len(by_id) != len(manifest) or len(result_by_id) != len(results): | |
| raise RuntimeError(f"Duplicate IDs in {benchmark}") | |
| if by_id.keys() != result_by_id.keys(): | |
| missing = sorted(by_id.keys() - result_by_id.keys())[:5] | |
| extra = sorted(result_by_id.keys() - by_id.keys())[:5] | |
| raise RuntimeError(f"ID mismatch in {benchmark}: missing={missing}, extra={extra}") | |
| errors = [] | |
| correct = 0 | |
| parsed = 0 | |
| config = defaultdict(lambda: [0, 0]) | |
| task = defaultdict(lambda: [0, 0]) | |
| token_lengths = [] | |
| for sample_id, source in by_id.items(): | |
| row = result_by_id[sample_id] | |
| expected = _expected(source["answer"]) | |
| if row["expected"] != expected: | |
| raise RuntimeError(f"Expected-answer mismatch for {benchmark}/{sample_id}") | |
| if row["error"] is not None: | |
| errors.append({"sample_id": sample_id, "error": row["error"]}) | |
| prediction = row["predicted"] | |
| if prediction is not None: | |
| parsed += 1 | |
| if not (len(prediction) == 1 and "A" <= prediction <= chr(ord("A") + len(source["choices"]) - 1)): | |
| raise RuntimeError(f"Illegal prediction for {benchmark}/{sample_id}: {prediction!r}") | |
| is_correct = prediction == expected | |
| if bool(row["correct"]) != is_correct: | |
| raise RuntimeError(f"Correct flag mismatch for {benchmark}/{sample_id}") | |
| correct += int(is_correct) | |
| config[str(source["config"])][0] += int(is_correct) | |
| config[str(source["config"])][1] += 1 | |
| task[str(source["task"])][0] += int(is_correct) | |
| task[str(source["task"])][1] += 1 | |
| token_lengths.append(len(row["generated_token_ids"])) | |
| sample_count = len(results) | |
| accuracy = correct / sample_count | |
| parse_rate = parsed / sample_count | |
| published = metrics_file["metrics"] | |
| for key, recomputed in { | |
| f"eval/{benchmark}_samples": float(sample_count), | |
| f"eval/{benchmark}_errors": float(len(errors)), | |
| f"eval/{benchmark}_accuracy": accuracy, | |
| f"eval/{benchmark}_parse_rate": parse_rate, | |
| }.items(): | |
| if key not in published or not math.isclose(float(published[key]), recomputed, abs_tol=1e-12): | |
| raise RuntimeError( | |
| f"Published metric mismatch for {key}: {published.get(key)!r} != {recomputed!r}" | |
| ) | |
| sorted_lengths = sorted(token_lengths) | |
| audit = { | |
| "benchmark": benchmark, | |
| "split": metrics_file["metadata"]["split"], | |
| "samples": sample_count, | |
| "correct": correct, | |
| "accuracy": accuracy, | |
| "parsed": parsed, | |
| "parse_rate": parse_rate, | |
| "errors": len(errors), | |
| "config_scores": { | |
| name: {"correct": values[0], "samples": values[1], "accuracy": values[0] / values[1]} | |
| for name, values in sorted(config.items()) | |
| }, | |
| "task_scores": { | |
| name: {"correct": values[0], "samples": values[1], "accuracy": values[0] / values[1]} | |
| for name, values in sorted(task.items()) | |
| }, | |
| "generated_tokens": { | |
| "min": min(sorted_lengths), | |
| "median": sorted_lengths[len(sorted_lengths) // 2], | |
| "p95": sorted_lengths[int(0.95 * (len(sorted_lengths) - 1))], | |
| "max": max(sorted_lengths), | |
| "at_256_cap": sum(length >= 256 for length in sorted_lengths), | |
| }, | |
| "checks": { | |
| "manifest_result_ids_exact": True, | |
| "expected_answers_exact": True, | |
| "correct_flags_recomputed": True, | |
| "published_core_metrics_recomputed": True, | |
| "all_inference_errors_zero": not errors, | |
| "no_generation_hit_cap": max(sorted_lengths) < 256, | |
| }, | |
| } | |
| if benchmark == "cvbench": | |
| audit["official_scoring"] = _cvbench_official_scores( | |
| manifest, result_by_id, raw_root | |
| ) | |
| if errors: | |
| raise RuntimeError(f"{benchmark} has {len(errors)} inference errors") | |
| if max(sorted_lengths) >= 256: | |
| raise RuntimeError(f"{benchmark} has generation(s) at the token cap") | |
| _write_json(result_dir / "audit.json", audit) | |
| return audit | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--root", type=Path, required=True) | |
| parser.add_argument("--prepared-root", type=Path, default=Path("data/benchmarks/prepared")) | |
| args = parser.parse_args() | |
| root = args.root.resolve() | |
| prepared = args.prepared_root.resolve() | |
| raw = prepared.parent / "raw" | |
| jobs = [ | |
| (prepared / "cvbench_full/manifest.jsonl", root / "cvbench_full_test"), | |
| (prepared / "blink_val/manifest.jsonl", root / "blink_full_val"), | |
| (prepared / "vstar_test/manifest.jsonl", root / "vstar_full_test"), | |
| ] | |
| audits = [audit_one(manifest, result, raw) for manifest, result in jobs] | |
| _write_json(root / "audit_summary.json", {row["benchmark"]: row for row in audits}) | |
| print(json.dumps({row["benchmark"]: row["accuracy"] for row in audits}, sort_keys=True)) | |
| if __name__ == "__main__": | |
| main() | |