Zhouhh123's picture
Implement ARK-Bench leaderboard
4cc47cb
Raw History Blame Contribute Delete
14.4 kB
"""Pure Python implementation of the topic-gallery-top20-v1 protocol.
Prediction JSONL rows contain exactly topic, query_id, and retrieved_ids.
Label rows contain topic, query_id, and target (private auxiliary fields are
ignored). IDs are topic-local nonnegative integers. No error includes input values.
"""
from __future__ import annotations
import json
import math
from copy import deepcopy
from datetime import datetime, timezone
from pathlib import Path
from src.leaderboard.schema import (
ACTIVE_BENCHMARK_VERSION, COGNITIVE_LABELS, KNOWLEDGE_GROUPS, KNOWLEDGE_LABELS,
METRICS, REASONING_LABELS,
)
from src.leaderboard.validate_scores import validate_score_data
MAX_BYTES = 2 * 1024 * 1024
MAX_ROW_BYTES = 16 * 1024
QUERY_COUNT = 1750
PROTOCOL = "topic-gallery-top20-v1"
_CONFIG_ERROR = "Invalid evaluation configuration."
SKILL_LABELS = {
"knowledge_reasoning": "KR", "spatial_reasoning": "SR",
"fine_grained_visual_reasoning": "FGVR", "logical_reasoning": "LR",
"symbolic_reasoning": "SyR", "conceptual_abstraction": "CA",
}
DEMAND_LABELS = {
frozenset(["None"]): "PO",
frozenset(["knowledge-intensive"]): "KI",
frozenset(["reasoning-intensive"]): "RI",
frozenset(["knowledge-intensive", "reasoning-intensive"]): "KRI",
}
class SubmissionError(ValueError):
"""A submission violates the public prediction protocol."""
class LabelSet(dict):
"""Private targets plus private grouping metadata used during scoring."""
def __init__(self, targets, annotations, knowledge_by_topic):
super().__init__(targets)
self.annotations = annotations
self.knowledge_by_topic = knowledge_by_topic
def _object(pairs):
result = {}
for key, value in pairs:
if key in result:
raise ValueError
result[key] = value
return result
def _constant(value):
raise ValueError
def _decode(line):
return json.loads(line.decode("utf-8"), object_pairs_hook=_object, parse_constant=_constant)
def _key(row):
if not isinstance(row, dict) or not isinstance(row.get("topic"), str) or type(row.get("query_id")) is not int or row["query_id"] < 0:
raise ValueError
return row["topic"], row["query_id"]
def _ids(values, gallery, count):
return (isinstance(values, list) and len(values) == count
and all(type(value) is int and value >= 0 and value in gallery for value in values)
and len(set(values)) == count)
def _manifest(manifest):
try:
if (manifest["benchmark_version"] != "1.0"
or manifest["benchmark_version"] != ACTIVE_BENCHMARK_VERSION
or type(manifest["top_k"]) is not int or manifest["top_k"] != 20
or not isinstance(manifest["dataset_revision"], str) or not manifest["dataset_revision"].strip()):
raise ValueError
topics = manifest["topics"]
if not isinstance(topics, dict) or len(topics) != len(KNOWLEDGE_LABELS):
raise ValueError
galleries = {}
knowledge = []
for topic, data in topics.items():
if not isinstance(topic, str) or not topic:
raise ValueError
values = data["gallery_ids"]
if (not isinstance(values, list) or len(values) < 20
or any(type(value) is not int or value < 0 for value in values) or len(set(values)) != len(values)):
raise ValueError
galleries[topic] = set(values)
knowledge.append(data["knowledge_label"])
if set(knowledge) != set(KNOWLEDGE_LABELS):
raise ValueError
rows = manifest["queries"]
if not isinstance(rows, list) or len(rows) != QUERY_COUNT:
raise ValueError
queries = {}
used_topics = set()
for row in rows:
key = _key(row)
if key in queries or key[0] not in galleries:
raise ValueError
queries[key] = row
used_topics.add(key[0])
if used_topics != set(topics):
raise ValueError
return queries, galleries
except (KeyError, TypeError, ValueError, AttributeError):
raise ValueError(_CONFIG_ERROR) from None
def validate_predictions(content: bytes, manifest: dict) -> dict[tuple[str, int], list[int]]:
"""Validate strict UTF-8 JSONL with complete coverage and 20 unique IDs.
LF and CRLF are accepted, including one optional final line ending. Blank
lines, BOMs, non-finite JSON numbers, and duplicate object keys are rejected.
Limits count bytes; the row limit excludes its LF/CRLF line ending.
"""
queries, galleries = _manifest(manifest)
if not isinstance(content, bytes) or len(content) > MAX_BYTES:
raise SubmissionError("Submission exceeds the byte limit or is not bytes.")
predictions = {}
lines = content.split(b"\n")
if lines[-1] == b"":
lines.pop()
try:
for line in lines:
if line.endswith(b"\r"):
line = line[:-1]
if len(line) > MAX_ROW_BYTES:
raise SubmissionError("Submission row exceeds the byte limit.")
row = _decode(line)
key = _key(row)
if set(row) != {"topic", "query_id", "retrieved_ids"}:
raise ValueError
if key not in queries or key in predictions or not _ids(row["retrieved_ids"], galleries[key[0]], 20):
raise ValueError
predictions[key] = row["retrieved_ids"]
if predictions.keys() != queries.keys():
raise ValueError
except SubmissionError:
raise
except (ValueError, TypeError, KeyError, RecursionError):
raise SubmissionError("Invalid prediction JSONL or query/candidate coverage.") from None
return predictions
def _knowledge_mapping(path: Path, topics: set[str]) -> dict[str, str]:
try:
data = json.loads(path.read_text(encoding="utf-8"), object_pairs_hook=_object,
parse_constant=_constant)
if not isinstance(data, dict) or set(data) != {
"Visual Cognition", "Natural Science", "Formal Science",
"Humanities & Social Science", "Engineering & Technology"}:
raise ValueError
mapping = {}
for domain, members in data.items():
if not isinstance(members, list) or not members or any(not isinstance(x, str) for x in members):
raise ValueError
for topic in members:
if topic in mapping:
raise ValueError
mapping[topic] = domain
if set(mapping) != topics:
raise ValueError
return mapping
except (OSError, ValueError, TypeError, json.JSONDecodeError):
raise ValueError(_CONFIG_ERROR) from None
def load_labels(path: Path, manifest: dict, knowledge_path: Path | None = None) -> LabelSet:
"""Read private targets and annotations; v1.0 requires one target per query."""
queries, galleries = _manifest(manifest)
labels = {}
annotations = {}
knowledge_by_topic = _knowledge_mapping(
knowledge_path or path.with_name("knowledge_domains.json"), set(galleries))
if any(manifest["topics"][topic]["knowledge_label"] not in KNOWLEDGE_GROUPS[domain]
for topic, domain in knowledge_by_topic.items()):
raise ValueError(_CONFIG_ERROR)
try:
with path.open("rb") as handle:
for line in handle:
row = _decode(line)
key = _key(row)
if key not in queries or key in labels or not _ids(row["target"], galleries[key[0]], 1):
raise ValueError
demand = DEMAND_LABELS.get(frozenset(row["reasoning_demand"]))
skills = row["reasoning_skill"]
if (demand is None or not isinstance(skills, list)
or any(skill not in SKILL_LABELS for skill in skills)
or len(set(skills)) != len(skills)):
raise ValueError
labels[key] = row["target"]
annotations[key] = {
"reasoning_skill": [SKILL_LABELS[skill] for skill in skills],
"cognitive_demand": demand,
}
if labels.keys() != queries.keys():
raise ValueError
except (OSError, ValueError, TypeError, KeyError, RecursionError):
raise ValueError(_CONFIG_ERROR) from None
if ({skill for row in annotations.values() for skill in row["reasoning_skill"]}
!= set(REASONING_LABELS)
or {row["cognitive_demand"] for row in annotations.values()} != set(COGNITIVE_LABELS)):
raise ValueError(_CONFIG_ERROR)
return LabelSet(labels, annotations, knowledge_by_topic)
def evaluate(predictions, labels, manifest, model: dict) -> dict:
"""Compute percentages without rounding; validate even direct API inputs."""
queries, galleries = _manifest(manifest)
try:
if (not isinstance(labels, LabelSet) or labels.keys() != queries.keys()
or labels.annotations.keys() != queries.keys()
or set(labels.knowledge_by_topic) != set(galleries)):
raise ValueError
for key, targets in labels.items():
if (not isinstance(key, tuple) or len(key) != 2 or type(key[1]) is not int or key[1] < 0
or not _ids(targets, galleries[key[0]], 1)):
raise ValueError
for annotation in labels.annotations.values():
skills = annotation["reasoning_skill"]
if (not isinstance(skills, list) or len(set(skills)) != len(skills)
or any(skill not in REASONING_LABELS for skill in skills)
or annotation["cognitive_demand"] not in COGNITIVE_LABELS):
raise ValueError
if ({skill for annotation in labels.annotations.values()
for skill in annotation["reasoning_skill"]} != set(REASONING_LABELS)
or {annotation["cognitive_demand"] for annotation in labels.annotations.values()}
!= set(COGNITIVE_LABELS)):
raise ValueError
except (ValueError, KeyError, TypeError, AttributeError):
raise ValueError(_CONFIG_ERROR) from None
try:
if not isinstance(predictions, dict) or predictions.keys() != queries.keys():
raise ValueError
for key, candidates in predictions.items():
if (not isinstance(key, tuple) or len(key) != 2 or type(key[1]) is not int or key[1] < 0
or not _ids(candidates, galleries[key[0]], 20)):
raise ValueError
except (ValueError, KeyError, TypeError):
raise SubmissionError("Invalid prediction query/candidate coverage.") from None
groups = {"knowledge_domain": KNOWLEDGE_LABELS, "reasoning_skill": REASONING_LABELS,
"cognitive_demand": COGNITIVE_LABELS}
results = {}
for family, metrics in METRICS.items():
buckets = {group: {label: {metric: [] for metric in metrics} for label in names}
for group, names in groups.items()}
for key in queries:
candidates = predictions[key]
target = labels[key][0]
rank = candidates.index(target) + 1 if target in candidates else 21
annotation = labels.annotations[key]
demand = annotation["cognitive_demand"]
# Dual-tagged queries belong to both inclusive groups and their intersection.
cognitive_membership = ["KI", "RI", "KRI"] if demand == "KRI" else [demand]
membership = {"knowledge_domain": [manifest["topics"][key[0]]["knowledge_label"]],
"reasoning_skill": annotation["reasoning_skill"],
"cognitive_demand": cognitive_membership}
for metric in metrics:
cutoff = int(metric.split("@")[1])
value = (100.0 if family == "Recall" else 100.0 / math.log2(rank + 1)) if rank <= cutoff else 0.0
for group, names in membership.items():
for name in names:
buckets[group][name][metric].append(value)
results[family] = {
group: {name: {metric: math.fsum(values) / len(values) for metric, values in scores.items()}
for name, scores in members.items()}
for group, members in buckets.items()
}
for group in ("knowledge_domain", "reasoning_skill"):
results[family][group]["macro_avg"] = {
metric: math.fsum(results[family][group][name][metric] for name in groups[group]) / len(groups[group])
for metric in metrics
}
score = {"model": deepcopy(model), "meta_data": {
"date": datetime.now(timezone.utc).date().isoformat(),
"benchmark_version": ACTIVE_BENCHMARK_VERSION,
"dataset_revision": manifest["dataset_revision"], "evaluation_protocol": PROTOCOL,
}, "results": results}
if validate_score_data(score, expected_benchmark_version=ACTIVE_BENCHMARK_VERSION):
raise ValueError("Invalid score metadata or evaluation configuration.")
return score
def summary(score) -> dict:
"""Return the leaderboard's equal-weight knowledge/reasoning overview."""
metrics = []
unavailable = []
for family, names in METRICS.items():
results = score["results"][family]
for group, labels in (("knowledge_domain", KNOWLEDGE_LABELS),
("reasoning_skill", REASONING_LABELS), ("cognitive_demand", COGNITIVE_LABELS)):
for label in labels:
if any(results.get(group, {}).get(label, {}).get(metric) is None for metric in names):
item = {"family": family, "group": group, "label": label}
unavailable.append(item)
for metric in names:
knowledge = results["knowledge_domain"]["macro_avg"][metric]
reasoning = results["reasoning_skill"]["macro_avg"][metric]
metrics.append({"family": family, "metric": metric, "knowledge_macro": knowledge,
"reasoning_macro": reasoning,
"overview": (knowledge + reasoning) / 2 if knowledge is not None and reasoning is not None else None})
return {"metrics": metrics, "unavailable_groups": unavailable}