diff --git a/judgearena/benchmarks/pairwise/baselines.py b/judgearena/benchmarks/pairwise/baselines.py index d2b2ae1..309cded 100644 --- a/judgearena/benchmarks/pairwise/baselines.py +++ b/judgearena/benchmarks/pairwise/baselines.py @@ -4,7 +4,6 @@ from collections.abc import Mapping -from judgearena.datasets.arena_hard import ARENA_HARD_BASELINES from judgearena.datasets.m_arenahard import ( M_ARENA_HARD_BASELINES, split_m_arena_hard_dataset, @@ -14,7 +13,6 @@ from judgearena.tasks.schema import CategoryDefaultsBaseline, TaskDefaultBaseline LEGACY_PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = { - **ARENA_HARD_BASELINES, **M_ARENA_HARD_BASELINES, **MT_BENCH_BASELINES, } diff --git a/judgearena/benchmarks/pairwise/runner.py b/judgearena/benchmarks/pairwise/runner.py index 6b41d30..7330337 100644 --- a/judgearena/benchmarks/pairwise/runner.py +++ b/judgearena/benchmarks/pairwise/runner.py @@ -16,10 +16,6 @@ from judgearena.benchmarks.mt_bench.mt_bench_utils import run_mt_bench from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline from judgearena.datasets import load_instructions -from judgearena.datasets.arena_hard import ( - download_arena_hard, - is_arena_hard_dataset, -) from judgearena.datasets.fluency import is_fluency_task as task_is_fluency from judgearena.datasets.fluency import load_fluency_contexts from judgearena.evaluate import judge_and_parse_prefs, resolve_run_judge_prompt @@ -58,17 +54,12 @@ def try_load_dataset_completions( local_path_tables = data_root / "tables" resolved_task = get_packaged_task(dataset) if resolved_task is not None: - from judgearena.datasets.judgearena_tables import load_task_model_outputs + from judgearena.datasets.registry import resolve_dataset_adapter - df_outputs = load_task_model_outputs(resolved_task, local_path_tables) + adapter = resolve_dataset_adapter(resolved_task.spec.dataset.adapter) + df_outputs = adapter.load_model_outputs(resolved_task, local_path_tables) if df_outputs is None: return None - elif is_arena_hard_dataset(dataset): - download_arena_hard(dataset=dataset, local_tables_path=local_path_tables) - output_path = local_path_tables / "model_outputs" / f"{dataset}.csv.zip" - if not output_path.exists(): - return None - df_outputs = read_df(output_path) else: download_hf(name=dataset, local_path=local_path_tables) output_path = local_path_tables / "model_outputs" / f"{dataset}.csv.zip" diff --git a/judgearena/datasets/__init__.py b/judgearena/datasets/__init__.py index 5c0a13f..39d8bb0 100644 --- a/judgearena/datasets/__init__.py +++ b/judgearena/datasets/__init__.py @@ -1,9 +1,5 @@ import pandas as pd -from judgearena.datasets.arena_hard import ( - download_arena_hard, - is_arena_hard_dataset, -) from judgearena.datasets.m_arenahard import ( load_m_arenahard, split_m_arena_hard_dataset, @@ -18,9 +14,10 @@ def load_instructions(dataset: str, n_instructions: int | None = None) -> pd.Dat resolved_task = get_packaged_task(dataset) if resolved_task is not None: from judgearena import utils as judgearena_utils - from judgearena.datasets.judgearena_tables import load_task_instructions + from judgearena.datasets.registry import resolve_dataset_adapter - df_instructions = load_task_instructions( + adapter = resolve_dataset_adapter(resolved_task.spec.dataset.adapter) + df_instructions = adapter.load_instructions( resolved_task, judgearena_utils.data_root / "tables" ) @@ -55,20 +52,7 @@ def load_instructions(dataset: str, n_instructions: int | None = None) -> pd.Dat ) else: - assert dataset in [ - "arena-hard-v0.1", - "arena-hard-v2.0", - ] - from judgearena import utils as judgearena_utils - - local_path_tables = judgearena_utils.data_root / "tables" - if is_arena_hard_dataset(dataset): - download_arena_hard(dataset=dataset, local_tables_path=local_path_tables) - else: - judgearena_utils.download_hf(name=dataset, local_path=local_path_tables) - df_instructions = judgearena_utils.read_df( - local_path_tables / "instructions" / f"{dataset}.csv" - ) + raise ValueError(f"Unsupported instruction dataset {dataset!r}.") df_instructions = df_instructions.set_index("instruction_index").sort_index() logger.info("Loaded %d instructions for %s.", len(df_instructions), dataset) diff --git a/judgearena/datasets/arena_hard.py b/judgearena/datasets/arena_hard.py index f00ed4a..6005004 100644 --- a/judgearena/datasets/arena_hard.py +++ b/judgearena/datasets/arena_hard.py @@ -5,91 +5,70 @@ import pandas as pd from huggingface_hub import snapshot_download -from judgearena.dataset_revisions import hf_revision - -ARENA_HARD_HF_REPO_ID = "lmarena-ai/arena-hard-auto" - -# Mirrors upstream's `JUDGE_SETTINGS` baseline assignment in -# `arena-hard-auto/utils/judge_utils.py` verbatim: v0.1 has a single flat -# baseline, v2.0 routes per question category. `is_arena_hard_dataset` and -# the dispatcher in `generate_and_evaluate.py` key off this map. -# -# Note: the released v2.0 `question.jsonl` only tags rows as `hard_prompt` -# (500) or `creative_writing` (250); `coding` and `math` are inert keys -# upstream ships for forward compatibility (no question carries those -# labels, so the dispatcher never looks them up). We keep them so any -# future re-tagging upstream lights up automatically without a code -# change here. -ARENA_HARD_BASELINES: dict[str, str | Mapping[str, str]] = { - "arena-hard-v0.1": "gpt-4-0314", - "arena-hard-v2.0": { - "hard_prompt": "o3-mini-2025-01-31", - "coding": "o3-mini-2025-01-31", - "math": "o3-mini-2025-01-31", - "creative_writing": "gemini-2.0-flash-001", - }, -} - -# Dataset name -> upstream HF `data//` directory. Kept private so the -# public API of this module is just the baseline map and helpers below. -_ARENA_HARD_HF_VARIANTS: dict[str, str] = { - "arena-hard-v0.1": "arena-hard-v0.1", - "arena-hard-v2.0": "arena-hard-v2.0", -} +from judgearena.tasks.registry import get_packaged_task +from judgearena.tasks.schema import HuggingFaceDatasetSource, ResolvedTaskSpec def is_arena_hard_dataset(dataset: str) -> bool: - return dataset in ARENA_HARD_BASELINES + task = get_packaged_task(dataset) + return task is not None and task.spec.dataset.adapter == "arena_hard" def arena_hard_native_baseline( dataset: str, ) -> str | Mapping[str, str] | None: - """Dataset-native baseline assignment. + """Return the YAML-defined baseline for an Arena-Hard task.""" + if not is_arena_hard_dataset(dataset): + return None + from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline - Returns a plain string for flat datasets (v0.1), a `{category: model}` - mapping for per-category datasets (v2.0), or `None` for datasets that - don't ship a native baseline. - """ - return ARENA_HARD_BASELINES.get(dataset) + return native_pairwise_baseline(dataset) def normalize_official_arena_hard( raw_df: pd.DataFrame, dataset: str ) -> tuple[pd.DataFrame, pd.DataFrame | None]: - if dataset not in _ARENA_HARD_HF_VARIANTS: + if not is_arena_hard_dataset(dataset): raise ValueError(f"Unsupported Arena-Hard dataset: {dataset}") df_instructions = _build_instructions(raw_df) df_model_outputs = _build_model_outputs(raw_df) return df_instructions, df_model_outputs -def download_arena_hard(dataset: str, local_tables_path: Path) -> None: - """Populate `{dataset}.csv` and `{dataset}.csv.zip` on disk if missing. +def _source(task: ResolvedTaskSpec) -> HuggingFaceDatasetSource: + source = task.spec.dataset.sources.get("examples") + if not isinstance(source, HuggingFaceDatasetSource) or source.config is None: + raise ValueError( + f"Task {task.task!r} must define an 'examples' Hugging Face source " + "with its Arena-Hard data variant in 'config'." + ) + return source + + +def download_task_sources(task: ResolvedTaskSpec, local_tables_path: Path) -> None: + """Populate canonical instruction and output tables for one task. Pulls the raw jsonl files directly via `snapshot_download` and reads them with pandas: upstream's per-row `messages[].content` oscillates between string and dict across answer files, so `datasets.load_dataset` can't materialize them into a single Arrow schema. - """ - if dataset not in _ARENA_HARD_HF_VARIANTS: - return + if task.spec.dataset.adapter != "arena_hard": + raise ValueError(f"Task {task.task!r} does not use the Arena-Hard adapter.") + dataset = task.task instructions_path = local_tables_path / "instructions" / f"{dataset}.csv" model_outputs_path = local_tables_path / "model_outputs" / f"{dataset}.csv.zip" if instructions_path.exists() and model_outputs_path.exists(): return - variant = _ARENA_HARD_HF_VARIANTS[dataset] + source = _source(task) + variant = source.config snapshot_root = snapshot_download( - repo_id=ARENA_HARD_HF_REPO_ID, + repo_id=source.repo_id, repo_type="dataset", - allow_patterns=[ - f"data/{variant}/question.jsonl", - f"data/{variant}/model_answer/*.jsonl", - ], + allow_patterns=list(source.allow_patterns) or None, force_download=False, - revision=hf_revision(ARENA_HARD_HF_REPO_ID), + revision=source.revision, ) raw_df = _read_arena_hard_jsonl_frames( variant_dir=Path(snapshot_root) / "data" / variant @@ -104,6 +83,30 @@ def download_arena_hard(dataset: str, local_tables_path: Path) -> None: df_model_outputs.to_csv(model_outputs_path, index=False) +def load_task_instructions( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame: + """Load normalized instructions for a registered Arena-Hard task.""" + download_task_sources(task, local_tables_path) + return pd.read_csv(local_tables_path / "instructions" / f"{task.task}.csv") + + +def load_task_model_outputs( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame | None: + """Load normalized reference outputs for a registered Arena-Hard task.""" + download_task_sources(task, local_tables_path) + path = local_tables_path / "model_outputs" / f"{task.task}.csv.zip" + return pd.read_csv(path) if path.exists() else None + + +def download_arena_hard(dataset: str, local_tables_path: Path) -> None: + """Compatibility wrapper around the registered Arena-Hard adapter.""" + task = get_packaged_task(dataset) + if task is not None and task.spec.dataset.adapter == "arena_hard": + download_task_sources(task, local_tables_path) + + def _read_arena_hard_jsonl_frames(variant_dir: Path) -> pd.DataFrame: frames: list[pd.DataFrame] = [] question_path = variant_dir / "question.jsonl" diff --git a/judgearena/datasets/registry.py b/judgearena/datasets/registry.py new file mode 100644 index 0000000..df2c76f --- /dev/null +++ b/judgearena/datasets/registry.py @@ -0,0 +1,50 @@ +"""Registry connecting task dataset-adapter IDs to implementations.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path + +import pandas as pd + +from judgearena.tasks.schema import ResolvedTaskSpec + +TaskDataFunction = Callable[[ResolvedTaskSpec, Path], pd.DataFrame | None] +TaskDownloadFunction = Callable[[ResolvedTaskSpec, Path], None] + + +@dataclass(frozen=True) +class DatasetAdapter: + name: str + download: TaskDownloadFunction + load_instructions: TaskDataFunction + load_model_outputs: TaskDataFunction + + +def dataset_adapters() -> tuple[DatasetAdapter, ...]: + """Return registered dataset implementations.""" + from judgearena.datasets import arena_hard, judgearena_tables + + return ( + DatasetAdapter( + "judgearena_tables", + judgearena_tables.download_task_sources, + judgearena_tables.load_task_instructions, + judgearena_tables.load_task_model_outputs, + ), + DatasetAdapter( + "arena_hard", + arena_hard.download_task_sources, + arena_hard.load_task_instructions, + arena_hard.load_task_model_outputs, + ), + ) + + +def resolve_dataset_adapter(name: str) -> DatasetAdapter: + """Return the implementation registered under ``name``.""" + for adapter in dataset_adapters(): + if adapter.name == name: + return adapter + raise ValueError(f"Unknown task dataset adapter {name!r}.") diff --git a/judgearena/prompts/registry.py b/judgearena/prompts/registry.py index d7f184b..f60ffa6 100644 --- a/judgearena/prompts/registry.py +++ b/judgearena/prompts/registry.py @@ -93,8 +93,6 @@ def metadata(self) -> dict[str, str | bool | None]: JUDGE_PROMPT_PRESETS = tuple(PRESETS) TASK_DEFAULT_PRESET: dict[str, str] = { - "arena-hard-v0.1": DEFAULT_JUDGE_PROMPT_PRESET, - "arena-hard-v2.0": DEFAULT_JUDGE_PROMPT_PRESET, "mt-bench": FASTCHAT_PAIRWISE_PROMPT_PRESET, } diff --git a/judgearena/tasks/cli.py b/judgearena/tasks/cli.py index 9dcba9f..6abf2c4 100644 --- a/judgearena/tasks/cli.py +++ b/judgearena/tasks/cli.py @@ -8,8 +8,7 @@ import yaml -from judgearena.tasks.loader import TaskDefinitionError -from judgearena.tasks.registry import load_tasks +from judgearena.tasks.registry import TaskDefinitionError, load_tasks from judgearena.tasks.schema import ResolvedTaskSpec diff --git a/judgearena/tasks/definitions/arena_hard/_base.yaml b/judgearena/tasks/definitions/arena_hard/_base.yaml new file mode 100644 index 0000000..e4b2810 --- /dev/null +++ b/judgearena/tasks/definitions/arena_hard/_base.yaml @@ -0,0 +1,31 @@ +schema_version: 1 +task_version: 1 +tags: [pairwise, instruction-following, arena-hard] + +dataset: + adapter: arena_hard + sources: + examples: + type: huggingface_dataset + repo_id: lmarena-ai/arena-hard-auto + revision: "15f3746e21432264ce9b453999bde4f3c946d2e6" + fields: + id: instruction_index + instruction: instruction + +protocol: + runner: pairwise + generation: + mode: single_turn_chat + judge: + default_prompt: default + parser: pairwise_preference + default_swap_mode: fixed + allowed_swap_modes: [fixed, both] + scoring: + adapter: pairwise_win_rate + primary_metric: winrate + higher_is_better: true + +metadata: + reference_implementation: https://github.com/lmarena/arena-hard-auto diff --git a/judgearena/tasks/definitions/arena_hard/arena-hard-v0.1.yaml b/judgearena/tasks/definitions/arena_hard/arena-hard-v0.1.yaml new file mode 100644 index 0000000..6edb003 --- /dev/null +++ b/judgearena/tasks/definitions/arena_hard/arena-hard-v0.1.yaml @@ -0,0 +1,17 @@ +extends: _base.yaml +task: arena-hard-v0.1 +description: Pairwise evaluation on Arena-Hard v0.1. + +dataset: + sources: + examples: + config: arena-hard-v0.1 + allow_patterns: + - data/arena-hard-v0.1/question.jsonl + - data/arena-hard-v0.1/model_answer/*.jsonl + +protocol: + baseline: + strategy: task_default + reference_id: gpt-4-0314 + allow_runtime_override: true diff --git a/judgearena/tasks/definitions/arena_hard/arena-hard-v2.0.yaml b/judgearena/tasks/definitions/arena_hard/arena-hard-v2.0.yaml new file mode 100644 index 0000000..5fe3410 --- /dev/null +++ b/judgearena/tasks/definitions/arena_hard/arena-hard-v2.0.yaml @@ -0,0 +1,24 @@ +extends: _base.yaml +task: arena-hard-v2.0 +description: Pairwise evaluation on Arena-Hard v2.0. + +dataset: + sources: + examples: + config: arena-hard-v2.0 + allow_patterns: + - data/arena-hard-v2.0/question.jsonl + - data/arena-hard-v2.0/model_answer/*.jsonl + fields: + category: category + +protocol: + baseline: + strategy: category_defaults + category_field: category + references: + hard_prompt: o3-mini-2025-01-31 + coding: o3-mini-2025-01-31 + math: o3-mini-2025-01-31 + creative_writing: gemini-2.0-flash-001 + allow_runtime_override: true diff --git a/judgearena/tasks/loader.py b/judgearena/tasks/loader.py deleted file mode 100644 index 2a1e267..0000000 --- a/judgearena/tasks/loader.py +++ /dev/null @@ -1,231 +0,0 @@ -"""Read task YAML, resolve private base files, and record stable hashes.""" - -from __future__ import annotations - -import hashlib -import json -import posixpath -from copy import deepcopy -from importlib.resources.abc import Traversable -from pathlib import Path, PurePosixPath -from typing import Any - -import yaml -from pydantic import ValidationError - -from judgearena.tasks.schema import ( - ResolvedTaskSpec, - ResourceDigest, - TaskProvenance, - TaskSpec, -) - - -class TaskDefinitionError(ValueError): - """A packaged task definition is malformed or unsafe.""" - - -class _UniqueKeySafeLoader(yaml.SafeLoader): - """Safe YAML loader that rejects duplicate mapping keys.""" - - pass - - -def _construct_unique_mapping( - loader: _UniqueKeySafeLoader, node: yaml.MappingNode, deep: bool = False -) -> dict[object, object]: - mapping: dict[object, object] = {} - for key_node, value_node in node.value: - key = loader.construct_object(key_node, deep=deep) - try: - duplicate = key in mapping - except TypeError as exc: - raise yaml.constructor.ConstructorError( - "while constructing a mapping", - node.start_mark, - "found an unhashable mapping key", - key_node.start_mark, - ) from exc - if duplicate: - raise yaml.constructor.ConstructorError( - "while constructing a mapping", - node.start_mark, - f"found duplicate key {key!r}", - key_node.start_mark, - ) - mapping[key] = loader.construct_object(value_node, deep=deep) - return mapping - - -_UniqueKeySafeLoader.add_constructor( - yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, - _construct_unique_mapping, -) - - -def _sha256(text: str) -> str: - return hashlib.sha256(text.encode("utf-8")).hexdigest() - - -def _canonical_sha256(data: dict[str, object]) -> str: - canonical = json.dumps( - data, - allow_nan=False, - ensure_ascii=False, - separators=(",", ":"), - sort_keys=True, - ) - return _sha256(canonical) - - -def _strict_yaml_mapping(text: str, *, path: str) -> dict[str, Any]: - try: - data = yaml.load(text, Loader=_UniqueKeySafeLoader) - except yaml.YAMLError as exc: - raise TaskDefinitionError(f"{path}: invalid YAML: {exc}") from exc - if not isinstance(data, dict): - raise TaskDefinitionError(f"{path}: task YAML must contain a mapping") - if any(not isinstance(key, str) for key in data): - raise TaskDefinitionError(f"{path}: task YAML keys must be strings") - return data - - -def _merge_mapping(parent: dict[str, Any], child: dict[str, Any]) -> dict[str, Any]: - """Apply the documented recursive-map/replace-list/null-delete merge.""" - merged = deepcopy(parent) - for key, value in child.items(): - if value is None: - merged.pop(key, None) - elif isinstance(value, dict) and isinstance(merged.get(key), dict): - merged[key] = _merge_mapping(merged[key], value) - else: - merged[key] = deepcopy(value) - return merged - - -class TaskLoader: - """Load and normalize task YAML from a filesystem or installed package.""" - - def __init__(self, root: Traversable): - self.root = root - - def discover(self) -> tuple[str, ...]: - """Return public task YAML paths; ``_*.yaml`` files are private bases.""" - discovered: list[str] = [] - - def walk(directory: Traversable, prefix: PurePosixPath) -> None: - for entry in sorted(directory.iterdir(), key=lambda item: item.name): - relative = prefix / entry.name - if entry.is_dir(): - walk(entry, relative) - elif ( - entry.is_file() - and entry.name.endswith((".yaml", ".yml")) - and not entry.name.startswith("_") - ): - discovered.append(relative.as_posix()) - - if not self.root.is_dir(): - raise TaskDefinitionError("task definitions root is not a directory") - walk(self.root, PurePosixPath()) - return tuple(discovered) - - def load(self, relative_path: str) -> ResolvedTaskSpec: - """Resolve, validate, and fingerprint one public task definition.""" - relative_path = self._normalize_root_path(relative_path) - if PurePosixPath(relative_path).name.startswith("_"): - raise TaskDefinitionError( - f"{relative_path}: private base files are not runnable tasks" - ) - resolved, resources = self._resolve(relative_path, chain=()) - if "task" not in resolved: - raise TaskDefinitionError(f"{relative_path}: public task must define task") - try: - spec = TaskSpec.model_validate(resolved) - except ValidationError as exc: - raise TaskDefinitionError(f"{relative_path}: {exc}") from exc - normalized = spec.model_dump(mode="json") - child_digest = next( - digest for digest in resources if digest.path == relative_path - ) - return ResolvedTaskSpec( - spec=spec, - provenance=TaskProvenance( - source_path=relative_path, - source_sha256=child_digest.sha256, - resolved_sha256=_canonical_sha256(normalized), - resources=resources, - ), - ) - - def _resolve( - self, relative_path: str, *, chain: tuple[str, ...] - ) -> tuple[dict[str, Any], tuple[ResourceDigest, ...]]: - normalized_path = self._normalize_root_path(relative_path) - if normalized_path in chain: - cycle = " -> ".join((*chain, normalized_path)) - raise TaskDefinitionError(f"task inheritance cycle: {cycle}") - - resource = self._resource(normalized_path) - if not resource.is_file(): - raise TaskDefinitionError(f"{normalized_path}: task file does not exist") - text = resource.read_text(encoding="utf-8") - data = _strict_yaml_mapping(text, path=normalized_path) - digest = ResourceDigest(normalized_path, _sha256(text)) - - is_base = PurePosixPath(normalized_path).name.startswith("_") - if is_base and "task" in data: - raise TaskDefinitionError( - f"{normalized_path}: private base files must not define task" - ) - - extends = data.pop("extends", None) - if extends is None: - return data, (digest,) - if not isinstance(extends, str) or not extends: - raise TaskDefinitionError( - f"{normalized_path}: extends must be one relative YAML path" - ) - base_path = self._resolve_extends(normalized_path, extends) - base, base_resources = self._resolve(base_path, chain=(*chain, normalized_path)) - return _merge_mapping(base, data), (*base_resources, digest) - - def _resolve_extends(self, child_path: str, extends: str) -> str: - requested = PurePosixPath(extends) - if requested.is_absolute() or requested.suffix not in {".yaml", ".yml"}: - raise TaskDefinitionError( - f"{child_path}: extends must reference a relative YAML file" - ) - if not requested.name.startswith("_"): - raise TaskDefinitionError( - f"{child_path}: extends may only reference a private _*.yaml base" - ) - joined = posixpath.normpath(str(PurePosixPath(child_path).parent / requested)) - return self._normalize_root_path(joined) - - def _normalize_root_path(self, relative_path: str) -> str: - pure = PurePosixPath(relative_path) - normalized = posixpath.normpath(pure.as_posix()) - if ( - pure.is_absolute() - or normalized in {"", ".", ".."} - or normalized.startswith("../") - ): - raise TaskDefinitionError( - f"{relative_path}: path escapes task definitions root" - ) - return PurePosixPath(normalized).as_posix() - - def _resource(self, relative_path: str) -> Traversable: - normalized = self._normalize_root_path(relative_path) - resource: Traversable = self.root - for part in PurePosixPath(normalized).parts: - resource = resource.joinpath(part) - if isinstance(self.root, Path) and isinstance(resource, Path): - root_path = self.root.resolve() - resource_path = resource.resolve() - if not resource_path.is_relative_to(root_path): - raise TaskDefinitionError( - f"{relative_path}: path escapes task definitions root" - ) - return resource diff --git a/judgearena/tasks/registry.py b/judgearena/tasks/registry.py index 2159e50..a70fd17 100644 --- a/judgearena/tasks/registry.py +++ b/judgearena/tasks/registry.py @@ -1,15 +1,242 @@ -"""Discover task definitions and validate their referenced component IDs.""" +"""Load packaged task definitions and validate their referenced component IDs.""" from __future__ import annotations +import hashlib +import json +import posixpath +from copy import deepcopy from dataclasses import dataclass from functools import cache from importlib.resources import files from importlib.resources.abc import Traversable +from pathlib import Path, PurePosixPath +from typing import Any +import yaml +from pydantic import ValidationError + +from judgearena.log import get_logger from judgearena.prompts.registry import JUDGE_PROMPT_PRESETS -from judgearena.tasks.loader import TaskDefinitionError, TaskLoader -from judgearena.tasks.schema import ResolvedTaskSpec +from judgearena.tasks.schema import ( + ResolvedTaskSpec, + ResourceDigest, + TaskProvenance, + TaskSpec, +) + +logger = get_logger(__name__) + + +class TaskDefinitionError(ValueError): + """A packaged task definition is malformed or unsafe.""" + + +class _UniqueKeySafeLoader(yaml.SafeLoader): + """Safe YAML loader that rejects duplicate mapping keys.""" + + pass + + +def _construct_unique_mapping( + loader: _UniqueKeySafeLoader, node: yaml.MappingNode, deep: bool = False +) -> dict[object, object]: + mapping: dict[object, object] = {} + for key_node, value_node in node.value: + key = loader.construct_object(key_node, deep=deep) + try: + duplicate = key in mapping + except TypeError as exc: + raise yaml.constructor.ConstructorError( + "while constructing a mapping", + node.start_mark, + "found an unhashable mapping key", + key_node.start_mark, + ) from exc + if duplicate: + raise yaml.constructor.ConstructorError( + "while constructing a mapping", + node.start_mark, + f"found duplicate key {key!r}", + key_node.start_mark, + ) + mapping[key] = loader.construct_object(value_node, deep=deep) + return mapping + + +_UniqueKeySafeLoader.add_constructor( + yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, + _construct_unique_mapping, +) + + +def _sha256(text: str) -> str: + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def _canonical_sha256(data: dict[str, object]) -> str: + canonical = json.dumps( + data, + allow_nan=False, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + ) + return _sha256(canonical) + + +def _strict_yaml_mapping(text: str, *, path: str) -> dict[str, Any]: + try: + data = yaml.load(text, Loader=_UniqueKeySafeLoader) + except yaml.YAMLError as exc: + raise TaskDefinitionError(f"{path}: invalid YAML: {exc}") from exc + if not isinstance(data, dict): + raise TaskDefinitionError(f"{path}: task YAML must contain a mapping") + if any(not isinstance(key, str) for key in data): + raise TaskDefinitionError(f"{path}: task YAML keys must be strings") + return data + + +def _merge_mapping(parent: dict[str, Any], child: dict[str, Any]) -> dict[str, Any]: + """Apply the documented recursive-map/replace-list/null-delete merge.""" + merged = deepcopy(parent) + for key, value in child.items(): + if value is None: + merged.pop(key, None) + elif isinstance(value, dict) and isinstance(merged.get(key), dict): + merged[key] = _merge_mapping(merged[key], value) + else: + merged[key] = deepcopy(value) + return merged + + +def _normalize_root_path(relative_path: str) -> str: + pure = PurePosixPath(relative_path) + normalized = posixpath.normpath(pure.as_posix()) + if ( + pure.is_absolute() + or normalized in {"", ".", ".."} + or normalized.startswith("../") + ): + raise TaskDefinitionError( + f"{relative_path}: path escapes task definitions root" + ) + return PurePosixPath(normalized).as_posix() + + +def _resource(root: Traversable, relative_path: str) -> Traversable: + normalized = _normalize_root_path(relative_path) + resource: Traversable = root + for part in PurePosixPath(normalized).parts: + resource = resource.joinpath(part) + if isinstance(root, Path) and isinstance(resource, Path): + root_path = root.resolve() + resource_path = resource.resolve() + if not resource_path.is_relative_to(root_path): + raise TaskDefinitionError( + f"{relative_path}: path escapes task definitions root" + ) + return resource + + +def _discover(root: Traversable) -> tuple[str, ...]: + """Public task YAMLs at exactly ``/.yaml``. + + ``_*.yaml`` files are private bases for ``extends`` and are not runnable. + """ + if not root.is_dir(): + raise TaskDefinitionError("task definitions root is not a directory") + discovered: list[str] = [] + for family in sorted(root.iterdir(), key=lambda entry: entry.name): + if not family.is_dir(): + if family.name.endswith((".yaml", ".yml")): + logger.warning( + "Ignoring %s: task YAML must live in a / subfolder.", + family.name, + ) + continue + for entry in sorted(family.iterdir(), key=lambda item: item.name): + if ( + entry.is_file() + and entry.name.endswith((".yaml", ".yml")) + and not entry.name.startswith("_") + ): + discovered.append(f"{family.name}/{entry.name}") + return tuple(discovered) + + +def _resolve( + root: Traversable, relative_path: str, *, chain: tuple[str, ...] +) -> tuple[dict[str, Any], tuple[ResourceDigest, ...]]: + normalized_path = _normalize_root_path(relative_path) + if normalized_path in chain: + cycle = " -> ".join((*chain, normalized_path)) + raise TaskDefinitionError(f"task inheritance cycle: {cycle}") + + resource = _resource(root, normalized_path) + if not resource.is_file(): + raise TaskDefinitionError(f"{normalized_path}: task file does not exist") + text = resource.read_text(encoding="utf-8") + data = _strict_yaml_mapping(text, path=normalized_path) + digest = ResourceDigest(normalized_path, _sha256(text)) + + is_base = PurePosixPath(normalized_path).name.startswith("_") + if is_base and "task" in data: + raise TaskDefinitionError( + f"{normalized_path}: private base files must not define task" + ) + + extends = data.pop("extends", None) + if extends is None: + return data, (digest,) + if not isinstance(extends, str) or not extends: + raise TaskDefinitionError( + f"{normalized_path}: extends must be one relative YAML path" + ) + base_path = _resolve_extends(normalized_path, extends) + base, base_resources = _resolve(root, base_path, chain=(*chain, normalized_path)) + return _merge_mapping(base, data), (*base_resources, digest) + + +def _resolve_extends(child_path: str, extends: str) -> str: + requested = PurePosixPath(extends) + if requested.is_absolute() or requested.suffix not in {".yaml", ".yml"}: + raise TaskDefinitionError( + f"{child_path}: extends must reference a relative YAML file" + ) + if not requested.name.startswith("_"): + raise TaskDefinitionError( + f"{child_path}: extends may only reference a private _*.yaml base" + ) + joined = posixpath.normpath(str(PurePosixPath(child_path).parent / requested)) + return _normalize_root_path(joined) + + +def _load_task(root: Traversable, relative_path: str) -> ResolvedTaskSpec: + """Resolve, validate, and fingerprint one public task definition.""" + relative_path = _normalize_root_path(relative_path) + if PurePosixPath(relative_path).name.startswith("_"): + raise TaskDefinitionError( + f"{relative_path}: private base files are not runnable tasks" + ) + resolved, resources = _resolve(root, relative_path, chain=()) + if "task" not in resolved: + raise TaskDefinitionError(f"{relative_path}: public task must define task") + try: + spec = TaskSpec.model_validate(resolved) + except ValidationError as exc: + raise TaskDefinitionError(f"{relative_path}: {exc}") from exc + normalized = spec.model_dump(mode="json") + child_digest = next(digest for digest in resources if digest.path == relative_path) + return ResolvedTaskSpec( + spec=spec, + provenance=TaskProvenance( + source_path=relative_path, + source_sha256=child_digest.sha256, + resolved_sha256=_canonical_sha256(normalized), + resources=resources, + ), + ) @dataclass(frozen=True) @@ -17,7 +244,7 @@ class AdapterCatalog: """Component IDs that task YAML files may reference.""" runners: frozenset[str] = frozenset({"pairwise"}) - datasets: frozenset[str] = frozenset({"judgearena_tables"}) + datasets: frozenset[str] = frozenset({"arena_hard", "judgearena_tables"}) prompts: frozenset[str] = frozenset(JUDGE_PROMPT_PRESETS) parsers: frozenset[str] = frozenset({"pairwise_preference"}) scorers: frozenset[str] = frozenset({"pairwise_win_rate"}) @@ -48,10 +275,9 @@ def _load_packaged_tasks() -> dict[str, ResolvedTaskSpec]: def _discover_tasks( definitions_root: Traversable, adapters: AdapterCatalog ) -> dict[str, ResolvedTaskSpec]: - loader = TaskLoader(definitions_root) tasks: dict[str, ResolvedTaskSpec] = {} - for relative_path in loader.discover(): - resolved = loader.load(relative_path) + for relative_path in _discover(definitions_root): + resolved = _load_task(definitions_root, relative_path) if resolved.task in tasks: other = tasks[resolved.task].provenance.source_path raise TaskDefinitionError( diff --git a/judgearena/tasks/schema.py b/judgearena/tasks/schema.py index 2715ef8..146791c 100644 --- a/judgearena/tasks/schema.py +++ b/judgearena/tasks/schema.py @@ -171,6 +171,12 @@ def _validate_task(self) -> TaskSpec: f"official baseline source {baseline.source!r} is not declared " "in dataset.sources" ) + if isinstance(baseline, CategoryDefaultsBaseline) and ( + self.dataset.fields.category != baseline.category_field + ): + raise ValueError( + "category-default baseline must use dataset.fields.category" + ) return self diff --git a/judgearena/utils/io.py b/judgearena/utils/io.py index 3eb5a6c..85c8efc 100644 --- a/judgearena/utils/io.py +++ b/judgearena/utils/io.py @@ -11,10 +11,6 @@ import pandas as pd from huggingface_hub import snapshot_download -from judgearena.datasets.arena_hard import ( - download_arena_hard, - is_arena_hard_dataset, -) from judgearena.log import get_logger logger = get_logger(__name__) @@ -38,9 +34,11 @@ def download_hf(name: str, local_path: Path): resolved_task = get_packaged_task(name) if resolved_task is not None: - from judgearena.datasets.judgearena_tables import download_task_sources + from judgearena.datasets.registry import resolve_dataset_adapter - download_task_sources(resolved_task, local_path) + resolve_dataset_adapter(resolved_task.spec.dataset.adapter).download( + resolved_task, local_path + ) else: local_path.mkdir(exist_ok=True, parents=True) snapshot_download( @@ -85,21 +83,10 @@ def download_all(): logger.info("Downloading all datasets in %s", data_root) local_path_tables = data_root / "tables" - packaged_table_tasks = tuple( - task_id - for task_id, resolved in load_tasks().items() - if resolved.spec.dataset.adapter == "judgearena_tables" - ) - for dataset in ( - *packaged_table_tasks, - "arena-hard-v0.1", - "arena-hard-v2.0", - *M_ARENA_HARD_BASELINES, - ): - if is_arena_hard_dataset(dataset): - download_arena_hard(dataset=dataset, local_tables_path=local_path_tables) - else: - download_hf(name=dataset, local_path=local_path_tables) + for task_id in load_tasks(): + download_hf(name=task_id, local_path=local_path_tables) + for dataset in M_ARENA_HARD_BASELINES: + download_hf(name=dataset, local_path=local_path_tables) download_fluency_dataset(data_root) diff --git a/tests/test_generate_and_evaluate.py b/tests/test_generate_and_evaluate.py index 57f58b4..e579b96 100644 --- a/tests/test_generate_and_evaluate.py +++ b/tests/test_generate_and_evaluate.py @@ -138,6 +138,11 @@ def test_alpaca_eval_baseline_is_not_duplicated_in_legacy_registry(): assert native_pairwise_baseline("alpaca-eval") == "gpt4_1106_preview" +def test_arena_hard_baselines_are_not_duplicated_in_legacy_registry(): + assert "arena-hard-v0.1" not in LEGACY_PAIRWISE_BASELINES + assert "arena-hard-v2.0" not in LEGACY_PAIRWISE_BASELINES + + def test_resolve_plan_explicit_model_b_overrides_native(): plan = _resolve_baseline_plan( task="arena-hard-v2.0", diff --git a/tests/test_instruction_dataset.py b/tests/test_instruction_dataset.py index 142b657..8dd1627 100644 --- a/tests/test_instruction_dataset.py +++ b/tests/test_instruction_dataset.py @@ -5,10 +5,9 @@ import judgearena.benchmarks.pairwise.runner as generate_and_evaluate import judgearena.datasets as instruction_dataset +import judgearena.datasets.arena_hard as arena_hard import judgearena.datasets.judgearena_tables as judgearena_tables -import judgearena.utils as judgearena_utils from judgearena.datasets.arena_hard import ( - ARENA_HARD_BASELINES, _build_instructions, _build_model_outputs, _extract_assistant_output, @@ -69,20 +68,13 @@ def test_arena_hard_native_baseline_v20_is_per_category_mapping(): assert native["creative_writing"] == "gemini-2.0-flash-001" -def test_arena_hard_baselines_mapping_matches_upstream(): - """Pin the exact baseline assignment so a silent edit to - ARENA_HARD_BASELINES can't drift away from upstream - (arena-hard-auto/utils/judge_utils.py::JUDGE_SETTINGS). - """ - assert ARENA_HARD_BASELINES == { - "arena-hard-v0.1": "gpt-4-0314", - "arena-hard-v2.0": { - "hard_prompt": "o3-mini-2025-01-31", - "coding": "o3-mini-2025-01-31", - "math": "o3-mini-2025-01-31", - "creative_writing": "gemini-2.0-flash-001", - }, - } +def test_arena_hard_source_is_owned_by_task_yaml(): + task = get_packaged_task("arena-hard-v2.0") + assert task is not None + source = task.spec.dataset.sources["examples"] + assert source.repo_id == "lmarena-ai/arena-hard-auto" + assert source.revision == "15f3746e21432264ce9b453999bde4f3c946d2e6" + assert source.config == "arena-hard-v2.0" def test_mt_bench_native_baseline_is_flat_string(): @@ -281,12 +273,9 @@ def test_build_instructions_drops_model_answer_rows(): def test_load_instructions_uses_explicit_version_filename(monkeypatch): captured = {} - def _fake_ensure(dataset: str, local_tables_path: Path): - captured["dataset"] = dataset + def _fake_load(task, local_tables_path: Path): + captured["dataset"] = task.task captured["local_tables_path"] = local_tables_path - - def _fake_read_df(path: Path): - captured["path"] = path return pd.DataFrame( { "instruction_index": ["0", "1"], @@ -294,12 +283,10 @@ def _fake_read_df(path: Path): } ) - monkeypatch.setattr(instruction_dataset, "download_arena_hard", _fake_ensure) - monkeypatch.setattr(judgearena_utils, "read_df", _fake_read_df) + monkeypatch.setattr(arena_hard, "load_task_instructions", _fake_load) df = instruction_dataset.load_instructions(dataset="arena-hard-v2.0") assert captured["dataset"] == "arena-hard-v2.0" - assert captured["path"].name == "arena-hard-v2.0.csv" assert df.index.tolist() == ["0", "1"] @@ -309,14 +296,9 @@ def test_load_instructions_surfaces_category_for_v20(monkeypatch): from the cached CSV. """ monkeypatch.setattr( - instruction_dataset, - "download_arena_hard", - lambda dataset, local_tables_path: None, - ) - monkeypatch.setattr( - judgearena_utils, - "read_df", - lambda path: pd.DataFrame( + arena_hard, + "load_task_instructions", + lambda task, path: pd.DataFrame( { "instruction_index": ["q1", "q2"], "instruction": ["a", "b"], @@ -346,9 +328,9 @@ def test_try_load_dataset_completions_uses_dataset_output_file(monkeypatch, tmp_ monkeypatch.setattr(generate_and_evaluate, "data_root", tmp_path) monkeypatch.setattr( - generate_and_evaluate, - "download_arena_hard", - lambda dataset, local_tables_path: None, + arena_hard, + "load_task_model_outputs", + lambda task, path: pd.read_csv(output_path), ) loaded = generate_and_evaluate.try_load_dataset_completions( diff --git a/tests/test_mt_bench_downloads.py b/tests/test_mt_bench_downloads.py index 6ae0178..84c4246 100644 --- a/tests/test_mt_bench_downloads.py +++ b/tests/test_mt_bench_downloads.py @@ -43,7 +43,6 @@ def _snapshot_download_stub(**_kwargs): def test_download_all_includes_mt_bench(tmp_path, monkeypatch): hf_datasets = [] - arena_hard_datasets = [] calls = {"contexts": 0, "mt_bench": 0} monkeypatch.setattr(utils_io, "data_root", tmp_path) @@ -52,13 +51,6 @@ def test_download_all_includes_mt_bench(tmp_path, monkeypatch): "download_hf", lambda name, local_path: hf_datasets.append((name, local_path)), ) - monkeypatch.setattr( - utils_io, - "download_arena_hard", - lambda dataset, local_tables_path: arena_hard_datasets.append( - (dataset, local_tables_path) - ), - ) def _contexts_snapshot_stub(**_kwargs): calls["contexts"] += 1 @@ -75,13 +67,12 @@ def _contexts_snapshot_stub(**_kwargs): tables_dir = tmp_path / "tables" assert [name for name, _ in hf_datasets] == [ "alpaca-eval", + "arena-hard-v0.1", + "arena-hard-v2.0", "m-arena-hard-v0.1", "m-arena-hard-v2.0", ] - assert arena_hard_datasets == [ - ("arena-hard-v0.1", tables_dir), - ("arena-hard-v2.0", tables_dir), - ] + assert all(path == tables_dir for _, path in hf_datasets) assert calls["contexts"] == 1 assert calls["mt_bench"] == 1 diff --git a/tests/test_prompt_registry.py b/tests/test_prompt_registry.py index 414bbdc..a4df403 100644 --- a/tests/test_prompt_registry.py +++ b/tests/test_prompt_registry.py @@ -35,6 +35,12 @@ def test_alpaca_eval_prompt_default_is_not_duplicated_in_legacy_registry(): assert default_preset_for_task("alpaca-eval") == "default" +def test_arena_hard_prompt_defaults_are_not_duplicated_in_legacy_registry(): + assert "arena-hard-v0.1" not in TASK_DEFAULT_PRESET + assert "arena-hard-v2.0" not in TASK_DEFAULT_PRESET + assert default_preset_for_task("arena-hard-v2.0") == "default" + + def test_default_preset_for_fluency_prefix(): assert default_preset_for_task("fluency-french") == FLUENCY_JUDGE_PROMPT_PRESET assert default_preset_for_task("fluency-spanish") == FLUENCY_JUDGE_PROMPT_PRESET diff --git a/tests/test_task_registry.py b/tests/test_task_registry.py index 835b8e5..1e81fb7 100644 --- a/tests/test_task_registry.py +++ b/tests/test_task_registry.py @@ -9,8 +9,7 @@ from judgearena import cli as cli_module from judgearena.tasks.cli import run_task_command -from judgearena.tasks.loader import TaskDefinitionError -from judgearena.tasks.registry import load_tasks +from judgearena.tasks.registry import TaskDefinitionError, load_tasks def _task_definition(task: str = "test-task") -> dict[str, object]: @@ -70,16 +69,30 @@ def _write_family( return path -def test_packaged_registry_discovers_alpaca_eval(): +def test_packaged_registry_discovers_versioned_tasks(): tasks = load_tasks() - resolved = tasks["alpaca-eval"] - - assert list(tasks) == ["alpaca-eval"] - assert resolved.spec.dataset.sources["examples"].revision == ( + alpaca = tasks["alpaca-eval"] + arena_v01 = tasks["arena-hard-v0.1"] + arena_v20 = tasks["arena-hard-v2.0"] + + assert list(tasks) == [ + "alpaca-eval", + "arena-hard-v0.1", + "arena-hard-v2.0", + ] + assert alpaca.spec.dataset.sources["examples"].revision == ( "004c4a992956eeefffd36b63ade470f32fd0a582" ) - assert resolved.spec.protocol.baseline.reference_id == "gpt4_1106_preview" - assert resolved.spec.protocol.scoring.primary_metric == "winrate" + assert alpaca.spec.protocol.baseline.reference_id == "gpt4_1106_preview" + assert arena_v01.spec.protocol.baseline.reference_id == "gpt-4-0314" + assert arena_v20.spec.protocol.baseline.references["hard_prompt"] == ( + "o3-mini-2025-01-31" + ) + assert [resource.path for resource in arena_v20.provenance.resources] == [ + "arena_hard/_base.yaml", + "arena_hard/arena-hard-v2.0.yaml", + ] + assert alpaca.spec.protocol.scoring.primary_metric == "winrate" def test_find_returns_none_for_unregistered_task(): @@ -227,6 +240,25 @@ def test_official_outputs_must_reference_declared_source(tmp_path): load_tasks(tmp_path) +def test_category_baseline_uses_declared_category_field(tmp_path): + definition = _task_definition() + definition["dataset"]["fields"]["category"] = "category" + definition["protocol"]["baseline"] = { + "strategy": "category_defaults", + "category_field": "other_category", + "references": {"test": "reference"}, + } + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=definition, + ) + + with pytest.raises(TaskDefinitionError, match="dataset.fields.category"): + load_tasks(tmp_path) + + def test_resolved_hash_ignores_yaml_formatting(tmp_path): definition = _task_definition() path = _write_family( diff --git a/tests/test_utils.py b/tests/test_utils.py index 3ff4c06..4bcd55a 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -31,7 +31,7 @@ def test_safe_parse_int(monkeypatch, raw, expected): assert safe_parse_int(var) == expected -def test_download_all_dispatches_arena_hard_versions(monkeypatch, tmp_path): +def test_download_all_dispatches_registered_tasks(monkeypatch, tmp_path): calls: list[tuple[str, str, object]] = [] monkeypatch.setattr(utils_io, "data_root", tmp_path) @@ -40,13 +40,6 @@ def test_download_all_dispatches_arena_hard_versions(monkeypatch, tmp_path): "download_hf", lambda name, local_path: calls.append(("hf", name, local_path)), ) - monkeypatch.setattr( - utils_io, - "download_arena_hard", - lambda dataset, local_tables_path: calls.append( - ("arena", dataset, local_tables_path) - ), - ) monkeypatch.setattr( fluency_mod, "snapshot_download", @@ -65,8 +58,8 @@ def test_download_all_dispatches_arena_hard_versions(monkeypatch, tmp_path): tables_dir = tmp_path / "tables" assert calls[:5] == [ ("hf", "alpaca-eval", tables_dir), - ("arena", "arena-hard-v0.1", tables_dir), - ("arena", "arena-hard-v2.0", tables_dir), + ("hf", "arena-hard-v0.1", tables_dir), + ("hf", "arena-hard-v2.0", tables_dir), ("hf", "m-arena-hard-v0.1", tables_dir), ("hf", "m-arena-hard-v2.0", tables_dir), ]