diff --git a/judgearena/__main__.py b/judgearena/__main__.py new file mode 100644 index 0000000..a6db602 --- /dev/null +++ b/judgearena/__main__.py @@ -0,0 +1,3 @@ +from judgearena.cli import cli + +cli() diff --git a/judgearena/artifacts/metadata.py b/judgearena/artifacts/metadata.py index 659f5c7..44b165f 100644 --- a/judgearena/artifacts/metadata.py +++ b/judgearena/artifacts/metadata.py @@ -221,6 +221,28 @@ def _compact_results(results: dict[str, Any] | None) -> dict[str, Any]: return payload +def _task_definition_metadata(run: dict[str, Any]) -> dict[str, Any] | None: + """Return compact provenance when ``run.task`` is a packaged task.""" + task_id = run.get("task") + if not isinstance(task_id, str): + return None + + from judgearena.tasks.registry import get_packaged_task + + resolved = get_packaged_task(task_id) + if resolved is None: + return None + return { + "schema_version": resolved.spec.schema_version, + "task_version": resolved.spec.task_version, + "resolved_sha256": resolved.provenance.resolved_sha256, + "resources": [ + {"path": resource.path, "sha256": resource.sha256} + for resource in resolved.provenance.resources + ], + } + + def write_run_metadata( *, output_dir: str | Path, @@ -262,6 +284,10 @@ def write_run_metadata( start_path=Path(__file__).resolve().parent ), } + task_definition = _task_definition_metadata(run) + if task_definition is not None: + metadata["task_definition"] = task_definition + git_hash = _get_git_hash(start_path=Path(__file__).resolve().parent) if git_hash: metadata["git_hash"] = git_hash diff --git a/judgearena/benchmarks/pairwise/baselines.py b/judgearena/benchmarks/pairwise/baselines.py index bc0e167..d2b2ae1 100644 --- a/judgearena/benchmarks/pairwise/baselines.py +++ b/judgearena/benchmarks/pairwise/baselines.py @@ -10,13 +10,10 @@ split_m_arena_hard_dataset, ) from judgearena.datasets.mt_bench import MT_BENCH_BASELINES +from judgearena.tasks.registry import get_packaged_task +from judgearena.tasks.schema import CategoryDefaultsBaseline, TaskDefaultBaseline -ALPACA_EVAL_BASELINES: dict[str, str] = { - "alpaca-eval": "gpt4_1106_preview", -} - -PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = { - **ALPACA_EVAL_BASELINES, +LEGACY_PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = { **ARENA_HARD_BASELINES, **M_ARENA_HARD_BASELINES, **MT_BENCH_BASELINES, @@ -24,11 +21,20 @@ def native_pairwise_baseline(task: str) -> str | Mapping[str, str] | None: - """Return the dataset-native pairwise baseline, if the task defines one.""" - if task in PAIRWISE_BASELINES: - return PAIRWISE_BASELINES[task] + """Return the task-defined baseline, with fallback for unmigrated tasks.""" + resolved = get_packaged_task(task) + if resolved is not None: + baseline = resolved.spec.protocol.baseline + if isinstance(baseline, TaskDefaultBaseline): + return baseline.reference_id + if isinstance(baseline, CategoryDefaultsBaseline): + return baseline.references + return None + + if task in LEGACY_PAIRWISE_BASELINES: + return LEGACY_PAIRWISE_BASELINES[task] parsed_m_arena_hard = split_m_arena_hard_dataset(task) if parsed_m_arena_hard is not None: version_key, _lang_or_subset = parsed_m_arena_hard - return PAIRWISE_BASELINES[version_key] + return LEGACY_PAIRWISE_BASELINES[version_key] return None diff --git a/judgearena/benchmarks/pairwise/runner.py b/judgearena/benchmarks/pairwise/runner.py index c91817f..6b41d30 100644 --- a/judgearena/benchmarks/pairwise/runner.py +++ b/judgearena/benchmarks/pairwise/runner.py @@ -25,6 +25,7 @@ from judgearena.evaluate import judge_and_parse_prefs, resolve_run_judge_prompt from judgearena.generate import generate_base, generate_instructions from judgearena.log import get_logger +from judgearena.tasks.registry import get_packaged_task from judgearena.utils import ( cache_function_dataframe, compute_pref_summary, @@ -55,14 +56,25 @@ def try_load_dataset_completions( or ``None`` when no pre-existing completions are found. """ local_path_tables = data_root / "tables" - if is_arena_hard_dataset(dataset): + resolved_task = get_packaged_task(dataset) + if resolved_task is not None: + from judgearena.datasets.judgearena_tables import load_task_model_outputs + + df_outputs = load_task_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" - if not output_path.exists(): - return None - df_outputs = read_df(output_path) + output_path = local_path_tables / "model_outputs" / f"{dataset}.csv.zip" + if not output_path.exists(): + return None + df_outputs = read_df(output_path) df_outputs.loc[:, "output"] = df_outputs.loc[:, "output"].fillna("") df_outputs = df_outputs.pivot_table( index="instruction_index", columns="model", values="output", aggfunc="last" diff --git a/judgearena/benchmarks/registry.py b/judgearena/benchmarks/registry.py index 525eb12..f5a1fbf 100644 --- a/judgearena/benchmarks/registry.py +++ b/judgearena/benchmarks/registry.py @@ -5,6 +5,8 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Protocol +from judgearena.tasks.registry import get_packaged_task + if TYPE_CHECKING: from judgearena.config import RunConfig @@ -39,8 +41,17 @@ def benchmark_adapters() -> tuple[BenchmarkAdapter, ...]: def resolve_benchmark_adapter(task: str) -> BenchmarkAdapter: - """Return the first adapter supporting ``task``.""" - for adapter in benchmark_adapters(): + """Resolve a YAML-selected runner, then fall back for unmigrated tasks.""" + adapters = benchmark_adapters() + resolved = get_packaged_task(task) + if resolved is not None: + runner_id = resolved.spec.protocol.runner + for adapter in adapters: + if adapter.name == runner_id: + return adapter + raise ValueError(f"Task {task!r} selects unavailable runner {runner_id!r}.") + + for adapter in adapters: if adapter.supports(task): return adapter raise ValueError(f"No generate-and-evaluate adapter supports task {task!r}.") diff --git a/judgearena/cli.py b/judgearena/cli.py index ba24376..4a7d339 100644 --- a/judgearena/cli.py +++ b/judgearena/cli.py @@ -1,12 +1,13 @@ -"""Unified CLI entrypoint for judgearena. +"""Unified CLI entrypoint for JudgeArena. -Builds a ``RunConfig`` from CLI flags (derived from the config model) and/or a -``--config_path`` YAML, then dispatches to the ELO or generate-and-judge flow -based on ``--task`` (``elo-`` prefix runs the ELO rating flow). +Task-management subcommands are handled before the existing model-driven run +configuration and benchmark dispatch. """ from __future__ import annotations +import sys + from pydantic import ValidationError from judgearena.benchmarks.elo.runner import main as main_elo @@ -27,17 +28,28 @@ def _format_config_error(exc: ValidationError) -> str: def cli(argv: list[str] | None = None) -> None: - try: - cfg = build_run_config(argv) - except ValidationError as exc: - raise SystemExit(_format_config_error(exc)) from exc - - configure_logging(cfg.run.verbosity, log_file=cfg.run.log_file) - logger.debug("Running with config: %s", cfg.model_dump()) - if cfg.task.startswith(ELO_TASK_PREFIX): - main_elo(cfg) + args = list(sys.argv[1:] if argv is None else argv) + # `judgearena tasks {list,show,validate}` inspects packaged task definitions + # instead of running an evaluation, so it has its own argparse grammar (see + # tasks/cli.py) and must be routed before build_run_config, which only parses + # run/eval flags and would reject the subcommand form. + is_task_cli = args[:1] == ["tasks"] + if is_task_cli: + from judgearena.tasks.cli import run_task_command + + run_task_command(args[1:]) else: - run_benchmark(cfg) + try: + cfg = build_run_config(args) + except ValidationError as exc: + raise SystemExit(_format_config_error(exc)) from exc + + configure_logging(cfg.run.verbosity, log_file=cfg.run.log_file) + logger.debug("Running with config: %s", cfg.model_dump()) + if cfg.task.startswith(ELO_TASK_PREFIX): + main_elo(cfg) + else: + run_benchmark(cfg) if __name__ == "__main__": diff --git a/judgearena/config.py b/judgearena/config.py index aac5a1f..79479a7 100644 --- a/judgearena/config.py +++ b/judgearena/config.py @@ -18,6 +18,7 @@ from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline from judgearena.constants import ELO_TASK_PREFIX, ELO_TASK_TO_ARENA +from judgearena.tasks.registry import get_packaged_task # Set by build_run_config() for the duration of RunConfig() construction. _ACTIVE_CONFIG_PATH: str | None = None @@ -366,9 +367,8 @@ class RunConfig(BaseSettings): ) task: str - """Benchmark to run. Generate+judge: ``alpaca-eval``, ``arena-hard-v2.0``, - ``m-arena-hard-*``, ``mt-bench``, ``fluency-*``. ELO: ``elo-lmarena-100k``, - ``elo-lmarena-140k``, ``elo-lmarena``, ``elo-comparia``.""" + """Benchmark task ID. Use ``judgearena tasks list`` for packaged tasks; + legacy ELO task IDs use the ``elo-*`` prefix.""" model: ModelArgs = Field(default_factory=ModelArgs) """Model(s) under evaluation and their generation settings.""" @@ -387,6 +387,33 @@ class RunConfig(BaseSettings): @model_validator(mode="after") def _validate(self) -> RunConfig: + resolved_task = get_packaged_task(self.task) + if resolved_task is not None: + task_judge = resolved_task.spec.protocol.judge + if "swap_mode" not in self.judge.model_fields_set: + self.judge.swap_mode = task_judge.default_swap_mode + if self.judge.swap_mode not in task_judge.allowed_swap_modes: + raise ValueError( + f"judge.swap_mode={self.judge.swap_mode!r} is not supported " + f"by task {self.task!r}; choose from " + f"{list(task_judge.allowed_swap_modes)}." + ) + if ( + self.judge.temperature is None + and task_judge.default_temperature is not None + ): + self.judge.temperature = task_judge.default_temperature + + baseline = resolved_task.spec.protocol.baseline + if ( + self.model.baseline is not None + and getattr(baseline, "allow_runtime_override", True) is False + ): + raise ValueError( + f"model.baseline cannot override the baseline defined by " + f"task {self.task!r}." + ) + is_elo = self.task.startswith(ELO_TASK_PREFIX) if is_elo: if self.elo is None: diff --git a/judgearena/dataset_revisions.py b/judgearena/dataset_revisions.py index d5132fa..96b220c 100644 --- a/judgearena/dataset_revisions.py +++ b/judgearena/dataset_revisions.py @@ -25,10 +25,6 @@ # m-ArenaHard (Cohere release) "CohereLabs/m-ArenaHard": "ab393a96cd0b134a1acfa96e080af31e5e73a393", "CohereLabs/m-ArenaHard-v2.0": "24c65eff42cec85e30dd5db99d1a702c7ebaa8ab", - # AlpacaEval instructions / model_outputs. Originally ``geoalgo/llmjudge``, - # since renamed to ``judge-arena/judge-arena-dataset`` upstream (the old id - # still redirects, and the commit SHA is shared across the rename). - "judge-arena/judge-arena-dataset": "004c4a992956eeefffd36b63ade470f32fd0a582", # MT-Bench questions (LMSYS Space). "lmsys/mt-bench": "a4b674ca573c24143824ac7f60d9173e7081e37d", # Multilingual fluency contexts (generated by scripts/fluency/generate_fluency.py). diff --git a/judgearena/datasets/__init__.py b/judgearena/datasets/__init__.py index 9e1b603..5c0a13f 100644 --- a/judgearena/datasets/__init__.py +++ b/judgearena/datasets/__init__.py @@ -9,12 +9,22 @@ split_m_arena_hard_dataset, ) from judgearena.log import get_logger +from judgearena.tasks.registry import get_packaged_task logger = get_logger(__name__) def load_instructions(dataset: str, n_instructions: int | None = None) -> pd.DataFrame: - if dataset == "mt-bench": + 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 + + df_instructions = load_task_instructions( + resolved_task, judgearena_utils.data_root / "tables" + ) + + elif dataset == "mt-bench": from judgearena.datasets.mt_bench import load_mt_bench df_instructions = load_mt_bench() @@ -46,7 +56,6 @@ def load_instructions(dataset: str, n_instructions: int | None = None) -> pd.Dat else: assert dataset in [ - "alpaca-eval", "arena-hard-v0.1", "arena-hard-v2.0", ] diff --git a/judgearena/datasets/judgearena_tables.py b/judgearena/datasets/judgearena_tables.py new file mode 100644 index 0000000..113eb88 --- /dev/null +++ b/judgearena/datasets/judgearena_tables.py @@ -0,0 +1,70 @@ +"""Dataset adapter for JudgeArena's packaged instruction/output tables.""" + +from __future__ import annotations + +from pathlib import Path + +import pandas as pd +from huggingface_hub import snapshot_download + +from judgearena.tasks.schema import HuggingFaceDatasetSource, ResolvedTaskSpec + + +def download_task_sources(task: ResolvedTaskSpec, local_dir: Path) -> None: + """Download every Hugging Face source declared by a table-backed task.""" + if task.spec.dataset.adapter != "judgearena_tables": + raise ValueError( + f"Task {task.task!r} uses dataset adapter " + f"{task.spec.dataset.adapter!r}, not 'judgearena_tables'." + ) + local_dir.mkdir(exist_ok=True, parents=True) + for name, source in task.spec.dataset.sources.items(): + if not isinstance(source, HuggingFaceDatasetSource): + raise ValueError( + f"Dataset source {name!r} for task {task.task!r} is not supported " + "by the 'judgearena_tables' adapter." + ) + snapshot_download( + repo_id=source.repo_id, + repo_type="dataset", + revision=source.revision, + allow_patterns=list(source.allow_patterns) or None, + local_dir=local_dir, + force_download=False, + ) + + +def load_task_instructions( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame: + """Load a task's table and map its declared fields to runner names.""" + download_task_sources(task, local_tables_path) + path = local_tables_path / "instructions" / f"{task.task}.csv" + if not path.exists(): + raise FileNotFoundError(f"Instruction table not found at {path}") + df = pd.read_csv(path) + + fields = task.spec.dataset.fields + field_mapping = { + fields.id: "instruction_index", + fields.instruction: "instruction", + } + if fields.category is not None: + field_mapping[fields.category] = "category" + missing = sorted(set(field_mapping) - set(df.columns)) + if missing: + raise ValueError( + f"Task {task.task!r} is missing declared dataset fields: {missing}." + ) + return df.rename(columns=field_mapping) + + +def load_task_model_outputs( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame | None: + """Load optional pre-generated model outputs for a table-backed task.""" + download_task_sources(task, local_tables_path) + path = local_tables_path / "model_outputs" / f"{task.task}.csv.zip" + if not path.exists(): + return None + return pd.read_csv(path) diff --git a/judgearena/paths.py b/judgearena/paths.py index e4a8c77..097bf21 100644 --- a/judgearena/paths.py +++ b/judgearena/paths.py @@ -7,7 +7,7 @@ Symbols here are re-exported from :mod:`judgearena.utils` for backward compatibility, so existing ``from judgearena.utils import data_root`` / -``from judgearena.utils import download_hf, read_df`` callers keep working. +``from judgearena.utils import read_df`` callers keep working. """ from __future__ import annotations @@ -16,9 +16,6 @@ from pathlib import Path import pandas as pd -from huggingface_hub import snapshot_download - -from judgearena.dataset_revisions import hf_revision def _data_root_path() -> Path: @@ -31,20 +28,6 @@ def _data_root_path() -> Path: data_root: Path = _data_root_path() -def download_hf(name: str, local_path: Path) -> None: - """Download AlpacaEval-style instruction/output tables into ``local_path``.""" - local_path.mkdir(exist_ok=True, parents=True) - repo_id = "judge-arena/judge-arena-dataset" - snapshot_download( - repo_id=repo_id, - repo_type="dataset", - allow_patterns=f"*{name}*", - local_dir=local_path, - force_download=False, - revision=hf_revision(repo_id), - ) - - def read_df(filename: Path, **pandas_kwargs) -> pd.DataFrame: """Read a CSV/CSV-zip/parquet dataframe from disk.""" assert filename.exists(), f"Dataframe file not found at {filename}" diff --git a/judgearena/prompts/registry.py b/judgearena/prompts/registry.py index 12f398b..d7f184b 100644 --- a/judgearena/prompts/registry.py +++ b/judgearena/prompts/registry.py @@ -93,7 +93,6 @@ def metadata(self) -> dict[str, str | bool | None]: JUDGE_PROMPT_PRESETS = tuple(PRESETS) TASK_DEFAULT_PRESET: dict[str, str] = { - "alpaca-eval": DEFAULT_JUDGE_PROMPT_PRESET, "arena-hard-v0.1": DEFAULT_JUDGE_PROMPT_PRESET, "arena-hard-v2.0": DEFAULT_JUDGE_PROMPT_PRESET, "mt-bench": FASTCHAT_PAIRWISE_PROMPT_PRESET, @@ -103,6 +102,12 @@ def metadata(self) -> dict[str, str | bool | None]: def default_preset_for_task(task: str | None) -> str: if task is None: return DEFAULT_JUDGE_PROMPT_PRESET + # Import lazily: task validation uses the prompt catalog in this module. + from judgearena.tasks.registry import get_packaged_task + + resolved = get_packaged_task(task) + if resolved is not None: + return resolved.spec.protocol.judge.default_prompt if task in TASK_DEFAULT_PRESET: return TASK_DEFAULT_PRESET[task] if task.startswith("m-arena-hard"): diff --git a/judgearena/tasks/__init__.py b/judgearena/tasks/__init__.py new file mode 100644 index 0000000..7a7a538 --- /dev/null +++ b/judgearena/tasks/__init__.py @@ -0,0 +1,6 @@ +"""Public API for declaring, validating, and discovering benchmark tasks.""" + +from judgearena.tasks.registry import get_packaged_task, load_tasks +from judgearena.tasks.schema import ResolvedTaskSpec, TaskSpec + +__all__ = ["ResolvedTaskSpec", "TaskSpec", "get_packaged_task", "load_tasks"] diff --git a/judgearena/tasks/cli.py b/judgearena/tasks/cli.py new file mode 100644 index 0000000..9dcba9f --- /dev/null +++ b/judgearena/tasks/cli.py @@ -0,0 +1,70 @@ +"""Implement ``judgearena tasks list|show|validate`` inspection commands.""" + +from __future__ import annotations + +import argparse +from collections.abc import Sequence +from dataclasses import asdict + +import yaml + +from judgearena.tasks.loader import TaskDefinitionError +from judgearena.tasks.registry import load_tasks +from judgearena.tasks.schema import ResolvedTaskSpec + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(prog="judgearena tasks") + commands = parser.add_subparsers(dest="command", required=True) + commands.add_parser("list", help="List packaged tasks.") + + show = commands.add_parser("show", help="Show a packaged task definition.") + show.add_argument("task") + show.add_argument( + "--resolved", + action="store_true", + help="Include inheritance and digest provenance.", + ) + + validate = commands.add_parser("validate", help="Validate packaged tasks.") + validate.add_argument("task", nargs="?") + return parser + + +def run_task_command( + argv: Sequence[str], *, tasks: dict[str, ResolvedTaskSpec] | None = None +) -> None: + """Run a task-registry command without starting an evaluation.""" + parser = _parser() + args = parser.parse_args(list(argv)) + try: + tasks = load_tasks() if tasks is None else tasks + if args.command == "list": + for resolved in tasks.values(): + spec = resolved.spec + print(f"{spec.task}\tv{spec.task_version}\t{spec.description}") + elif args.command == "show": + resolved = _require(parser, tasks, args.task) + output = resolved.model_dump() + if args.resolved: + output["_provenance"] = asdict(resolved.provenance) + print(yaml.safe_dump(output, sort_keys=False).rstrip()) + elif args.command == "validate": + if args.task is not None: + _require(parser, tasks, args.task) + print(f"Validated task {args.task!r}.") + else: + print(f"Validated {len(tasks)} task(s).") + except TaskDefinitionError as exc: + parser.error(str(exc)) + + +def _require( + parser: argparse.ArgumentParser, + tasks: dict[str, ResolvedTaskSpec], + task_id: str, +) -> ResolvedTaskSpec: + if task_id not in tasks: + known = ", ".join(sorted(tasks)) or "none" + parser.error(f"unknown task {task_id!r}; registered tasks: {known}") + return tasks[task_id] diff --git a/judgearena/tasks/definitions/alpaca_eval/alpaca-eval.yaml b/judgearena/tasks/definitions/alpaca_eval/alpaca-eval.yaml new file mode 100644 index 0000000..fdd972e --- /dev/null +++ b/judgearena/tasks/definitions/alpaca_eval/alpaca-eval.yaml @@ -0,0 +1,38 @@ +schema_version: 1 +task: alpaca-eval +task_version: 1 +description: Pairwise evaluation on the AlpacaEval instruction set. +tags: [pairwise, instruction-following] + +dataset: + adapter: judgearena_tables + sources: + examples: + type: huggingface_dataset + repo_id: judge-arena/judge-arena-dataset + revision: "004c4a992956eeefffd36b63ade470f32fd0a582" + allow_patterns: ["*alpaca-eval*"] + fields: + id: instruction_index + instruction: instruction + +protocol: + runner: pairwise + generation: + mode: single_turn_chat + baseline: + strategy: task_default + reference_id: gpt4_1106_preview + allow_runtime_override: true + 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/tatsu-lab/alpaca_eval diff --git a/judgearena/tasks/loader.py b/judgearena/tasks/loader.py new file mode 100644 index 0000000..2a1e267 --- /dev/null +++ b/judgearena/tasks/loader.py @@ -0,0 +1,231 @@ +"""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 new file mode 100644 index 0000000..2159e50 --- /dev/null +++ b/judgearena/tasks/registry.py @@ -0,0 +1,83 @@ +"""Discover task definitions and validate their referenced component IDs.""" + +from __future__ import annotations + +from dataclasses import dataclass +from functools import cache +from importlib.resources import files +from importlib.resources.abc import Traversable + +from judgearena.prompts.registry import JUDGE_PROMPT_PRESETS +from judgearena.tasks.loader import TaskDefinitionError, TaskLoader +from judgearena.tasks.schema import ResolvedTaskSpec + + +@dataclass(frozen=True) +class AdapterCatalog: + """Component IDs that task YAML files may reference.""" + + runners: frozenset[str] = frozenset({"pairwise"}) + datasets: frozenset[str] = frozenset({"judgearena_tables"}) + prompts: frozenset[str] = frozenset(JUDGE_PROMPT_PRESETS) + parsers: frozenset[str] = frozenset({"pairwise_preference"}) + scorers: frozenset[str] = frozenset({"pairwise_win_rate"}) + + +def load_tasks( + definitions_root: Traversable | None = None, + *, + adapters: AdapterCatalog | None = None, +) -> dict[str, ResolvedTaskSpec]: + """Discover, validate, and return all packaged tasks keyed by task ID. + + With no arguments this reads JudgeArena's installed definitions and caches + the result. An explicit ``definitions_root`` (used by tests) is never cached. + """ + if definitions_root is None and adapters is None: + return _load_packaged_tasks() + root = definitions_root or files("judgearena.tasks").joinpath("definitions") + return _discover_tasks(root, adapters or AdapterCatalog()) + + +@cache +def _load_packaged_tasks() -> dict[str, ResolvedTaskSpec]: + root = files("judgearena.tasks").joinpath("definitions") + return _discover_tasks(root, AdapterCatalog()) + + +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) + if resolved.task in tasks: + other = tasks[resolved.task].provenance.source_path + raise TaskDefinitionError( + f"Duplicate task ID {resolved.task!r} in {other} and {relative_path}" + ) + _validate_adapter_ids(resolved, adapters) + tasks[resolved.task] = resolved + return dict(sorted(tasks.items())) + + +def _validate_adapter_ids(resolved: ResolvedTaskSpec, adapters: AdapterCatalog) -> None: + spec = resolved.spec + references = { + "runner": (spec.protocol.runner, adapters.runners), + "dataset adapter": (spec.dataset.adapter, adapters.datasets), + "prompt": (spec.protocol.judge.default_prompt, adapters.prompts), + "parser": (spec.protocol.judge.parser, adapters.parsers), + "scorer": (spec.protocol.scoring.adapter, adapters.scorers), + } + for kind, (adapter_id, available) in references.items(): + if adapter_id not in available: + raise TaskDefinitionError( + f"{resolved.provenance.source_path}: unknown {kind} {adapter_id!r}" + ) + + +def get_packaged_task(task_id: str) -> ResolvedTaskSpec | None: + """Look up a task from JudgeArena's installed YAML definitions.""" + return load_tasks().get(task_id) diff --git a/judgearena/tasks/schema.py b/judgearena/tasks/schema.py new file mode 100644 index 0000000..2715ef8 --- /dev/null +++ b/judgearena/tasks/schema.py @@ -0,0 +1,208 @@ +"""Typed contract for task YAML; this module performs no loading or execution.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Annotated, Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator + + +class _StrictFrozenModel(BaseModel): + """Immutable schema node that rejects unknown YAML fields.""" + + model_config = ConfigDict(extra="forbid", frozen=True) + + +class HuggingFaceDatasetSource(_StrictFrozenModel): + type: Literal["huggingface_dataset"] + repo_id: str = Field(min_length=1) + revision: str = Field(pattern=r"^[0-9a-fA-F]{40}$") + config: str | None = None + split: str | None = None + allow_patterns: tuple[str, ...] = () + + +class HuggingFaceSpaceSource(_StrictFrozenModel): + type: Literal["huggingface_space"] + repo_id: str = Field(min_length=1) + revision: str = Field(pattern=r"^[0-9a-fA-F]{40}$") + allow_patterns: tuple[str, ...] = () + + +class GitRawSource(_StrictFrozenModel): + type: Literal["git_raw"] + repository: str = Field(min_length=1) + revision: str = Field(pattern=r"^[0-9a-fA-F]{40}$") + path: str = Field(min_length=1) + + +class LocalSource(_StrictFrozenModel): + type: Literal["local"] + path: str = Field(min_length=1) + format: Literal["csv", "json", "jsonl", "parquet"] + sha256: str = Field(pattern=r"^[0-9a-fA-F]{64}$") + + +SourceSpec = Annotated[ + HuggingFaceDatasetSource | HuggingFaceSpaceSource | GitRawSource | LocalSource, + Field(discriminator="type"), +] + + +class DatasetFields(_StrictFrozenModel): + id: str = Field(min_length=1) + instruction: str = Field(min_length=1) + category: str | None = None + + +class DatasetSpec(_StrictFrozenModel): + """Dataset sources, loader adapter, and canonical field mapping.""" + + adapter: str = Field(min_length=1) + sources: dict[str, SourceSpec] = Field(min_length=1) + fields: DatasetFields + + +class SingleTurnGeneration(_StrictFrozenModel): + mode: Literal["single_turn_chat"] + + +class NoBaseline(_StrictFrozenModel): + strategy: Literal["none"] + + +class RuntimeRequiredBaseline(_StrictFrozenModel): + strategy: Literal["runtime_required"] + + +class TaskDefaultBaseline(_StrictFrozenModel): + strategy: Literal["task_default"] + reference_id: str = Field(min_length=1) + allow_runtime_override: bool = True + + +class CategoryDefaultsBaseline(_StrictFrozenModel): + strategy: Literal["category_defaults"] + category_field: str = Field(min_length=1) + references: dict[str, str] = Field(min_length=1) + allow_runtime_override: bool = True + + +class OfficialOutputsBaseline(_StrictFrozenModel): + strategy: Literal["official_outputs"] + source: str = Field(min_length=1) + + +BaselineSpec = Annotated[ + NoBaseline + | RuntimeRequiredBaseline + | TaskDefaultBaseline + | CategoryDefaultsBaseline + | OfficialOutputsBaseline, + Field(discriminator="strategy"), +] + +SwapMode = Literal["fixed", "both"] + + +class PairwiseJudgeSpec(_StrictFrozenModel): + default_prompt: str = Field(min_length=1) + parser: str = Field(min_length=1) + default_swap_mode: SwapMode = "fixed" + allowed_swap_modes: tuple[SwapMode, ...] = ("fixed", "both") + default_temperature: float | None = None + + @model_validator(mode="after") + def _default_must_be_allowed(self) -> PairwiseJudgeSpec: + if not self.allowed_swap_modes: + raise ValueError("allowed_swap_modes must not be empty") + if self.default_swap_mode not in self.allowed_swap_modes: + raise ValueError("default_swap_mode must be present in allowed_swap_modes") + if len(set(self.allowed_swap_modes)) != len(self.allowed_swap_modes): + raise ValueError("allowed_swap_modes must not contain duplicates") + return self + + +class ScoringSpec(_StrictFrozenModel): + adapter: str = Field(min_length=1) + primary_metric: str = Field(min_length=1) + higher_is_better: bool + + +class PairwiseProtocol(_StrictFrozenModel): + """Task-owned generation, baseline, judging, and scoring behavior.""" + + runner: Literal["pairwise"] + generation: SingleTurnGeneration + baseline: BaselineSpec + judge: PairwiseJudgeSpec + scoring: ScoringSpec + + +class TaskMetadata(_StrictFrozenModel): + reference_implementation: str | None = None + paper: str | None = None + + +class TaskSpec(_StrictFrozenModel): + """Complete validated definition of one registered task.""" + + schema_version: Literal[1] + task: str = Field(pattern=r"^[A-Za-z0-9][A-Za-z0-9._-]*$") + task_version: int = Field(ge=1) + description: str = Field(min_length=1) + tags: tuple[str, ...] = () + dataset: DatasetSpec + protocol: PairwiseProtocol + metadata: TaskMetadata = Field(default_factory=TaskMetadata) + + @model_validator(mode="after") + def _validate_task(self) -> TaskSpec: + if len(set(self.tags)) != len(self.tags): + raise ValueError("tags must not contain duplicates") + source_names = set(self.dataset.sources) + baseline = self.protocol.baseline + if ( + isinstance(baseline, OfficialOutputsBaseline) + and baseline.source not in source_names + ): + raise ValueError( + f"official baseline source {baseline.source!r} is not declared " + "in dataset.sources" + ) + return self + + +@dataclass(frozen=True) +class ResourceDigest: + """Hash of one YAML resource used to construct a task.""" + + path: str + sha256: str + + +@dataclass(frozen=True) +class TaskProvenance: + """Source paths and hashes needed to identify a resolved definition.""" + + source_path: str + source_sha256: str + resolved_sha256: str + resources: tuple[ResourceDigest, ...] + + +@dataclass(frozen=True) +class ResolvedTaskSpec: + """Validated task plus the provenance of its resolved YAML.""" + + spec: TaskSpec + provenance: TaskProvenance + + @property + def task(self) -> str: + return self.spec.task + + def model_dump(self) -> dict[str, object]: + """Return the normalized task definition without provenance.""" + return self.spec.model_dump(mode="json") diff --git a/judgearena/utils/io.py b/judgearena/utils/io.py index da9e9f5..3eb5a6c 100644 --- a/judgearena/utils/io.py +++ b/judgearena/utils/io.py @@ -31,15 +31,25 @@ def _data_root_path() -> Path: def download_hf(name: str, local_path: Path): - local_path.mkdir(exist_ok=True, parents=True) - # downloads the model from huggingface into `local_path` folder - snapshot_download( - repo_id="judge-arena/judge-arena-dataset", - repo_type="dataset", - allow_patterns=f"*{name}*", - local_dir=local_path, - force_download=False, - ) + # A `name` is either a packaged task (download the HF sources declared in + # its task YAML) or a legacy dataset name (pattern-matched out of the + # monolithic judge-arena-dataset repo). + from judgearena.tasks.registry import get_packaged_task + + resolved_task = get_packaged_task(name) + if resolved_task is not None: + from judgearena.datasets.judgearena_tables import download_task_sources + + download_task_sources(resolved_task, local_path) + else: + local_path.mkdir(exist_ok=True, parents=True) + snapshot_download( + repo_id="judge-arena/judge-arena-dataset", + repo_type="dataset", + allow_patterns=f"*{name}*", + local_dir=local_path, + force_download=False, + ) def read_df(filename: Path, **pandas_kwargs) -> pd.DataFrame: @@ -71,11 +81,17 @@ def safe_parse_int(env_var: str) -> int | None: def download_all(): from judgearena.datasets.fluency import download_fluency_dataset from judgearena.datasets.m_arenahard import M_ARENA_HARD_BASELINES + from judgearena.tasks.registry import load_tasks 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 ( - "alpaca-eval", + *packaged_table_tasks, "arena-hard-v0.1", "arena-hard-v2.0", *M_ARENA_HARD_BASELINES, diff --git a/pyproject.toml b/pyproject.toml index ba0e0a2..1c7c701 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,6 +54,7 @@ exclude = ["slurmpilot_scripts*"] [tool.setuptools.package-data] "judgearena.criteria" = ["data/*.yaml"] "judgearena.prompts" = ["*.txt", "*/*.txt"] +"judgearena.tasks" = ["definitions/*/*.yaml"] [dependency-groups] dev = [ diff --git a/tests/test_config.py b/tests/test_config.py index 39df6d5..f53ded7 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1,6 +1,9 @@ +from types import SimpleNamespace + import pytest from pydantic import ValidationError +import judgearena.config as config_module from judgearena import cli as cli_module from judgearena.config import RunConfig @@ -29,6 +32,65 @@ def test_generate_config_constructs(): assert cfg.elo is None +def _registered_task( + *, + default_swap_mode: str = "both", + allowed_swap_modes: tuple[str, ...] = ("both",), + default_temperature: float | None = 0.25, + allow_runtime_override: bool = True, +): + return SimpleNamespace( + spec=SimpleNamespace( + protocol=SimpleNamespace( + judge=SimpleNamespace( + default_swap_mode=default_swap_mode, + allowed_swap_modes=allowed_swap_modes, + default_temperature=default_temperature, + ), + baseline=SimpleNamespace(allow_runtime_override=allow_runtime_override), + ) + ) + ) + + +def test_registered_task_applies_judge_defaults(monkeypatch): + monkeypatch.setattr( + config_module, "get_packaged_task", lambda _task: _registered_task() + ) + data = _base_generate() + data["task"] = "yaml-task" + + cfg = RunConfig(**data) + + assert cfg.judge.swap_mode == "both" + assert cfg.judge.temperature == 0.25 + + +def test_registered_task_rejects_unsupported_swap_mode(monkeypatch): + monkeypatch.setattr( + config_module, "get_packaged_task", lambda _task: _registered_task() + ) + data = _base_generate() + data["task"] = "yaml-task" + data["judge"]["swap_mode"] = "fixed" + + with pytest.raises(ValidationError, match="not supported"): + RunConfig(**data) + + +def test_registered_task_can_forbid_baseline_override(monkeypatch): + monkeypatch.setattr( + config_module, + "get_packaged_task", + lambda _task: _registered_task(allow_runtime_override=False), + ) + data = _base_generate() + data["task"] = "yaml-task" + + with pytest.raises(ValidationError, match="cannot override"): + RunConfig(**data) + + def test_elo_config_derives_arena(): cfg = RunConfig(**_base_elo()) assert cfg.elo is not None diff --git a/tests/test_generate_and_evaluate.py b/tests/test_generate_and_evaluate.py index 3e824b9..57f58b4 100644 --- a/tests/test_generate_and_evaluate.py +++ b/tests/test_generate_and_evaluate.py @@ -1,15 +1,21 @@ +from types import SimpleNamespace + import pandas as pd import pytest import judgearena.benchmarks.execution as benchmark_execution import judgearena.benchmarks.pairwise.runner as generate_and_evaluate -from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline +import judgearena.benchmarks.registry as benchmark_registry +from judgearena.benchmarks.pairwise.baselines import ( + LEGACY_PAIRWISE_BASELINES, + native_pairwise_baseline, +) from judgearena.benchmarks.pairwise.runner import ( BaselinePlan, _resolve_baseline_plan, run_pairwise, ) -from judgearena.benchmarks.registry import resolve_benchmark_adapter +from judgearena.benchmarks.registry import BenchmarkAdapter, resolve_benchmark_adapter from judgearena.config import RunConfig @@ -127,6 +133,11 @@ def test_resolve_plan_alpaca_eval_uses_native_baseline(): assert plan.single_model == "gpt4_1106_preview" +def test_alpaca_eval_baseline_is_not_duplicated_in_legacy_registry(): + assert "alpaca-eval" not in LEGACY_PAIRWISE_BASELINES + assert native_pairwise_baseline("alpaca-eval") == "gpt4_1106_preview" + + def test_resolve_plan_explicit_model_b_overrides_native(): plan = _resolve_baseline_plan( task="arena-hard-v2.0", @@ -157,6 +168,20 @@ def test_benchmark_adapter_resolution(): assert resolve_benchmark_adapter("alpaca-eval").name == "pairwise" +def test_registered_task_runner_wins_over_legacy_fallback(monkeypatch): + fallback = BenchmarkAdapter("fallback", None, lambda _cfg: None) + pairwise = BenchmarkAdapter("pairwise", frozenset(), lambda _cfg: None) + resolved = SimpleNamespace( + spec=SimpleNamespace(protocol=SimpleNamespace(runner="pairwise")) + ) + monkeypatch.setattr( + benchmark_registry, "benchmark_adapters", lambda: (fallback, pairwise) + ) + monkeypatch.setattr(benchmark_registry, "get_packaged_task", lambda _task: resolved) + + assert benchmark_registry.resolve_benchmark_adapter("yaml-task") is pairwise + + def test_resolve_plan_task_without_native_baseline_requires_model_b(): with pytest.raises(ValueError, match="baseline"): _resolve_baseline_plan( diff --git a/tests/test_instruction_dataset.py b/tests/test_instruction_dataset.py index 9551691..142b657 100644 --- a/tests/test_instruction_dataset.py +++ b/tests/test_instruction_dataset.py @@ -5,6 +5,7 @@ import judgearena.benchmarks.pairwise.runner as generate_and_evaluate import judgearena.datasets as instruction_dataset +import judgearena.datasets.judgearena_tables as judgearena_tables import judgearena.utils as judgearena_utils from judgearena.datasets.arena_hard import ( ARENA_HARD_BASELINES, @@ -14,6 +15,45 @@ arena_hard_native_baseline, normalize_official_arena_hard, ) +from judgearena.tasks.registry import get_packaged_task + + +def test_alpaca_eval_table_download_uses_yaml_source(monkeypatch, tmp_path): + captured = {} + monkeypatch.setattr( + judgearena_tables, + "snapshot_download", + lambda **kwargs: captured.update(kwargs), + ) + task = get_packaged_task("alpaca-eval") + assert task is not None + + judgearena_tables.download_task_sources(task, tmp_path) + + assert captured["repo_id"] == "judge-arena/judge-arena-dataset" + assert captured["revision"] == "004c4a992956eeefffd36b63ade470f32fd0a582" + assert captured["allow_patterns"] == ["*alpaca-eval*"] + + +def test_alpaca_eval_table_loader_uses_yaml_fields(monkeypatch, tmp_path): + monkeypatch.setattr( + judgearena_tables, "download_task_sources", lambda _task, _path: None + ) + instructions_dir = tmp_path / "instructions" + instructions_dir.mkdir() + pd.DataFrame( + { + "instruction_index": [1, 2], + "instruction": ["First", "Second"], + } + ).to_csv(instructions_dir / "alpaca-eval.csv", index=False) + task = get_packaged_task("alpaca-eval") + assert task is not None + + loaded = judgearena_tables.load_task_instructions(task, tmp_path) + + assert loaded["instruction_index"].tolist() == [1, 2] + assert loaded["instruction"].tolist() == ["First", "Second"] def test_arena_hard_native_baseline_v01_is_flat_string(): diff --git a/tests/test_prompt_registry.py b/tests/test_prompt_registry.py index 604c54e..414bbdc 100644 --- a/tests/test_prompt_registry.py +++ b/tests/test_prompt_registry.py @@ -30,6 +30,11 @@ def test_default_preset_for_task_known_keys(): assert default_preset_for_task(task) == preset +def test_alpaca_eval_prompt_default_is_not_duplicated_in_legacy_registry(): + assert "alpaca-eval" not in TASK_DEFAULT_PRESET + assert default_preset_for_task("alpaca-eval") == "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_repro.py b/tests/test_repro.py index dd0ee40..e0d1560 100644 --- a/tests/test_repro.py +++ b/tests/test_repro.py @@ -95,3 +95,20 @@ def test_write_run_metadata_omits_optional_fields_when_inputs_missing( assert "instruction_indices_sha256" not in metadata assert "judge_system_prompt_sha256" not in metadata assert "judge_user_prompt_template_sha256" not in metadata + + +def test_write_run_metadata_records_packaged_task_provenance(tmp_path, monkeypatch): + monkeypatch.setattr(repro, "_get_dependency_versions", lambda *args, **kwargs: {}) + monkeypatch.setattr(repro, "_get_git_hash", lambda *args, **kwargs: None) + + metadata_path = repro.write_run_metadata( + output_dir=tmp_path, + entrypoint="judgearena.test.entrypoint", + run={"task": "alpaca-eval"}, + ) + + task_definition = json.loads(metadata_path.read_text())["task_definition"] + assert task_definition["schema_version"] == 1 + assert task_definition["task_version"] == 1 + assert len(task_definition["resolved_sha256"]) == 64 + assert task_definition["resources"][0]["path"] == ("alpaca_eval/alpaca-eval.yaml") diff --git a/tests/test_task_registry.py b/tests/test_task_registry.py new file mode 100644 index 0000000..835b8e5 --- /dev/null +++ b/tests/test_task_registry.py @@ -0,0 +1,290 @@ +"""Tests for declarative task loading, discovery, and static commands.""" + +from __future__ import annotations + +from pathlib import Path + +import pytest +import yaml + +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 + + +def _task_definition(task: str = "test-task") -> dict[str, object]: + return { + "schema_version": 1, + "task": task, + "task_version": 1, + "description": "Test pairwise task.", + "tags": ["pairwise", "test"], + "dataset": { + "adapter": "judgearena_tables", + "sources": { + "examples": { + "type": "huggingface_dataset", + "repo_id": "example/tasks", + "revision": "a" * 40, + "allow_patterns": [f"*{task}*"], + } + }, + "fields": {"id": "id", "instruction": "prompt"}, + }, + "protocol": { + "runner": "pairwise", + "generation": {"mode": "single_turn_chat"}, + "baseline": { + "strategy": "task_default", + "reference_id": "reference-output", + "allow_runtime_override": True, + }, + "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, + }, + }, + } + + +def _write_family( + root: Path, + *, + family: str, + filename: str, + definition: dict[str, object] | str, +) -> Path: + family_dir = root / family + family_dir.mkdir(parents=True, exist_ok=True) + path = family_dir / filename + text = definition if isinstance(definition, str) else yaml.safe_dump(definition) + path.write_text(text) + return path + + +def test_packaged_registry_discovers_alpaca_eval(): + tasks = load_tasks() + resolved = tasks["alpaca-eval"] + + assert list(tasks) == ["alpaca-eval"] + assert resolved.spec.dataset.sources["examples"].revision == ( + "004c4a992956eeefffd36b63ade470f32fd0a582" + ) + assert resolved.spec.protocol.baseline.reference_id == "gpt4_1106_preview" + assert resolved.spec.protocol.scoring.primary_metric == "winrate" + + +def test_find_returns_none_for_unregistered_task(): + assert load_tasks().get("not-packaged-yet") is None + + +def test_registry_rejects_unpinned_remote_source(tmp_path): + definition = _task_definition() + definition["dataset"]["sources"]["examples"]["revision"] = "main" + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=definition, + ) + + with pytest.raises(TaskDefinitionError, match="revision"): + load_tasks(tmp_path) + + +def test_registry_rejects_duplicate_yaml_keys(tmp_path): + text = yaml.safe_dump(_task_definition()) + "task: duplicate\n" + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=text, + ) + + with pytest.raises(TaskDefinitionError, match="duplicate key 'task'"): + load_tasks(tmp_path) + + +def test_registry_resolves_private_base_and_records_provenance(tmp_path): + definition = _task_definition() + child_task = definition.pop("task") + definition.pop("description") + definition["tags"] = ["base"] + _write_family( + tmp_path, + family="example", + filename="_base.yaml", + definition=definition, + ) + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition={ + "extends": "_base.yaml", + "task": child_task, + "description": "Resolved child.", + "tags": ["child"], + }, + ) + + resolved = load_tasks(tmp_path)["test-task"] + + assert resolved.spec.description == "Resolved child." + assert resolved.spec.tags == ("child",) + assert [item.path for item in resolved.provenance.resources] == [ + "example/_base.yaml", + "example/test-task.yaml", + ] + assert len(resolved.provenance.resolved_sha256) == 64 + + +def test_registry_rejects_inheritance_cycle(tmp_path): + _write_family( + tmp_path, + family="example", + filename="_a.yaml", + definition={"extends": "_b.yaml"}, + ) + _write_family( + tmp_path, + family="example", + filename="_b.yaml", + definition={"extends": "_a.yaml"}, + ) + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition={"extends": "_a.yaml", "task": "test-task"}, + ) + + with pytest.raises(TaskDefinitionError, match="inheritance cycle"): + load_tasks(tmp_path) + + +def test_registry_rejects_extends_path_escape(tmp_path): + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition={"extends": "../../_base.yaml", "task": "test-task"}, + ) + + with pytest.raises(TaskDefinitionError, match="path escapes"): + load_tasks(tmp_path) + + +def test_registry_rejects_duplicate_task_ids(tmp_path): + for family in ("one", "two"): + _write_family( + tmp_path, + family=family, + filename=f"{family}.yaml", + definition=_task_definition("same-task"), + ) + + with pytest.raises(TaskDefinitionError, match="Duplicate task ID 'same-task'"): + load_tasks(tmp_path) + + +def test_registry_rejects_unknown_adapter_id(tmp_path): + definition = _task_definition() + definition["dataset"]["adapter"] = "missing_loader" + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=definition, + ) + + with pytest.raises(TaskDefinitionError, match="unknown dataset adapter"): + load_tasks(tmp_path) + + +def test_official_outputs_must_reference_declared_source(tmp_path): + definition = _task_definition() + definition["protocol"]["baseline"] = { + "strategy": "official_outputs", + "source": "missing_outputs", + } + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=definition, + ) + + with pytest.raises(TaskDefinitionError, match="not declared in dataset.sources"): + load_tasks(tmp_path) + + +def test_resolved_hash_ignores_yaml_formatting(tmp_path): + definition = _task_definition() + path = _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=definition, + ) + first = load_tasks(tmp_path)["test-task"] + + path.write_text("# formatting-only change\n" + yaml.safe_dump(definition)) + second = load_tasks(tmp_path)["test-task"] + + assert first.provenance.source_sha256 != second.provenance.source_sha256 + assert first.provenance.resolved_sha256 == second.provenance.resolved_sha256 + + +def test_unknown_task_lists_registered_tasks(tmp_path, capsys): + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=_task_definition(), + ) + tasks = load_tasks(tmp_path) + + with pytest.raises(SystemExit): + run_task_command(["show", "missing"], tasks=tasks) + assert "test-task" in capsys.readouterr().err + + +def test_task_commands_list_show_and_validate(tmp_path, capsys, caplog): + _write_family( + tmp_path, + family="example", + filename="test-task.yaml", + definition=_task_definition(), + ) + tasks = load_tasks(tmp_path) + + run_task_command(["list"], tasks=tasks) + assert capsys.readouterr().out.startswith("test-task\tv1\t") + + run_task_command(["show", "test-task", "--resolved"], tasks=tasks) + shown = yaml.safe_load(capsys.readouterr().out) + assert shown["task"] == "test-task" + assert shown["_provenance"]["resolved_sha256"] + + run_task_command(["validate"], tasks=tasks) + assert "Validated 1 task(s)." in capsys.readouterr().out + + +def test_main_cli_intercepts_task_commands(monkeypatch, capsys): + def unexpected_run_config(_argv): + raise AssertionError("task commands must not construct RunConfig") + + monkeypatch.setattr(cli_module, "build_run_config", unexpected_run_config) + + cli_module.cli(["tasks", "list"]) + + assert "alpaca-eval" in capsys.readouterr().out