From cd2d223772729bddbbd9edcf03e5d63c6598d23a Mon Sep 17 00:00:00 2001 From: bora kargi Date: Mon, 20 Jul 2026 14:28:50 +0200 Subject: [PATCH 1/3] Package MT-Bench task definition Move source pins, baseline, multi-turn policy, reference routing, and FastChat judge defaults into YAML and route the task through registered adapters. --- judgearena/benchmarks/mt_bench/common.py | 15 +- .../benchmarks/mt_bench/fastchat_compat.py | 8 +- .../benchmarks/mt_bench/mt_bench_utils.py | 45 ++- .../benchmarks/mt_bench/preset_judging.py | 8 +- judgearena/benchmarks/mt_bench/runner.py | 37 ++ judgearena/benchmarks/pairwise/baselines.py | 5 +- judgearena/benchmarks/pairwise/runner.py | 19 - judgearena/benchmarks/registry.py | 12 +- judgearena/dataset_revisions.py | 4 +- judgearena/datasets/__init__.py | 5 - judgearena/datasets/mt_bench.py | 345 +++++++++++------- judgearena/datasets/registry.py | 8 +- judgearena/prompts/registry.py | 4 +- .../tasks/definitions/mt_bench/mt-bench.yaml | 60 +++ judgearena/tasks/registry.py | 8 +- judgearena/tasks/schema.py | 55 ++- judgearena/utils/io.py | 4 - tests/test_generate_and_evaluate.py | 15 +- tests/test_instruction_dataset.py | 2 - tests/test_mt_bench_downloads.py | 76 +++- tests/test_mt_bench_fastchat_compat.py | 9 +- tests/test_mt_bench_preset_judging.py | 6 + tests/test_prompt_registry.py | 1 + tests/test_task_registry.py | 15 + tests/test_utils.py | 12 +- 25 files changed, 535 insertions(+), 243 deletions(-) create mode 100644 judgearena/benchmarks/mt_bench/runner.py create mode 100644 judgearena/tasks/definitions/mt_bench/mt-bench.yaml diff --git a/judgearena/benchmarks/mt_bench/common.py b/judgearena/benchmarks/mt_bench/common.py index 4223b90..a272ea1 100644 --- a/judgearena/benchmarks/mt_bench/common.py +++ b/judgearena/benchmarks/mt_bench/common.py @@ -1,22 +1,17 @@ from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Collection, Iterator from dataclasses import dataclass import pandas as pd from judgearena.utils import safe_text -MT_BENCH_REFERENCE_CATEGORIES: set[str] = { - "math", - "reasoning", - "coding", - "arena-hard-200", -} - -def is_reference_based_category(category: str | None) -> bool: - return (category or "") in MT_BENCH_REFERENCE_CATEGORIES +def is_reference_based_category( + category: str | None, reference_categories: Collection[str] +) -> bool: + return (category or "") in reference_categories def resolve_mt_bench_turn_flags(turns_mode: str) -> tuple[bool, bool]: diff --git a/judgearena/benchmarks/mt_bench/fastchat_compat.py b/judgearena/benchmarks/mt_bench/fastchat_compat.py index afac451..69a675d 100644 --- a/judgearena/benchmarks/mt_bench/fastchat_compat.py +++ b/judgearena/benchmarks/mt_bench/fastchat_compat.py @@ -3,6 +3,7 @@ from __future__ import annotations import math +from collections.abc import Collection from dataclasses import dataclass from typing import Any, Literal @@ -220,6 +221,7 @@ def _select_prompt( category: str | None, multi_turn: bool, *, + reference_categories: Collection[str], prompt_preset: str = DEFAULT_JUDGE_PROMPT_PRESET, ) -> FastChatPairwisePrompt: prompt_variants = _FASTCHAT_PROMPT_PRESET_REGISTRY.get(prompt_preset) @@ -228,7 +230,7 @@ def _select_prompt( raise ValueError( f"Unsupported MT-Bench prompt preset '{prompt_preset}'. Choose from: {supported}." ) - needs_ref = is_reference_based_category(category) + needs_ref = is_reference_based_category(category, reference_categories) if needs_ref and multi_turn: return prompt_variants["multi_ref"] if needs_ref: @@ -246,6 +248,7 @@ def _build_fastchat_judge_items( eval_single: bool, eval_multi: bool, truncate_input_chars: int | None, + reference_categories: Collection[str], prompt_preset: str = DEFAULT_JUDGE_PROMPT_PRESET, strip_thinking_before_judging: bool = False, ) -> list[MTBenchJudgeItem]: @@ -259,6 +262,7 @@ def _build_fastchat_judge_items( select_prompt=lambda category, multi_turn: _select_prompt( category, multi_turn=multi_turn, + reference_categories=reference_categories, prompt_preset=prompt_preset, ), strip_thinking_before_judging=strip_thinking_before_judging, @@ -345,6 +349,7 @@ def judge_mt_bench_pairwise_fastchat( swap_mode: str, truncate_input_chars: int | None, use_tqdm: bool, + reference_categories: Collection[str], prompt_preset: str = DEFAULT_JUDGE_PROMPT_PRESET, strip_thinking_before_judging: bool = False, ) -> tuple[pd.Series, list[dict[str, Any]], list[dict[str, object]], int]: @@ -359,6 +364,7 @@ def judge_mt_bench_pairwise_fastchat( eval_single=eval_single, eval_multi=eval_multi, truncate_input_chars=truncate_input_chars, + reference_categories=reference_categories, prompt_preset=prompt_preset, strip_thinking_before_judging=strip_thinking_before_judging, ) diff --git a/judgearena/benchmarks/mt_bench/mt_bench_utils.py b/judgearena/benchmarks/mt_bench/mt_bench_utils.py index c62b474..fa16dd5 100644 --- a/judgearena/benchmarks/mt_bench/mt_bench_utils.py +++ b/judgearena/benchmarks/mt_bench/mt_bench_utils.py @@ -15,23 +15,20 @@ from judgearena.artifacts import prepare_run_directory, write_run_metadata_safely from judgearena.benchmarks.mt_bench.fastchat_compat import ( - FASTCHAT_TEMPERATURE_CONFIG, judge_mt_bench_pairwise_fastchat, ) from judgearena.benchmarks.mt_bench.preset_judging import judge_mt_bench_with_preset +from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline from judgearena.datasets import load_instructions from judgearena.datasets.mt_bench import ( load_mt_bench_model_answers, - mt_bench_native_baseline, ) from judgearena.generate import generate_multiturn from judgearena.log import get_logger from judgearena.models import is_thinking_model, make_model -from judgearena.prompts.registry import ( - DEFAULT_JUDGE_PROMPT_PRESET, - ResolvedJudgePrompt, - resolve_run_judge_prompt, -) +from judgearena.prompts.registry import ResolvedJudgePrompt, resolve_run_judge_prompt +from judgearena.tasks.registry import get_packaged_task +from judgearena.tasks.schema import MTBenchProtocol from judgearena.utils import ( cache_function_dataframe, compute_pref_summary, @@ -45,6 +42,13 @@ from judgearena.config import RunConfig +def _task_protocol(task_id: str) -> MTBenchProtocol: + task = get_packaged_task(task_id) + if task is None or not isinstance(task.spec.protocol, MTBenchProtocol): + raise ValueError(f"Task {task_id!r} does not define an MT-Bench protocol.") + return task.spec.protocol + + def _align_mt_bench_completions( *, questions_df: pd.DataFrame, completions: pd.DataFrame, model_name: str ) -> pd.DataFrame: @@ -89,7 +93,8 @@ def _generate_mt_bench_completions( questions_df: pd.DataFrame, ignore_cache: bool, ) -> tuple[pd.DataFrame, pd.DataFrame]: - cache_prefix = "mt-bench" + cache_prefix = cfg.task + protocol = _task_protocol(cfg.task) def _run_generation( model_name: str, *, generation_kwargs: dict[str, object] @@ -98,7 +103,9 @@ def _run_generation( # not explicitly pinned a per-role temperature; otherwise the config # override should win for reproducibility. temperature_config = ( - None if "temperature" in generation_kwargs else FASTCHAT_TEMPERATURE_CONFIG + None + if "temperature" in generation_kwargs + else dict(protocol.generation.category_temperatures) ) return generate_multiturn( questions=questions_df, @@ -259,6 +266,7 @@ def _run_mt_bench_fastchat( fastchat_prompt_preset: str, started_at_utc: datetime, ) -> pd.Series: + protocol = _task_protocol(cfg.task) prefs, annotations, combined_metadata, num_inconsistent = ( judge_mt_bench_pairwise_fastchat( judge_chat_model=judge_chat_model, @@ -268,10 +276,11 @@ def _run_mt_bench_fastchat( completions_b=completions_b, model_a=cfg.model.name, model_b=cfg.model.baseline, - turns_mode="both", + turns_mode=protocol.judge.turns_mode, swap_mode=cfg.judge.swap_mode, truncate_input_chars=cfg.generation.truncate_judge_input_chars, use_tqdm=cfg.run.use_tqdm, + reference_categories=protocol.judge.reference_categories, prompt_preset=fastchat_prompt_preset, strip_thinking_before_judging=cfg.judge.strip_thinking_before_judging, ) @@ -304,6 +313,7 @@ def _run_mt_bench_preset( resolved_prompt: ResolvedJudgePrompt, started_at_utc: datetime, ) -> pd.Series: + protocol = _task_protocol(cfg.task) prefs, annotations, combined_metadata = judge_mt_bench_with_preset( judge_chat_model=judge_chat_model, judge_model=cfg.judge.model, @@ -312,10 +322,11 @@ def _run_mt_bench_preset( completions_b=completions_b, model_a=cfg.model.name, model_b=cfg.model.baseline, - turns_mode="both", + turns_mode=protocol.judge.turns_mode, swap_mode=cfg.judge.swap_mode, truncate_input_chars=cfg.generation.truncate_judge_input_chars, use_tqdm=cfg.run.use_tqdm, + reference_categories=protocol.judge.reference_categories, prompt_preset=cfg.judge.prompt_preset or resolved_prompt.preset_name, provide_explanation=cfg.judge.provide_explanation, system_file=cfg.judge.system_prompt_file, @@ -346,15 +357,17 @@ def run_mt_bench( ): """MT-Bench pipeline with preset or FastChat-original pairwise judging.""" run_started_at = datetime.now(UTC) + protocol = _task_protocol(cfg.task) if cfg.model.baseline is None: - cfg.model.baseline = mt_bench_native_baseline(cfg.task) + baseline = native_pairwise_baseline(cfg.task) + cfg.model.baseline = baseline if isinstance(baseline, str) else None if cfg.model.baseline is None: raise ValueError( f"--model_B is required for dataset '{cfg.task}'; " "no dataset-native baseline registered." ) questions_df = load_instructions( - "mt-bench", n_instructions=cfg.generation.n_instructions + cfg.task, n_instructions=cfg.generation.n_instructions ) logger.info( "Generating multi-turn completions for MT-Bench with %s and %s.", @@ -377,7 +390,9 @@ def run_mt_bench( fallback_chat_template=cfg.model.chat_template, ) if resolved_prompt.delegated and cfg.judge.temperature is None: - judge_model_kwargs.setdefault("temperature", 0.0) + judge_model_kwargs.setdefault( + "temperature", protocol.judge.fastchat_temperature + ) judge_chat_model = make_model(model=cfg.judge.model, **judge_model_kwargs) if resolved_prompt.delegated: return _run_mt_bench_fastchat( @@ -389,7 +404,7 @@ def run_mt_bench( completions_b=completions_b, judge_chat_model=judge_chat_model, resolved_prompt=resolved_prompt, - fastchat_prompt_preset=DEFAULT_JUDGE_PROMPT_PRESET, + fastchat_prompt_preset=protocol.judge.fastchat_prompt_preset, started_at_utc=run_started_at, ) return _run_mt_bench_preset( diff --git a/judgearena/benchmarks/mt_bench/preset_judging.py b/judgearena/benchmarks/mt_bench/preset_judging.py index 412fc5d..907b68f 100644 --- a/judgearena/benchmarks/mt_bench/preset_judging.py +++ b/judgearena/benchmarks/mt_bench/preset_judging.py @@ -1,6 +1,7 @@ from __future__ import annotations import math +from collections.abc import Collection from dataclasses import dataclass from typing import Any @@ -63,12 +64,13 @@ def _select_preset_prompt( category: str | None, multi_turn: bool, *, + reference_categories: Collection[str], prompt_preset: str = DEFAULT_JUDGE_PROMPT_PRESET, provide_explanation: bool, system_file: str | None = None, user_file: str | None = None, ) -> MTBenchPresetPrompt: - ref_based = is_reference_based_category(category) + ref_based = is_reference_based_category(category, reference_categories) resolved_prompt = resolve_judge_prompt( preset=prompt_preset, system_file=system_file, @@ -107,6 +109,7 @@ def _build_mt_bench_preset_items( eval_single: bool, eval_multi: bool, truncate_input_chars: int | None, + reference_categories: Collection[str], prompt_preset: str, provide_explanation: bool, system_file: str | None = None, @@ -123,6 +126,7 @@ def _build_mt_bench_preset_items( select_prompt=lambda category, multi_turn: _select_preset_prompt( category, multi_turn=multi_turn, + reference_categories=reference_categories, prompt_preset=prompt_preset, provide_explanation=provide_explanation, system_file=system_file, @@ -151,6 +155,7 @@ def judge_mt_bench_with_preset( swap_mode: str, truncate_input_chars: int | None, use_tqdm: bool, + reference_categories: Collection[str], prompt_preset: str = DEFAULT_JUDGE_PROMPT_PRESET, provide_explanation: bool = False, system_file: str | None = None, @@ -167,6 +172,7 @@ def judge_mt_bench_with_preset( eval_single=eval_single, eval_multi=eval_multi, truncate_input_chars=truncate_input_chars, + reference_categories=reference_categories, prompt_preset=prompt_preset, provide_explanation=provide_explanation, system_file=system_file, diff --git a/judgearena/benchmarks/mt_bench/runner.py b/judgearena/benchmarks/mt_bench/runner.py new file mode 100644 index 0000000..e5fd471 --- /dev/null +++ b/judgearena/benchmarks/mt_bench/runner.py @@ -0,0 +1,37 @@ +"""Registered entry point for the specialized MT-Bench pipeline.""" + +from __future__ import annotations + +from datetime import UTC, datetime +from pathlib import Path +from typing import TYPE_CHECKING + +from judgearena.artifacts import prepare_run_directory +from judgearena.benchmarks.mt_bench.mt_bench_utils import run_mt_bench +from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline + +if TYPE_CHECKING: + from judgearena.config import RunConfig + + +def run_mt_bench_benchmark(cfg: RunConfig): + """Prepare one run directory and execute the YAML-selected MT-Bench runner.""" + baseline = cfg.model.baseline or native_pairwise_baseline(cfg.task) + if not isinstance(baseline, str): + raise ValueError("MT-Bench requires a flat native baseline.") + + result_name = ( + f"{cfg.task}-{cfg.model.name}-{baseline}-{cfg.judge.model}-" + f"{cfg.judge.swap_mode}" + ).replace("/", "_") + run_timestamp = datetime.now(UTC).strftime("%Y%m%d_%H%M%S") + result_folder = prepare_run_directory( + cfg, + Path(cfg.run.result_folder) / f"{result_name}-{run_timestamp}", + ) + return run_mt_bench( + cfg, + cfg.run.ignore_cache, + res_folder=result_folder, + result_name=result_name, + ) diff --git a/judgearena/benchmarks/pairwise/baselines.py b/judgearena/benchmarks/pairwise/baselines.py index 1f78693..829c00e 100644 --- a/judgearena/benchmarks/pairwise/baselines.py +++ b/judgearena/benchmarks/pairwise/baselines.py @@ -4,13 +4,10 @@ from collections.abc import Mapping -from judgearena.datasets.mt_bench import MT_BENCH_BASELINES from judgearena.tasks.registry import get_packaged_task from judgearena.tasks.schema import CategoryDefaultsBaseline, TaskDefaultBaseline -LEGACY_PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = { - **MT_BENCH_BASELINES, -} +LEGACY_PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = {} def native_pairwise_baseline(task: str) -> str | Mapping[str, str] | None: diff --git a/judgearena/benchmarks/pairwise/runner.py b/judgearena/benchmarks/pairwise/runner.py index 7330337..09937f4 100644 --- a/judgearena/benchmarks/pairwise/runner.py +++ b/judgearena/benchmarks/pairwise/runner.py @@ -13,7 +13,6 @@ from judgearena.artifacts import prepare_run_directory, write_run_metadata_safely from judgearena.benchmarks.execution import build_generation_kwargs, build_judge -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.fluency import is_fluency_task as task_is_fluency @@ -180,24 +179,6 @@ def run_pairwise(cfg: "RunConfig"): # set_langchain_cache() ignore_cache = cfg.run.ignore_cache - if cfg.task == "mt-bench": - model_b = cfg.model.baseline or native_pairwise_baseline(cfg.task) - if not isinstance(model_b, str): - raise ValueError("MT-Bench requires a flat native baseline.") - name = f"{cfg.task}-{cfg.model.name}-{model_b}-{cfg.judge.model}" - name += f"-{cfg.judge.swap_mode}" - name = name.replace("/", "_") - run_ts = run_started_at.strftime("%Y%m%d_%H%M%S") - res_folder = prepare_run_directory( - cfg, Path(cfg.run.result_folder) / f"{name}-{run_ts}" - ) - return run_mt_bench( - cfg, - ignore_cache, - res_folder=res_folder, - result_name=name, - ) - # Currrently, we run context evaluation is_fluency_task = task_is_fluency(cfg.task) if is_fluency_task: diff --git a/judgearena/benchmarks/registry.py b/judgearena/benchmarks/registry.py index ef2e325..7003910 100644 --- a/judgearena/benchmarks/registry.py +++ b/judgearena/benchmarks/registry.py @@ -34,10 +34,14 @@ def supports(self, task: str) -> bool: def benchmark_adapters() -> tuple[BenchmarkAdapter, ...]: - """Return the registered benchmark implementations.""" + """Return registered benchmark implementations, specific first.""" + from judgearena.benchmarks.mt_bench.runner import run_mt_bench_benchmark from judgearena.benchmarks.pairwise.runner import run_pairwise - return (BenchmarkAdapter("pairwise", None, run_pairwise),) + return ( + BenchmarkAdapter("mt_bench", frozenset(), run_mt_bench_benchmark), + BenchmarkAdapter("pairwise", None, run_pairwise), + ) def resolve_benchmark_adapter(task: str) -> BenchmarkAdapter: @@ -49,9 +53,7 @@ def resolve_benchmark_adapter(task: str) -> BenchmarkAdapter: for adapter in adapters: if adapter.name == runner_id: return adapter - raise ValueError( - f"Task {task!r} selects unavailable runner {runner_id!r}." - ) + raise ValueError(f"Task {task!r} selects unavailable runner {runner_id!r}.") for adapter in adapters: if adapter.supports(task): diff --git a/judgearena/dataset_revisions.py b/judgearena/dataset_revisions.py index 3d13f48..481124c 100644 --- a/judgearena/dataset_revisions.py +++ b/judgearena/dataset_revisions.py @@ -34,9 +34,7 @@ # Raw-URL pins (e.g. FastChat reference answers fetched as a raw GitHub URL). # Mapping is "logical name" -> commit SHA on the upstream repo. The downloader # rewrites the URL to point at the pinned SHA. -RAW_URL_REVISIONS: dict[str, str | None] = { - "lm-sys/FastChat": "587d5cfa1609a43d192cedb8441cac3c17db105d", -} +RAW_URL_REVISIONS: dict[str, str | None] = {} def hf_revision(repo_id: str) -> str | None: diff --git a/judgearena/datasets/__init__.py b/judgearena/datasets/__init__.py index 0bfd674..961f255 100644 --- a/judgearena/datasets/__init__.py +++ b/judgearena/datasets/__init__.py @@ -17,11 +17,6 @@ def load_instructions(dataset: str, n_instructions: int | None = None) -> pd.Dat 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() - else: raise ValueError(f"Unsupported instruction dataset {dataset!r}.") diff --git a/judgearena/datasets/mt_bench.py b/judgearena/datasets/mt_bench.py index 137af0f..878cc5d 100644 --- a/judgearena/datasets/mt_bench.py +++ b/judgearena/datasets/mt_bench.py @@ -1,46 +1,69 @@ +"""Dataset adapter for the YAML-defined MT-Bench task.""" + +from __future__ import annotations + import warnings +from collections.abc import Mapping from pathlib import Path from urllib.request import urlretrieve import pandas as pd from huggingface_hub import snapshot_download -from judgearena.dataset_revisions import RAW_URL_REVISIONS, hf_revision from judgearena.paths import data_root +from judgearena.tasks.registry import get_packaged_task +from judgearena.tasks.schema import ( + GitRawSource, + HuggingFaceSpaceSource, + ResolvedTaskSpec, + TaskDefaultBaseline, +) -MT_BENCH_SPACE_ID = "lmsys/mt-bench" -MT_BENCH_QUESTION_PATTERN = "data/mt_bench/question.jsonl" -MT_BENCH_MODEL_ANSWER_DIR = Path("data") / "mt_bench" / "model_answer" +def _task(task_id: str = "mt-bench") -> ResolvedTaskSpec: + task = get_packaged_task(task_id) + if task is None or task.spec.dataset.adapter != "mt_bench": + raise ValueError(f"Unsupported MT-Bench task: {task_id!r}.") + return task -def _fastchat_reference_url() -> str: - """URL for FastChat MT-Bench GPT-4 references, pinned when available.""" - revision = RAW_URL_REVISIONS.get("lm-sys/FastChat") - rev = revision if revision else "main" - return ( - f"https://raw.githubusercontent.com/lm-sys/FastChat/{rev}/" - "fastchat/llm_judge/data/mt_bench/reference_answer/gpt-4.jsonl" - ) +def is_mt_bench_dataset(dataset: str) -> bool: + task = get_packaged_task(dataset) + return task is not None and task.spec.dataset.adapter == "mt_bench" -FASTCHAT_GPT4_REFERENCE_URL = _fastchat_reference_url() -# Mirrors ``ARENA_HARD_BASELINES`` / ``M_ARENA_HARD_BASELINES``: dataset name -> -# dataset-native pairwise baseline. MT-Bench ships only one variant today, and -# ``gpt-4`` is the stronger-reference choice (FastChat's own ``pairwise-baseline`` -# default is ``gpt-3.5-turbo``; we deliberately diverge here). -MT_BENCH_BASELINES: dict[str, str] = { - "mt-bench": "gpt-4", -} +def mt_bench_native_baseline( + dataset: str, +) -> str | Mapping[str, str] | None: + """Return the task-defined MT-Bench baseline.""" + if not is_mt_bench_dataset(dataset): + return None + from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline + return native_pairwise_baseline(dataset) -def is_mt_bench_dataset(dataset: str) -> bool: - return dataset in MT_BENCH_BASELINES + +def _space_source(task: ResolvedTaskSpec) -> HuggingFaceSpaceSource: + source = task.spec.dataset.sources.get("benchmark") + if not isinstance(source, HuggingFaceSpaceSource): + raise ValueError( + f"Task {task.task!r} must define a Hugging Face Space source " + "named 'benchmark'." + ) + return source + + +def _reference_source(task: ResolvedTaskSpec) -> GitRawSource: + source = task.spec.dataset.sources.get("references") + if not isinstance(source, GitRawSource): + raise ValueError( + f"Task {task.task!r} must define a Git raw source named 'references'." + ) + return source -def mt_bench_native_baseline(dataset: str) -> str | None: - """Baseline for a dataset name, or ``None`` if it isn't mt-bench.""" - return MT_BENCH_BASELINES.get(dataset) +def _task_cache_dir(task: ResolvedTaskSpec, local_tables_path: Path) -> Path: + return local_tables_path / "_sources" / task.definition_task def _normalize_question_id(question_id: object) -> object: @@ -52,98 +75,132 @@ def _normalize_question_id(question_id: object) -> object: def _snapshot_mt_bench_files( *, + task: ResolvedTaskSpec, local_dir: Path, allow_patterns: list[str], expected_path: Path, description: str, ) -> None: + source = _space_source(task) try: snapshot_download( - repo_id=MT_BENCH_SPACE_ID, + repo_id=source.repo_id, repo_type="space", allow_patterns=allow_patterns, local_dir=local_dir, force_download=False, - revision=hf_revision(MT_BENCH_SPACE_ID), + revision=source.revision, ) - except Exception as e: + except Exception as exc: raise RuntimeError( - f"Failed to download {description} from HuggingFace space " - f"'{MT_BENCH_SPACE_ID}'. If you're in an offline / restricted-network " - f"environment, pre-download the space snapshot and place the file at " - f"{expected_path}, or set OPENJURY_DATA to point to that directory." - ) from e + f"Failed to download {description} from Hugging Face Space " + f"{source.repo_id!r}. If you are offline, place the file at " + f"{expected_path}." + ) from exc if not expected_path.exists(): raise FileNotFoundError( - f"Could not locate {description} after download. " - f"Expected file at {expected_path}." + f"Could not locate {description} after download. Expected {expected_path}." + ) + + +def _git_raw_url(source: GitRawSource) -> str: + repository = source.repository.rstrip("/") + github_prefix = "https://github.com/" + if repository.startswith(github_prefix): + project = repository.removeprefix(github_prefix) + return ( + f"https://raw.githubusercontent.com/{project}/{source.revision}/" + f"{source.path}" ) + return f"{repository}/raw/{source.revision}/{source.path}" -def _download_gpt4_references(local_dir: Path) -> Path | None: +def _download_references(task: ResolvedTaskSpec, local_dir: Path) -> Path | None: reference_dir = local_dir / "reference_answer" reference_dir.mkdir(parents=True, exist_ok=True) - gpt4_reference_path = reference_dir / "gpt-4.jsonl" - if gpt4_reference_path.exists(): - return gpt4_reference_path + reference_path = reference_dir / "gpt-4.jsonl" + if reference_path.exists(): + return reference_path + source = _reference_source(task) try: - urlretrieve(FASTCHAT_GPT4_REFERENCE_URL, gpt4_reference_path) - except Exception as e: + urlretrieve(_git_raw_url(source), reference_path) + except Exception as exc: warnings.warn( - "Could not download MT-Bench GPT-4 reference answers from FastChat. " - f"Falling back to inline references from question.jsonl: {e}", + "Could not download MT-Bench GPT-4 reference answers. Falling back " + f"to inline references from question.jsonl: {exc}", RuntimeWarning, stacklevel=2, ) return None - return gpt4_reference_path + return reference_path -def download_mt_bench(local_dir: Path | None = None) -> tuple[Path, Path | None]: - """Download MT-Bench questions and GPT-4 references if missing.""" - if local_dir is None: - local_dir = data_root / "mt-bench" +def _download_mt_bench( + task: ResolvedTaskSpec, local_dir: Path +) -> tuple[Path, Path | None]: try: local_dir.mkdir(parents=True, exist_ok=True) - except PermissionError as e: + except PermissionError as exc: raise PermissionError( - f"Cannot create MT-Bench cache directory at {local_dir}. " - "Set environment variable OPENJURY_DATA to a writable location." - ) from e + f"Cannot create MT-Bench cache directory at {local_dir}. Set " + "JUDGEARENA_DATA to a writable location." + ) from exc question_path = local_dir / "data" / "mt_bench" / "question.jsonl" if not question_path.exists(): _snapshot_mt_bench_files( + task=task, local_dir=local_dir, - allow_patterns=[MT_BENCH_QUESTION_PATTERN], + allow_patterns=[question_path.relative_to(local_dir).as_posix()], expected_path=question_path, description="MT-Bench questions", ) + return question_path, _download_references(task, local_dir) + - gpt4_reference_path = _download_gpt4_references(local_dir) - return question_path, gpt4_reference_path +def download_mt_bench(local_dir: Path | None = None) -> tuple[Path, Path | None]: + """Compatibility wrapper downloading the registered MT-Bench sources.""" + return _download_mt_bench(_task(), local_dir or data_root / "mt-bench") def download_mt_bench_model_answer( - model_id: str, local_dir: Path | None = None + model_id: str, + local_dir: Path | None = None, + *, + task: ResolvedTaskSpec | None = None, ) -> Path: - """Download a cached MT-Bench baseline answer file if missing.""" - if local_dir is None: - local_dir = data_root / "mt-bench" - answer_path = local_dir / MT_BENCH_MODEL_ANSWER_DIR / f"{model_id}.jsonl" + """Download a cached MT-Bench model-answer file if missing.""" + resolved = task or _task() + root = local_dir or data_root / "mt-bench" + answer_path = root / "data" / "mt_bench" / "model_answer" / f"{model_id}.jsonl" if answer_path.exists(): return answer_path answer_path.parent.mkdir(parents=True, exist_ok=True) - allow_pattern = (MT_BENCH_MODEL_ANSWER_DIR / f"{model_id}.jsonl").as_posix() _snapshot_mt_bench_files( - local_dir=local_dir, - allow_patterns=[allow_pattern], + task=resolved, + local_dir=root, + allow_patterns=[answer_path.relative_to(root).as_posix()], expected_path=answer_path, - description=f"MT-Bench model answers for '{model_id}'", + description=f"MT-Bench model answers for {model_id!r}", ) return answer_path +def download_task_sources(task: ResolvedTaskSpec, local_tables_path: Path) -> None: + """Download every source required by the registered MT-Bench task.""" + if task.spec.dataset.adapter != "mt_bench": + raise ValueError(f"Task {task.task!r} does not use the MT-Bench adapter.") + local_dir = _task_cache_dir(task, local_tables_path) + _download_mt_bench(task, local_dir) + baseline = task.spec.protocol.baseline + if isinstance(baseline, TaskDefaultBaseline): + download_mt_bench_model_answer( + baseline.reference_id, + local_dir=local_dir, + task=task, + ) + + def _extract_answer_turns(record: dict, source_name: str) -> tuple[object, list[str]]: question_id = record.get("question_id", record.get("id")) if question_id is None: @@ -153,20 +210,20 @@ def _extract_answer_turns(record: dict, source_name: str) -> tuple[object, list[ choices = record.get("choices") if not (isinstance(choices, list) and choices): raise ValueError( - f"MT-Bench answer record for question {question_id} in {source_name} is " - "missing a non-empty choices list." + f"MT-Bench answer record for question {question_id} in {source_name} " + "is missing a non-empty choices list." ) first_choice = choices[0] if not isinstance(first_choice, dict): raise ValueError( - f"MT-Bench answer record for question {question_id} in {source_name} has " - "a malformed first choice entry." + f"MT-Bench answer record for question {question_id} in {source_name} " + "has a malformed first choice entry." ) turns = first_choice.get("turns") if not isinstance(turns, list): raise ValueError( - f"MT-Bench answer record for question {question_id} in {source_name} is " - "missing a turns list." + f"MT-Bench answer record for question {question_id} in {source_name} " + "is missing a turns list." ) return _normalize_question_id(question_id), turns @@ -175,127 +232,133 @@ def load_mt_bench_model_answers( model: str, n_instructions: int | None = None, local_dir: Path | None = None, + *, + task: ResolvedTaskSpec | None = None, ) -> pd.DataFrame | None: - """Load pre-generated MT-Bench answers from a local file or cached model id.""" + """Load pre-generated MT-Bench answers from a path or cached model ID.""" local_path = Path(model) if local_path.exists(): answer_path = local_path elif "/" not in model: answer_path = download_mt_bench_model_answer( - model_id=model, local_dir=local_dir + model_id=model, + local_dir=local_dir, + task=task, ) else: return None answer_records = pd.read_json(answer_path, lines=True).to_dict(orient="records") rows = [] - for rec in answer_records: - question_id, turns = _extract_answer_turns(rec, str(answer_path)) + for record in answer_records: + question_id, turns = _extract_answer_turns(record, str(answer_path)) rows.append( { "instruction_index": question_id, - "completion_turn_1": turns[0] if len(turns) > 0 else "", + "completion_turn_1": turns[0] if turns else "", "completion_turn_2": turns[1] if len(turns) > 1 else "", } ) df_answers = pd.DataFrame(rows) if df_answers.empty: - raise ValueError( - f"MT-Bench answer file {answer_path} did not contain any rows." - ) + raise ValueError(f"MT-Bench answer file {answer_path} contained no rows.") df_answers.sort_values("instruction_index", inplace=True) - if n_instructions is not None: - df_answers = df_answers.head(n_instructions) - return df_answers - + return df_answers.head(n_instructions) if n_instructions is not None else df_answers -def load_mt_bench() -> pd.DataFrame: - """Load MT-Bench questions and reference answers. - - Downloads MT-Bench questions from the HuggingFace LMSYS space and tries to - load GPT-4 references from FastChat GitHub. If GPT-4 references cannot be - downloaded or parsed, falls back to inline references from question.jsonl. - """ - question_path, ref_path = download_mt_bench() +def _load_mt_bench(task: ResolvedTaskSpec, local_dir: Path) -> pd.DataFrame: + question_path, reference_path = _download_mt_bench(task, local_dir) questions = pd.read_json(question_path, lines=True).to_dict(orient="records") - ref_by_id: dict[int | str, list[str]] = {} - use_inline_reference_fallback = ref_path is None - if ref_path is not None: + references_by_id: dict[int | str, list[str]] = {} + use_inline_references = reference_path is None + if reference_path is not None: try: - reference_records = pd.read_json(ref_path, lines=True).to_dict( - orient="records" - ) - for rec in reference_records: - qid = rec.get("question_id", rec.get("id")) - if qid is None: - continue - choices = rec.get("choices") - if not (isinstance(choices, list) and choices): + records = pd.read_json(reference_path, lines=True).to_dict(orient="records") + for record in records: + question_id = record.get("question_id", record.get("id")) + choices = record.get("choices") + if question_id is None or not (isinstance(choices, list) and choices): continue first_choice = choices[0] - if not isinstance(first_choice, dict): - continue - turns = first_choice.get("turns") + turns = ( + first_choice.get("turns") + if isinstance(first_choice, dict) + else None + ) if not isinstance(turns, list): continue - ref_by_id[qid] = turns - try: - ref_by_id[int(qid)] = turns - except Exception: - pass - except Exception as e: + references_by_id[question_id] = turns + references_by_id[_normalize_question_id(question_id)] = turns + except Exception as exc: warnings.warn( - "Failed to parse GPT-4 reference answers from FastChat. " - f"Falling back to inline references from question.jsonl: {e}", + "Failed to parse MT-Bench GPT-4 references. Falling back to " + f"inline references from question.jsonl: {exc}", RuntimeWarning, stacklevel=2, ) - use_inline_reference_fallback = True + use_inline_references = True + fields = task.spec.dataset.fields rows = [] - for rec in questions: - qid_raw = rec.get("question_id", rec.get("id")) - if qid_raw is None: + for record in questions: + question_id_raw = record.get(fields.id, record.get("id")) + if question_id_raw is None: raise ValueError( - f"MT-Bench question record missing question_id/id: keys={list(rec.keys())}" + f"MT-Bench question is missing field {fields.id!r}: keys={list(record)}" ) - qid = _normalize_question_id(qid_raw) - - category = rec.get("category") - turns = rec.get("turns") + question_id = _normalize_question_id(question_id_raw) + turns = record.get(fields.instruction) if isinstance(turns, list): - turn_1 = turns[0] if len(turns) > 0 else None + turn_1 = turns[0] if turns else None turn_2 = turns[1] if len(turns) > 1 else None else: - turn_1 = rec.get("turn_1", rec.get("instruction")) - turn_2 = rec.get("turn_2") + turn_1 = turns + turn_2 = record.get("turn_2") - ref_turns = ref_by_id.get(qid_raw) or ref_by_id.get(qid) - if ref_turns is None and use_inline_reference_fallback: - inline_ref = rec.get("reference") - if isinstance(inline_ref, list): - ref_turns = inline_ref - - ref_turn_1 = ( - ref_turns[0] if isinstance(ref_turns, list) and len(ref_turns) > 0 else None - ) - ref_turn_2 = ( - ref_turns[1] if isinstance(ref_turns, list) and len(ref_turns) > 1 else None + reference_turns = references_by_id.get(question_id_raw) or references_by_id.get( + question_id ) + if reference_turns is None and use_inline_references: + inline_reference = record.get("reference") + if isinstance(inline_reference, list): + reference_turns = inline_reference rows.append( { - "instruction_index": qid, - "category": category, + "instruction_index": question_id, + "category": ( + record.get(fields.category) if fields.category is not None else None + ), "turn_1": turn_1, "turn_2": turn_2, - "reference_turn_1": ref_turn_1, - "reference_turn_2": ref_turn_2, + "reference_turn_1": (reference_turns[0] if reference_turns else None), + "reference_turn_2": ( + reference_turns[1] + if reference_turns is not None and len(reference_turns) > 1 + else None + ), "instruction": turn_1, } ) - return pd.DataFrame(rows) + + +def load_task_instructions( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame: + """Load normalized MT-Bench questions through the dataset registry.""" + return _load_mt_bench(task, _task_cache_dir(task, local_tables_path)) + + +def load_task_model_outputs( + task: ResolvedTaskSpec, local_tables_path: Path +) -> pd.DataFrame | None: + """MT-Bench loads two-turn model answers through its specialized runner.""" + return None + + +def load_mt_bench() -> pd.DataFrame: + """Compatibility wrapper loading the registered MT-Bench task.""" + return _load_mt_bench(_task(), data_root / "mt-bench") diff --git a/judgearena/datasets/registry.py b/judgearena/datasets/registry.py index 0f24d77..9c9016f 100644 --- a/judgearena/datasets/registry.py +++ b/judgearena/datasets/registry.py @@ -24,7 +24,7 @@ class DatasetAdapter: def dataset_adapters() -> tuple[DatasetAdapter, ...]: """Return registered dataset implementations.""" - from judgearena.datasets import arena_hard, judgearena_tables, m_arenahard + from judgearena.datasets import arena_hard, judgearena_tables, m_arenahard, mt_bench return ( DatasetAdapter( @@ -45,6 +45,12 @@ def dataset_adapters() -> tuple[DatasetAdapter, ...]: m_arenahard.load_task_instructions, m_arenahard.load_task_model_outputs, ), + DatasetAdapter( + "mt_bench", + mt_bench.download_task_sources, + mt_bench.load_task_instructions, + mt_bench.load_task_model_outputs, + ), ) diff --git a/judgearena/prompts/registry.py b/judgearena/prompts/registry.py index 0342bc4..2c6698a 100644 --- a/judgearena/prompts/registry.py +++ b/judgearena/prompts/registry.py @@ -92,9 +92,7 @@ def metadata(self) -> dict[str, str | bool | None]: JUDGE_PROMPT_PRESETS = tuple(PRESETS) -TASK_DEFAULT_PRESET: dict[str, str] = { - "mt-bench": FASTCHAT_PAIRWISE_PROMPT_PRESET, -} +TASK_DEFAULT_PRESET: dict[str, str] = {} def default_preset_for_task(task: str | None) -> str: diff --git a/judgearena/tasks/definitions/mt_bench/mt-bench.yaml b/judgearena/tasks/definitions/mt_bench/mt-bench.yaml new file mode 100644 index 0000000..099ee12 --- /dev/null +++ b/judgearena/tasks/definitions/mt_bench/mt-bench.yaml @@ -0,0 +1,60 @@ +schema_version: 1 +task: mt-bench +task_version: 1 +description: FastChat-compatible pairwise evaluation on the MT-Bench question set. +tags: [pairwise, instruction-following, multi-turn, mt-bench] + +dataset: + adapter: mt_bench + sources: + benchmark: + type: huggingface_space + repo_id: lmsys/mt-bench + revision: "a4b674ca573c24143824ac7f60d9173e7081e37d" + allow_patterns: + - data/mt_bench/question.jsonl + - data/mt_bench/model_answer/gpt-4.jsonl + references: + type: git_raw + repository: https://github.com/lm-sys/FastChat + revision: "587d5cfa1609a43d192cedb8441cac3c17db105d" + path: fastchat/llm_judge/data/mt_bench/reference_answer/gpt-4.jsonl + fields: + id: question_id + instruction: turns + category: category + +protocol: + runner: mt_bench + generation: + mode: multi_turn_chat + category_temperatures: + writing: 0.7 + roleplay: 0.7 + extraction: 0.0 + math: 0.0 + coding: 0.0 + reasoning: 0.0 + stem: 0.1 + humanities: 0.1 + arena-hard-200: 0.0 + baseline: + strategy: task_default + reference_id: gpt-4 + allow_runtime_override: true + judge: + default_prompt: fastchat-pairwise + parser: fastchat_pairwise_verdict + default_swap_mode: fixed + allowed_swap_modes: [fixed, both] + turns_mode: both + fastchat_prompt_preset: default + fastchat_temperature: 0.0 + reference_categories: [math, reasoning, coding, arena-hard-200] + scoring: + adapter: pairwise_win_rate + primary_metric: winrate + higher_is_better: true + +metadata: + reference_implementation: https://github.com/lm-sys/FastChat/tree/master/fastchat/llm_judge diff --git a/judgearena/tasks/registry.py b/judgearena/tasks/registry.py index cbb02e9..b13d526 100644 --- a/judgearena/tasks/registry.py +++ b/judgearena/tasks/registry.py @@ -16,12 +16,14 @@ class AdapterCatalog: """Component IDs that task YAML files may reference.""" - runners: frozenset[str] = frozenset({"pairwise"}) + runners: frozenset[str] = frozenset({"mt_bench", "pairwise"}) datasets: frozenset[str] = frozenset( - {"arena_hard", "judgearena_tables", "m_arena_hard"} + {"arena_hard", "judgearena_tables", "m_arena_hard", "mt_bench"} ) prompts: frozenset[str] = frozenset(JUDGE_PROMPT_PRESETS) - parsers: frozenset[str] = frozenset({"pairwise_preference"}) + parsers: frozenset[str] = frozenset( + {"fastchat_pairwise_verdict", "pairwise_preference"} + ) scorers: frozenset[str] = frozenset({"pairwise_win_rate"}) diff --git a/judgearena/tasks/schema.py b/judgearena/tasks/schema.py index 46d65a8..637361a 100644 --- a/judgearena/tasks/schema.py +++ b/judgearena/tasks/schema.py @@ -68,6 +68,24 @@ class SingleTurnGeneration(_StrictFrozenModel): mode: Literal["single_turn_chat"] +class MultiTurnGeneration(_StrictFrozenModel): + mode: Literal["multi_turn_chat"] + category_temperatures: dict[str, float] = Field(default_factory=dict) + + @model_validator(mode="after") + def _validate_temperatures(self) -> MultiTurnGeneration: + invalid = { + category: temperature + for category, temperature in self.category_temperatures.items() + if not category or temperature < 0 + } + if invalid: + raise ValueError( + f"category temperatures require names and non-negative values: {invalid}" + ) + return self + + class NoBaseline(_StrictFrozenModel): strategy: Literal["none"] @@ -140,6 +158,37 @@ class PairwiseProtocol(_StrictFrozenModel): scoring: ScoringSpec +class MTBenchJudgeSpec(PairwiseJudgeSpec): + turns_mode: Literal["both", "single", "multi"] = "both" + fastchat_prompt_preset: str = Field(min_length=1) + fastchat_temperature: float = Field(ge=0) + reference_categories: tuple[str, ...] = () + + @model_validator(mode="after") + def _validate_reference_categories(self) -> MTBenchJudgeSpec: + if any(not category for category in self.reference_categories): + raise ValueError("reference categories must not be empty") + if len(set(self.reference_categories)) != len(self.reference_categories): + raise ValueError("reference categories must not contain duplicates") + return self + + +class MTBenchProtocol(_StrictFrozenModel): + """Task policy used by the specialized multi-turn MT-Bench runner.""" + + runner: Literal["mt_bench"] + generation: MultiTurnGeneration + baseline: BaselineSpec + judge: MTBenchJudgeSpec + scoring: ScoringSpec + + +ProtocolSpec = Annotated[ + PairwiseProtocol | MTBenchProtocol, + Field(discriminator="runner"), +] + + class TaskMetadata(_StrictFrozenModel): reference_implementation: str | None = None paper: str | None = None @@ -168,9 +217,7 @@ def _validate_variants(self) -> SuffixVariants: if not members: raise ValueError(f"variant group {group!r} must not be empty") if len(set(members)) != len(members): - raise ValueError( - f"variant group {group!r} must not contain duplicates" - ) + raise ValueError(f"variant group {group!r} must not contain duplicates") unknown = sorted(set(members) - known) if unknown: raise ValueError( @@ -188,7 +235,7 @@ class TaskSpec(_StrictFrozenModel): description: str = Field(min_length=1) tags: tuple[str, ...] = () dataset: DatasetSpec - protocol: PairwiseProtocol + protocol: ProtocolSpec variants: SuffixVariants | None = None metadata: TaskMetadata = Field(default_factory=TaskMetadata) diff --git a/judgearena/utils/io.py b/judgearena/utils/io.py index b90cd0e..9e5c7f9 100644 --- a/judgearena/utils/io.py +++ b/judgearena/utils/io.py @@ -87,10 +87,6 @@ def download_all(): download_fluency_dataset(data_root) - from judgearena.datasets.mt_bench import download_mt_bench - - download_mt_bench() - class Timeblock: """Timer context manager""" diff --git a/tests/test_generate_and_evaluate.py b/tests/test_generate_and_evaluate.py index bb545c7..840f099 100644 --- a/tests/test_generate_and_evaluate.py +++ b/tests/test_generate_and_evaluate.py @@ -148,6 +148,10 @@ def test_m_arena_hard_baselines_are_not_duplicated_in_legacy_registry(): assert "m-arena-hard-v2.0" not in LEGACY_PAIRWISE_BASELINES +def test_mt_bench_baseline_is_not_duplicated_in_legacy_registry(): + assert "mt-bench" not in LEGACY_PAIRWISE_BASELINES + + def test_resolve_plan_explicit_model_b_overrides_native(): plan = _resolve_baseline_plan( task="arena-hard-v2.0", @@ -174,8 +178,15 @@ def test_native_pairwise_baseline_resolves_registered_tasks(task: str, expected: assert native_pairwise_baseline(task) == expected -def test_benchmark_adapter_resolution(): - assert resolve_benchmark_adapter("alpaca-eval").name == "pairwise" +@pytest.mark.parametrize( + ("task", "expected"), + [ + ("alpaca-eval", "pairwise"), + ("mt-bench", "mt_bench"), + ], +) +def test_benchmark_adapter_resolution(task: str, expected: str): + assert resolve_benchmark_adapter(task).name == expected def test_registered_task_runner_wins_over_legacy_fallback(monkeypatch): diff --git a/tests/test_instruction_dataset.py b/tests/test_instruction_dataset.py index e3fa0bb..521195f 100644 --- a/tests/test_instruction_dataset.py +++ b/tests/test_instruction_dataset.py @@ -151,7 +151,6 @@ def test_m_arena_hard_adapter_loads_invocation_specific_outputs(monkeypatch, tmp def test_mt_bench_native_baseline_is_flat_string(): from judgearena.datasets.mt_bench import ( - MT_BENCH_BASELINES, is_mt_bench_dataset, mt_bench_native_baseline, ) @@ -160,7 +159,6 @@ def test_mt_bench_native_baseline_is_flat_string(): assert is_mt_bench_dataset("alpaca-eval") is False assert mt_bench_native_baseline("mt-bench") == "gpt-4" assert mt_bench_native_baseline("alpaca-eval") is None - assert MT_BENCH_BASELINES == {"mt-bench": "gpt-4"} def test_normalize_official_arena_hard_v01_drops_no_category(): diff --git a/tests/test_mt_bench_downloads.py b/tests/test_mt_bench_downloads.py index 9e12376..1580199 100644 --- a/tests/test_mt_bench_downloads.py +++ b/tests/test_mt_bench_downloads.py @@ -9,6 +9,59 @@ import judgearena.utils.io as utils_io from judgearena.config import RunConfig from judgearena.prompts.registry import FASTCHAT_PAIRWISE_PROMPT_PRESET +from judgearena.tasks.registry import get_packaged_task + + +def test_mt_bench_sources_are_owned_by_task_yaml(): + task = get_packaged_task("mt-bench") + assert task is not None + + benchmark = task.spec.dataset.sources["benchmark"] + references = task.spec.dataset.sources["references"] + assert benchmark.repo_id == "lmsys/mt-bench" + assert benchmark.revision == "a4b674ca573c24143824ac7f60d9173e7081e37d" + assert benchmark.allow_patterns == ( + "data/mt_bench/question.jsonl", + "data/mt_bench/model_answer/gpt-4.jsonl", + ) + assert references.repository == "https://github.com/lm-sys/FastChat" + assert references.revision == "587d5cfa1609a43d192cedb8441cac3c17db105d" + assert mt_bench._git_raw_url(references).endswith( + "/587d5cfa1609a43d192cedb8441cac3c17db105d/" + "fastchat/llm_judge/data/mt_bench/reference_answer/gpt-4.jsonl" + ) + + +def test_mt_bench_adapter_normalizes_questions_and_references(monkeypatch, tmp_path): + task = get_packaged_task("mt-bench") + assert task is not None + question_path = tmp_path / "question.jsonl" + reference_path = tmp_path / "reference.jsonl" + question_path.write_text( + '{"question_id": 1, "category": "math", "turns": ["Q1", "Q2"]}\n' + ) + reference_path.write_text( + '{"question_id": 1, "choices": [{"turns": ["R1", "R2"]}]}\n' + ) + monkeypatch.setattr( + mt_bench, + "_download_mt_bench", + lambda _task, _local_dir: (question_path, reference_path), + ) + + loaded = mt_bench.load_task_instructions(task, tmp_path) + + assert loaded.to_dict(orient="records") == [ + { + "instruction_index": 1, + "category": "math", + "turn_1": "Q1", + "turn_2": "Q2", + "reference_turn_1": "R1", + "reference_turn_2": "R2", + "instruction": "Q1", + } + ] def test_download_mt_bench_skips_question_download_if_cached(tmp_path, monkeypatch): @@ -28,8 +81,8 @@ def _snapshot_download_stub(**_kwargs): monkeypatch.setattr(mt_bench, "snapshot_download", _snapshot_download_stub) monkeypatch.setattr( mt_bench, - "_download_gpt4_references", - lambda _local_dir: reference_path, + "_download_references", + lambda _task, _local_dir: reference_path, ) downloaded_question_path, downloaded_reference_path = mt_bench.download_mt_bench( @@ -43,7 +96,7 @@ def _snapshot_download_stub(**_kwargs): def test_download_all_includes_mt_bench(tmp_path, monkeypatch): hf_datasets = [] - calls = {"contexts": 0, "mt_bench": 0} + calls = {"contexts": 0} monkeypatch.setattr(utils_io, "data_root", tmp_path) monkeypatch.setattr( @@ -51,16 +104,11 @@ def test_download_all_includes_mt_bench(tmp_path, monkeypatch): "download_hf", lambda name, local_path: hf_datasets.append((name, local_path)), ) + def _contexts_snapshot_stub(**_kwargs): calls["contexts"] += 1 monkeypatch.setattr(fluency_mod, "snapshot_download", _contexts_snapshot_stub) - monkeypatch.setattr( - mt_bench, - "download_mt_bench", - lambda: calls.__setitem__("mt_bench", calls["mt_bench"] + 1), - ) - utils_io.download_all() tables_dir = tmp_path / "tables" @@ -70,10 +118,10 @@ def _contexts_snapshot_stub(**_kwargs): "arena-hard-v2.0", "m-arena-hard-v0.1", "m-arena-hard-v2.0", + "mt-bench", ] assert all(path == tables_dir for _, path in hf_datasets) assert calls["contexts"] == 1 - assert calls["mt_bench"] == 1 def test_load_mt_bench_model_answers_reads_cached_baseline_file(tmp_path): @@ -510,6 +558,8 @@ def fake_generate_multiturn(**kwargs): assert thinking_call["strip_thinking_before_turn_2_prompt"] is True assert thinking_call["thinking_token_budget"] == 8192 + assert thinking_call["temperature_config"]["writing"] == 0.7 + assert thinking_call["temperature_config"]["math"] == 0.0 assert plain_call["strip_thinking_before_turn_2_prompt"] is True assert "thinking_token_budget" not in plain_call @@ -567,3 +617,9 @@ def fake_judge(**kwargs): ) assert captured["judge"]["strip_thinking_before_judging"] is True + assert captured["judge"]["reference_categories"] == ( + "math", + "reasoning", + "coding", + "arena-hard-200", + ) diff --git a/tests/test_mt_bench_fastchat_compat.py b/tests/test_mt_bench_fastchat_compat.py index 7d83a4d..b70df9e 100644 --- a/tests/test_mt_bench_fastchat_compat.py +++ b/tests/test_mt_bench_fastchat_compat.py @@ -11,6 +11,8 @@ judge_mt_bench_pairwise_fastchat, ) +REFERENCE_CATEGORIES = ("math", "reasoning", "coding", "arena-hard-200") + class SequenceJudge: def __init__(self, outputs: list[str]): @@ -85,7 +87,11 @@ def test_select_prompt_variants( expected_name: str, expected_ref_based: bool, ): - prompt = _select_prompt(category, multi_turn=multi_turn) + prompt = _select_prompt( + category, + multi_turn=multi_turn, + reference_categories=REFERENCE_CATEGORIES, + ) assert prompt.name == expected_name assert prompt.ref_based is expected_ref_based @@ -108,6 +114,7 @@ def test_judge_mt_bench_pairwise_fastchat_swap_mode_both_is_conservative(): swap_mode="both", truncate_input_chars=None, use_tqdm=False, + reference_categories=REFERENCE_CATEGORIES, ) assert num_inconsistent == 0 diff --git a/tests/test_mt_bench_preset_judging.py b/tests/test_mt_bench_preset_judging.py index c6c60cf..803323e 100644 --- a/tests/test_mt_bench_preset_judging.py +++ b/tests/test_mt_bench_preset_judging.py @@ -10,6 +10,8 @@ ) from judgearena.prompts.registry import FASTCHAT_PAIRWISE_PROMPT_PRESET +REFERENCE_CATEGORIES = ("math", "reasoning", "coding", "arena-hard-200") + class SequenceJudge: def __init__(self, outputs: list[str]): @@ -51,6 +53,7 @@ def test_select_preset_prompt_rejects_delegated_preset(): _select_preset_prompt( "writing", multi_turn=False, + reference_categories=REFERENCE_CATEGORIES, prompt_preset=FASTCHAT_PAIRWISE_PROMPT_PRESET, provide_explanation=False, ) @@ -74,6 +77,7 @@ def test_select_preset_prompt_variants( prompt = _select_preset_prompt( category, multi_turn=multi_turn, + reference_categories=REFERENCE_CATEGORIES, prompt_preset="default", provide_explanation=False, ) @@ -93,6 +97,7 @@ def test_build_mt_bench_preset_items_adds_turn_and_reference_kwargs(): eval_single=True, eval_multi=True, truncate_input_chars=None, + reference_categories=REFERENCE_CATEGORIES, prompt_preset="default", provide_explanation=False, ) @@ -136,6 +141,7 @@ def test_judge_mt_bench_with_preset_parses_and_inverts_swapped_scores(): swap_mode="both", truncate_input_chars=None, use_tqdm=False, + reference_categories=REFERENCE_CATEGORIES, prompt_preset="default", ) diff --git a/tests/test_prompt_registry.py b/tests/test_prompt_registry.py index ecfe7b1..65ee364 100644 --- a/tests/test_prompt_registry.py +++ b/tests/test_prompt_registry.py @@ -58,6 +58,7 @@ def test_default_preset_for_unknown_task(): def test_mt_bench_default_is_delegated_fastchat(): + assert "mt-bench" not in TASK_DEFAULT_PRESET resolved = resolve_judge_prompt(task="mt-bench") assert resolved.preset_name == FASTCHAT_PAIRWISE_PROMPT_PRESET diff --git a/tests/test_task_registry.py b/tests/test_task_registry.py index 08fa6e3..c8d9440 100644 --- a/tests/test_task_registry.py +++ b/tests/test_task_registry.py @@ -80,6 +80,7 @@ def test_packaged_registry_discovers_versioned_tasks(): m_arena_v01 = tasks["m-arena-hard-v0.1"] m_arena_eu = resolve_task(tasks, "m-arena-hard-v2.0-EU") assert m_arena_eu is not None + mt_bench = tasks["mt-bench"] assert list(tasks) == [ "alpaca-eval", @@ -87,6 +88,7 @@ def test_packaged_registry_discovers_versioned_tasks(): "arena-hard-v2.0", "m-arena-hard-v0.1", "m-arena-hard-v2.0", + "mt-bench", ] assert alpaca.spec.dataset.sources["examples"].revision == ( "004c4a992956eeefffd36b63ade470f32fd0a582" @@ -123,6 +125,19 @@ def test_packaged_registry_discovers_versioned_tasks(): "ro", "uk", ) + assert mt_bench.spec.protocol.runner == "mt_bench" + assert mt_bench.spec.protocol.generation.mode == "multi_turn_chat" + assert mt_bench.spec.protocol.baseline.reference_id == "gpt-4" + assert mt_bench.spec.protocol.judge.default_prompt == "fastchat-pairwise" + assert mt_bench.spec.protocol.judge.reference_categories == ( + "math", + "reasoning", + "coding", + "arena-hard-200", + ) + assert mt_bench.spec.dataset.sources["benchmark"].revision == ( + "a4b674ca573c24143824ac7f60d9173e7081e37d" + ) assert alpaca.spec.protocol.scoring.primary_metric == "winrate" diff --git a/tests/test_utils.py b/tests/test_utils.py index 4bcd55a..2d5f26f 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,7 +1,6 @@ import pytest import judgearena.datasets.fluency as fluency_mod -import judgearena.datasets.mt_bench as mt_bench_mod import judgearena.models as utils_models import judgearena.utils as utils import judgearena.utils.io as utils_io @@ -47,23 +46,18 @@ def test_download_all_dispatches_registered_tasks(monkeypatch, tmp_path): ("snapshot", kwargs["repo_id"], kwargs["local_dir"]) ), ) - monkeypatch.setattr( - mt_bench_mod, - "download_mt_bench", - lambda local_dir=None: None, - ) - utils_io.download_all() tables_dir = tmp_path / "tables" - assert calls[:5] == [ + assert calls[:6] == [ ("hf", "alpaca-eval", 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), + ("hf", "mt-bench", tables_dir), ] - assert calls[5] == ( + assert calls[6] == ( "snapshot", "geoalgo/multilingual-fluency", tmp_path / "multilingual-fluency", From c1d99850723aa40941945a502701aad399bfae48 Mon Sep 17 00:00:00 2001 From: bora kargi Date: Mon, 20 Jul 2026 14:32:09 +0200 Subject: [PATCH 2/3] Remove legacy baseline registry Resolve native baselines exclusively from validated task definitions and drop tests for the obsolete compatibility map. --- judgearena/benchmarks/pairwise/baselines.py | 6 +----- tests/test_generate_and_evaluate.py | 24 +-------------------- 2 files changed, 2 insertions(+), 28 deletions(-) diff --git a/judgearena/benchmarks/pairwise/baselines.py b/judgearena/benchmarks/pairwise/baselines.py index 829c00e..8ea3bbd 100644 --- a/judgearena/benchmarks/pairwise/baselines.py +++ b/judgearena/benchmarks/pairwise/baselines.py @@ -7,11 +7,9 @@ from judgearena.tasks.registry import get_packaged_task from judgearena.tasks.schema import CategoryDefaultsBaseline, TaskDefaultBaseline -LEGACY_PAIRWISE_BASELINES: dict[str, str | Mapping[str, str]] = {} - def native_pairwise_baseline(task: str) -> str | Mapping[str, str] | None: - """Return the task-defined baseline, with fallback for unmigrated tasks.""" + """Return the baseline declared by a registered task.""" resolved = get_packaged_task(task) if resolved is not None: baseline = resolved.spec.protocol.baseline @@ -21,6 +19,4 @@ def native_pairwise_baseline(task: str) -> str | Mapping[str, str] | None: return baseline.references return None - if task in LEGACY_PAIRWISE_BASELINES: - return LEGACY_PAIRWISE_BASELINES[task] return None diff --git a/tests/test_generate_and_evaluate.py b/tests/test_generate_and_evaluate.py index 840f099..f650c24 100644 --- a/tests/test_generate_and_evaluate.py +++ b/tests/test_generate_and_evaluate.py @@ -6,10 +6,7 @@ import judgearena.benchmarks.execution as benchmark_execution import judgearena.benchmarks.pairwise.runner as generate_and_evaluate import judgearena.benchmarks.registry as benchmark_registry -from judgearena.benchmarks.pairwise.baselines import ( - LEGACY_PAIRWISE_BASELINES, - native_pairwise_baseline, -) +from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline from judgearena.benchmarks.pairwise.runner import ( BaselinePlan, _resolve_baseline_plan, @@ -133,25 +130,6 @@ 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_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_m_arena_hard_baselines_are_not_duplicated_in_legacy_registry(): - assert "m-arena-hard-v0.1" not in LEGACY_PAIRWISE_BASELINES - assert "m-arena-hard-v2.0" not in LEGACY_PAIRWISE_BASELINES - - -def test_mt_bench_baseline_is_not_duplicated_in_legacy_registry(): - assert "mt-bench" not in LEGACY_PAIRWISE_BASELINES - - def test_resolve_plan_explicit_model_b_overrides_native(): plan = _resolve_baseline_plan( task="arena-hard-v2.0", From 54f58994cbd894d6b8c11423716db649b7c16f6e Mon Sep 17 00:00:00 2001 From: bora kargi Date: Tue, 21 Jul 2026 13:49:23 +0200 Subject: [PATCH 3/3] Consolidate MT-Bench runner ownership Move the complete MT-Bench lifecycle into its registered runner and remove duplicate wrapper parameters. --- .../benchmarks/mt_bench/mt_bench_utils.py | 420 ------------------ judgearena/benchmarks/mt_bench/runner.py | 412 ++++++++++++++++- tests/test_mt_bench_downloads.py | 108 ++--- 3 files changed, 439 insertions(+), 501 deletions(-) delete mode 100644 judgearena/benchmarks/mt_bench/mt_bench_utils.py diff --git a/judgearena/benchmarks/mt_bench/mt_bench_utils.py b/judgearena/benchmarks/mt_bench/mt_bench_utils.py deleted file mode 100644 index fa16dd5..0000000 --- a/judgearena/benchmarks/mt_bench/mt_bench_utils.py +++ /dev/null @@ -1,420 +0,0 @@ -"""MT-Bench evaluation pipeline. - -Orchestrates multi-turn generation, FastChat-compatible pairwise judging, -and result saving for the MT-Bench benchmark. -""" - -from __future__ import annotations - -import os -from datetime import UTC, datetime -from pathlib import Path -from typing import TYPE_CHECKING - -import pandas as pd - -from judgearena.artifacts import prepare_run_directory, write_run_metadata_safely -from judgearena.benchmarks.mt_bench.fastchat_compat import ( - judge_mt_bench_pairwise_fastchat, -) -from judgearena.benchmarks.mt_bench.preset_judging import judge_mt_bench_with_preset -from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline -from judgearena.datasets import load_instructions -from judgearena.datasets.mt_bench import ( - load_mt_bench_model_answers, -) -from judgearena.generate import generate_multiturn -from judgearena.log import get_logger -from judgearena.models import is_thinking_model, make_model -from judgearena.prompts.registry import ResolvedJudgePrompt, resolve_run_judge_prompt -from judgearena.tasks.registry import get_packaged_task -from judgearena.tasks.schema import MTBenchProtocol -from judgearena.utils import ( - cache_function_dataframe, - compute_pref_summary, - generation_cache_token, -) -from judgearena.utils.eval import BattleReport, _compute_grouped_stats - -logger = get_logger(__name__) - -if TYPE_CHECKING: - from judgearena.config import RunConfig - - -def _task_protocol(task_id: str) -> MTBenchProtocol: - task = get_packaged_task(task_id) - if task is None or not isinstance(task.spec.protocol, MTBenchProtocol): - raise ValueError(f"Task {task_id!r} does not define an MT-Bench protocol.") - return task.spec.protocol - - -def _align_mt_bench_completions( - *, questions_df: pd.DataFrame, completions: pd.DataFrame, model_name: str -) -> pd.DataFrame: - """Align cached or generated MT-Bench completions to the question order.""" - indexed = completions.set_index("instruction_index") - missing_ids = questions_df.index.difference(indexed.index) - if not missing_ids.empty: - missing_ids_preview = ", ".join(str(x) for x in missing_ids[:5]) - raise ValueError( - f"MT-Bench completions for '{model_name}' are missing " - f"{len(missing_ids)} question(s). First missing ids: {missing_ids_preview}." - ) - return indexed.loc[questions_df.index] - - -def _build_mt_bench_generation_kwargs( - *, cfg: RunConfig, model_spec: str, role: str -) -> dict[str, object]: - """Battle-model kwargs, adding a thinking-token sub-budget when requested.""" - if role == "A": - generation_kwargs = cfg.model.evaluated_generation_kwargs() - elif role == "B": - generation_kwargs = cfg.model.baseline_generation_kwargs() - else: - raise ValueError(f"Unknown generation role: {role!r}") - provider, _, model_name = model_spec.partition("/") - if ( - cfg.judge.battle_thinking_token_budget is not None - and provider == "VLLM" - and is_thinking_model(model_name) - ): - max_tokens = int(generation_kwargs.get("max_tokens", cfg.model.max_out_tokens)) - generation_kwargs["thinking_token_budget"] = min( - int(cfg.judge.battle_thinking_token_budget), - max_tokens, - ) - return generation_kwargs - - -def _generate_mt_bench_completions( - cfg: RunConfig, - questions_df: pd.DataFrame, - ignore_cache: bool, -) -> tuple[pd.DataFrame, pd.DataFrame]: - cache_prefix = cfg.task - protocol = _task_protocol(cfg.task) - - def _run_generation( - model_name: str, *, generation_kwargs: dict[str, object] - ) -> pd.DataFrame: - # MT-Bench's category-aware temperatures only kick in when the user has - # not explicitly pinned a per-role temperature; otherwise the config - # override should win for reproducibility. - temperature_config = ( - None - if "temperature" in generation_kwargs - else dict(protocol.generation.category_temperatures) - ) - return generate_multiturn( - questions=questions_df, - model=model_name, - truncate_input_chars=cfg.generation.truncate_all_input_chars, - use_tqdm=cfg.run.use_tqdm, - temperature_config=temperature_config, - strip_thinking_before_turn_2_prompt=cfg.judge.strip_thinking_before_judging, - **generation_kwargs, - ) - - def _load_or_generate(model_name: str, *, role: str) -> pd.DataFrame: - loaded_answers = load_mt_bench_model_answers( - model_name, n_instructions=cfg.generation.n_instructions - ) - if loaded_answers is not None: - return _align_mt_bench_completions( - questions_df=questions_df, - completions=loaded_answers, - model_name=model_name, - ) - # Fold the resolved generation kwargs into the cache key so changing any - # sampling param busts cached completions instead of reusing a stale run. - generation_kwargs = _build_mt_bench_generation_kwargs( - cfg=cfg, model_spec=model_name, role=role - ) - sampling_token = generation_cache_token(generation_kwargs) - generated_answers = cache_function_dataframe( - lambda: _run_generation(model_name, generation_kwargs=generation_kwargs), - ignore_cache=ignore_cache, - cache_name=( - f"{cache_prefix}_{model_name}_{cfg.generation.n_instructions}_" - f"{sampling_token}" - ), - ) - return _align_mt_bench_completions( - questions_df=questions_df, - completions=generated_answers, - model_name=model_name, - ) - - return _load_or_generate(cfg.model.name, role="A"), _load_or_generate( - cfg.model.baseline, role="B" - ) - - -def _build_mt_bench_input_payloads( - *, - questions_df: pd.DataFrame, - completions_a: pd.DataFrame, - completions_b: pd.DataFrame, -) -> dict[str, object]: - return { - "instruction_index": questions_df.index.tolist(), - "turn_1": questions_df["turn_1"].tolist(), - "turn_2": questions_df["turn_2"].tolist(), - "completion_turn_1_A": completions_a["completion_turn_1"].tolist(), - "completion_turn_2_A": completions_a["completion_turn_2"].tolist(), - "completion_turn_1_B": completions_b["completion_turn_1"].tolist(), - "completion_turn_2_B": completions_b["completion_turn_2"].tolist(), - } - - -def _save_mt_bench_results( - *, - cfg: RunConfig, - res_folder: Path, - result_name: str, - results: dict[str, object], - annotations_df: pd.DataFrame, - started_at_utc: datetime, - input_payloads: dict[str, object], - judge_system_prompt: str | None = None, - judge_user_prompt_template: str | None = None, -) -> None: - """Persist MT-Bench arguments, annotations, aggregate results, and metadata.""" - prepare_run_directory(cfg, res_folder, attach_log=False) - - annotations_df.to_csv(res_folder / f"{result_name}-annotations.csv", index=False) - - write_run_metadata_safely( - output_dir=res_folder, - entrypoint="judgearena.benchmarks.mt_bench.mt_bench_utils.run_mt_bench", - run=cfg.model_dump(), - results=results, - input_payloads=input_payloads, - judge_system_prompt=judge_system_prompt, - judge_user_prompt_template=judge_user_prompt_template, - started_at_utc=started_at_utc, - ) - - -def _finalize_mt_bench_run( - *, - cfg: RunConfig, - res_folder: Path, - result_name: str, - prefs: pd.Series, - annotations: list[dict[str, object]], - combined_metadata: list[dict[str, object]], - resolved_prompt: ResolvedJudgePrompt, - questions_df: pd.DataFrame, - completions_a: pd.DataFrame, - completions_b: pd.DataFrame, - started_at_utc: datetime, - extra_result_fields: dict[str, object] | None = None, -) -> pd.Series: - stats = compute_pref_summary(prefs) - report = BattleReport( - task=cfg.task, - model_a=cfg.model.name, - model_b=cfg.model.baseline, - judge_model=cfg.judge.model, - summary=stats, - per_category=_compute_grouped_stats(prefs, combined_metadata, "category"), - per_turn=_compute_grouped_stats(prefs, combined_metadata, "turn"), - preferences=prefs.tolist(), - metadata={ - **resolved_prompt.metadata(), - "battle_thinking_token_budget": cfg.judge.battle_thinking_token_budget, - "strip_thinking_before_judging": cfg.judge.strip_thinking_before_judging, - **(extra_result_fields or {}), - "date": datetime.now(UTC).isoformat(), - "user": os.getenv("USER", ""), - }, - ) - results = report.to_dict() - report.render() - report.save(res_folder / f"results-{result_name}.json") - _save_mt_bench_results( - cfg=cfg, - res_folder=res_folder, - result_name=result_name, - results=results, - annotations_df=pd.DataFrame(annotations), - started_at_utc=started_at_utc, - input_payloads=_build_mt_bench_input_payloads( - questions_df=questions_df, - completions_a=completions_a, - completions_b=completions_b, - ), - judge_system_prompt=resolved_prompt.system_prompt, - judge_user_prompt_template=resolved_prompt.user_prompt_template, - ) - return prefs - - -def _run_mt_bench_fastchat( - *, - cfg: RunConfig, - res_folder: Path, - result_name: str, - questions_df: pd.DataFrame, - completions_a: pd.DataFrame, - completions_b: pd.DataFrame, - judge_chat_model, - resolved_prompt: ResolvedJudgePrompt, - fastchat_prompt_preset: str, - started_at_utc: datetime, -) -> pd.Series: - protocol = _task_protocol(cfg.task) - prefs, annotations, combined_metadata, num_inconsistent = ( - judge_mt_bench_pairwise_fastchat( - judge_chat_model=judge_chat_model, - judge_model=cfg.judge.model, - questions=questions_df, - completions_a=completions_a, - completions_b=completions_b, - model_a=cfg.model.name, - model_b=cfg.model.baseline, - turns_mode=protocol.judge.turns_mode, - swap_mode=cfg.judge.swap_mode, - truncate_input_chars=cfg.generation.truncate_judge_input_chars, - use_tqdm=cfg.run.use_tqdm, - reference_categories=protocol.judge.reference_categories, - prompt_preset=fastchat_prompt_preset, - strip_thinking_before_judging=cfg.judge.strip_thinking_before_judging, - ) - ) - return _finalize_mt_bench_run( - cfg=cfg, - res_folder=res_folder, - result_name=result_name, - prefs=prefs, - annotations=annotations, - combined_metadata=combined_metadata, - resolved_prompt=resolved_prompt, - questions_df=questions_df, - completions_a=completions_a, - completions_b=completions_b, - started_at_utc=started_at_utc, - extra_result_fields={"num_inconsistent": num_inconsistent}, - ) - - -def _run_mt_bench_preset( - *, - cfg: RunConfig, - res_folder: Path, - result_name: str, - questions_df: pd.DataFrame, - completions_a: pd.DataFrame, - completions_b: pd.DataFrame, - judge_chat_model, - resolved_prompt: ResolvedJudgePrompt, - started_at_utc: datetime, -) -> pd.Series: - protocol = _task_protocol(cfg.task) - prefs, annotations, combined_metadata = judge_mt_bench_with_preset( - judge_chat_model=judge_chat_model, - judge_model=cfg.judge.model, - questions=questions_df, - completions_a=completions_a, - completions_b=completions_b, - model_a=cfg.model.name, - model_b=cfg.model.baseline, - turns_mode=protocol.judge.turns_mode, - swap_mode=cfg.judge.swap_mode, - truncate_input_chars=cfg.generation.truncate_judge_input_chars, - use_tqdm=cfg.run.use_tqdm, - reference_categories=protocol.judge.reference_categories, - prompt_preset=cfg.judge.prompt_preset or resolved_prompt.preset_name, - provide_explanation=cfg.judge.provide_explanation, - system_file=cfg.judge.system_prompt_file, - user_file=cfg.judge.user_prompt_file, - strip_thinking_before_judging=cfg.judge.strip_thinking_before_judging, - ) - return _finalize_mt_bench_run( - cfg=cfg, - res_folder=res_folder, - result_name=result_name, - prefs=prefs, - annotations=annotations, - combined_metadata=combined_metadata, - resolved_prompt=resolved_prompt, - questions_df=questions_df, - completions_a=completions_a, - completions_b=completions_b, - started_at_utc=started_at_utc, - ) - - -def run_mt_bench( - cfg: RunConfig, - ignore_cache: bool, - *, - res_folder: Path, - result_name: str, -): - """MT-Bench pipeline with preset or FastChat-original pairwise judging.""" - run_started_at = datetime.now(UTC) - protocol = _task_protocol(cfg.task) - if cfg.model.baseline is None: - baseline = native_pairwise_baseline(cfg.task) - cfg.model.baseline = baseline if isinstance(baseline, str) else None - if cfg.model.baseline is None: - raise ValueError( - f"--model_B is required for dataset '{cfg.task}'; " - "no dataset-native baseline registered." - ) - questions_df = load_instructions( - cfg.task, n_instructions=cfg.generation.n_instructions - ) - logger.info( - "Generating multi-turn completions for MT-Bench with %s and %s.", - cfg.model.name, - cfg.model.baseline, - ) - completions_a, completions_b = _generate_mt_bench_completions( - cfg=cfg, - questions_df=questions_df, - ignore_cache=ignore_cache, - ) - resolved_prompt = resolve_run_judge_prompt(cfg.task, cfg.judge, multi_turn=True) - if resolved_prompt.delegated and not cfg.judge.provide_explanation: - logger.info( - "MT-Bench keeps the original FastChat-style explanation-plus-verdict " - "prompt when delegated to FastChat compatibility mode." - ) - judge_model_kwargs = cfg.judge.model_kwargs( - base_engine_kwargs=cfg.model.engine_kwargs, - fallback_chat_template=cfg.model.chat_template, - ) - if resolved_prompt.delegated and cfg.judge.temperature is None: - judge_model_kwargs.setdefault( - "temperature", protocol.judge.fastchat_temperature - ) - judge_chat_model = make_model(model=cfg.judge.model, **judge_model_kwargs) - if resolved_prompt.delegated: - return _run_mt_bench_fastchat( - cfg=cfg, - res_folder=res_folder, - result_name=result_name, - questions_df=questions_df, - completions_a=completions_a, - completions_b=completions_b, - judge_chat_model=judge_chat_model, - resolved_prompt=resolved_prompt, - fastchat_prompt_preset=protocol.judge.fastchat_prompt_preset, - started_at_utc=run_started_at, - ) - return _run_mt_bench_preset( - cfg=cfg, - res_folder=res_folder, - result_name=result_name, - questions_df=questions_df, - completions_a=completions_a, - completions_b=completions_b, - judge_chat_model=judge_chat_model, - resolved_prompt=resolved_prompt, - started_at_utc=run_started_at, - ) diff --git a/judgearena/benchmarks/mt_bench/runner.py b/judgearena/benchmarks/mt_bench/runner.py index e5fd471..5e05e97 100644 --- a/judgearena/benchmarks/mt_bench/runner.py +++ b/judgearena/benchmarks/mt_bench/runner.py @@ -1,37 +1,421 @@ -"""Registered entry point for the specialized MT-Bench pipeline.""" +"""Registered MT-Bench runner and evaluation pipeline. + +Orchestrates multi-turn generation, FastChat-compatible pairwise judging, +and result saving for the MT-Bench benchmark. +""" from __future__ import annotations +import os from datetime import UTC, datetime from pathlib import Path from typing import TYPE_CHECKING -from judgearena.artifacts import prepare_run_directory -from judgearena.benchmarks.mt_bench.mt_bench_utils import run_mt_bench +import pandas as pd + +from judgearena.artifacts import prepare_run_directory, write_run_metadata_safely +from judgearena.benchmarks.mt_bench.fastchat_compat import ( + judge_mt_bench_pairwise_fastchat, +) +from judgearena.benchmarks.mt_bench.preset_judging import judge_mt_bench_with_preset from judgearena.benchmarks.pairwise.baselines import native_pairwise_baseline +from judgearena.datasets import load_instructions +from judgearena.datasets.mt_bench import ( + load_mt_bench_model_answers, +) +from judgearena.generate import generate_multiturn +from judgearena.log import get_logger +from judgearena.models import is_thinking_model, make_model +from judgearena.prompts.registry import ResolvedJudgePrompt, resolve_run_judge_prompt +from judgearena.tasks.registry import get_packaged_task +from judgearena.tasks.schema import MTBenchProtocol +from judgearena.utils import ( + cache_function_dataframe, + compute_pref_summary, + generation_cache_token, +) +from judgearena.utils.eval import BattleReport, _compute_grouped_stats + +logger = get_logger(__name__) if TYPE_CHECKING: from judgearena.config import RunConfig +def _task_protocol(task_id: str) -> MTBenchProtocol: + task = get_packaged_task(task_id) + if task is None or not isinstance(task.spec.protocol, MTBenchProtocol): + raise ValueError(f"Task {task_id!r} does not define an MT-Bench protocol.") + return task.spec.protocol + + +def _align_mt_bench_completions( + *, questions_df: pd.DataFrame, completions: pd.DataFrame, model_name: str +) -> pd.DataFrame: + """Align cached or generated MT-Bench completions to the question order.""" + indexed = completions.set_index("instruction_index") + missing_ids = questions_df.index.difference(indexed.index) + if not missing_ids.empty: + missing_ids_preview = ", ".join(str(x) for x in missing_ids[:5]) + raise ValueError( + f"MT-Bench completions for '{model_name}' are missing " + f"{len(missing_ids)} question(s). First missing ids: {missing_ids_preview}." + ) + return indexed.loc[questions_df.index] + + +def _build_mt_bench_generation_kwargs( + *, cfg: RunConfig, model_spec: str, role: str +) -> dict[str, object]: + """Battle-model kwargs, adding a thinking-token sub-budget when requested.""" + if role == "A": + generation_kwargs = cfg.model.evaluated_generation_kwargs() + elif role == "B": + generation_kwargs = cfg.model.baseline_generation_kwargs() + else: + raise ValueError(f"Unknown generation role: {role!r}") + provider, _, model_name = model_spec.partition("/") + if ( + cfg.judge.battle_thinking_token_budget is not None + and provider == "VLLM" + and is_thinking_model(model_name) + ): + max_tokens = int(generation_kwargs.get("max_tokens", cfg.model.max_out_tokens)) + generation_kwargs["thinking_token_budget"] = min( + int(cfg.judge.battle_thinking_token_budget), + max_tokens, + ) + return generation_kwargs + + +def _generate_mt_bench_completions( + cfg: RunConfig, + questions_df: pd.DataFrame, +) -> tuple[pd.DataFrame, pd.DataFrame]: + cache_prefix = cfg.task + protocol = _task_protocol(cfg.task) + + def _run_generation( + model_name: str, *, generation_kwargs: dict[str, object] + ) -> pd.DataFrame: + # MT-Bench's category-aware temperatures only kick in when the user has + # not explicitly pinned a per-role temperature; otherwise the config + # override should win for reproducibility. + temperature_config = ( + None + if "temperature" in generation_kwargs + else dict(protocol.generation.category_temperatures) + ) + return generate_multiturn( + questions=questions_df, + model=model_name, + truncate_input_chars=cfg.generation.truncate_all_input_chars, + use_tqdm=cfg.run.use_tqdm, + temperature_config=temperature_config, + strip_thinking_before_turn_2_prompt=cfg.judge.strip_thinking_before_judging, + **generation_kwargs, + ) + + def _load_or_generate(model_name: str, *, role: str) -> pd.DataFrame: + loaded_answers = load_mt_bench_model_answers( + model_name, n_instructions=cfg.generation.n_instructions + ) + if loaded_answers is not None: + return _align_mt_bench_completions( + questions_df=questions_df, + completions=loaded_answers, + model_name=model_name, + ) + # Fold the resolved generation kwargs into the cache key so changing any + # sampling param busts cached completions instead of reusing a stale run. + generation_kwargs = _build_mt_bench_generation_kwargs( + cfg=cfg, model_spec=model_name, role=role + ) + sampling_token = generation_cache_token(generation_kwargs) + generated_answers = cache_function_dataframe( + lambda: _run_generation(model_name, generation_kwargs=generation_kwargs), + ignore_cache=cfg.run.ignore_cache, + cache_name=( + f"{cache_prefix}_{model_name}_{cfg.generation.n_instructions}_" + f"{sampling_token}" + ), + ) + return _align_mt_bench_completions( + questions_df=questions_df, + completions=generated_answers, + model_name=model_name, + ) + + return _load_or_generate(cfg.model.name, role="A"), _load_or_generate( + cfg.model.baseline, role="B" + ) + + +def _build_mt_bench_input_payloads( + *, + questions_df: pd.DataFrame, + completions_a: pd.DataFrame, + completions_b: pd.DataFrame, +) -> dict[str, object]: + return { + "instruction_index": questions_df.index.tolist(), + "turn_1": questions_df["turn_1"].tolist(), + "turn_2": questions_df["turn_2"].tolist(), + "completion_turn_1_A": completions_a["completion_turn_1"].tolist(), + "completion_turn_2_A": completions_a["completion_turn_2"].tolist(), + "completion_turn_1_B": completions_b["completion_turn_1"].tolist(), + "completion_turn_2_B": completions_b["completion_turn_2"].tolist(), + } + + +def _save_mt_bench_results( + *, + cfg: RunConfig, + res_folder: Path, + result_name: str, + results: dict[str, object], + annotations_df: pd.DataFrame, + started_at_utc: datetime, + input_payloads: dict[str, object], + judge_system_prompt: str | None = None, + judge_user_prompt_template: str | None = None, +) -> None: + """Persist MT-Bench arguments, annotations, aggregate results, and metadata.""" + annotations_df.to_csv(res_folder / f"{result_name}-annotations.csv", index=False) + + write_run_metadata_safely( + output_dir=res_folder, + entrypoint="judgearena.benchmarks.mt_bench.runner.run_mt_bench_benchmark", + run=cfg.model_dump(), + results=results, + input_payloads=input_payloads, + judge_system_prompt=judge_system_prompt, + judge_user_prompt_template=judge_user_prompt_template, + started_at_utc=started_at_utc, + ) + + +def _finalize_mt_bench_run( + *, + cfg: RunConfig, + res_folder: Path, + result_name: str, + prefs: pd.Series, + annotations: list[dict[str, object]], + combined_metadata: list[dict[str, object]], + resolved_prompt: ResolvedJudgePrompt, + questions_df: pd.DataFrame, + completions_a: pd.DataFrame, + completions_b: pd.DataFrame, + started_at_utc: datetime, + extra_result_fields: dict[str, object] | None = None, +) -> pd.Series: + stats = compute_pref_summary(prefs) + report = BattleReport( + task=cfg.task, + model_a=cfg.model.name, + model_b=cfg.model.baseline, + judge_model=cfg.judge.model, + summary=stats, + per_category=_compute_grouped_stats(prefs, combined_metadata, "category"), + per_turn=_compute_grouped_stats(prefs, combined_metadata, "turn"), + preferences=prefs.tolist(), + metadata={ + **resolved_prompt.metadata(), + "battle_thinking_token_budget": cfg.judge.battle_thinking_token_budget, + "strip_thinking_before_judging": cfg.judge.strip_thinking_before_judging, + **(extra_result_fields or {}), + "date": datetime.now(UTC).isoformat(), + "user": os.getenv("USER", ""), + }, + ) + results = report.to_dict() + report.render() + report.save(res_folder / f"results-{result_name}.json") + _save_mt_bench_results( + cfg=cfg, + res_folder=res_folder, + result_name=result_name, + results=results, + annotations_df=pd.DataFrame(annotations), + started_at_utc=started_at_utc, + input_payloads=_build_mt_bench_input_payloads( + questions_df=questions_df, + completions_a=completions_a, + completions_b=completions_b, + ), + judge_system_prompt=resolved_prompt.system_prompt, + judge_user_prompt_template=resolved_prompt.user_prompt_template, + ) + return prefs + + +def _run_mt_bench_fastchat( + *, + cfg: RunConfig, + res_folder: Path, + result_name: str, + questions_df: pd.DataFrame, + completions_a: pd.DataFrame, + completions_b: pd.DataFrame, + judge_chat_model, + resolved_prompt: ResolvedJudgePrompt, + fastchat_prompt_preset: str, + started_at_utc: datetime, +) -> pd.Series: + protocol = _task_protocol(cfg.task) + prefs, annotations, combined_metadata, num_inconsistent = ( + judge_mt_bench_pairwise_fastchat( + judge_chat_model=judge_chat_model, + judge_model=cfg.judge.model, + questions=questions_df, + completions_a=completions_a, + completions_b=completions_b, + model_a=cfg.model.name, + model_b=cfg.model.baseline, + turns_mode=protocol.judge.turns_mode, + swap_mode=cfg.judge.swap_mode, + truncate_input_chars=cfg.generation.truncate_judge_input_chars, + use_tqdm=cfg.run.use_tqdm, + reference_categories=protocol.judge.reference_categories, + prompt_preset=fastchat_prompt_preset, + strip_thinking_before_judging=cfg.judge.strip_thinking_before_judging, + ) + ) + return _finalize_mt_bench_run( + cfg=cfg, + res_folder=res_folder, + result_name=result_name, + prefs=prefs, + annotations=annotations, + combined_metadata=combined_metadata, + resolved_prompt=resolved_prompt, + questions_df=questions_df, + completions_a=completions_a, + completions_b=completions_b, + started_at_utc=started_at_utc, + extra_result_fields={"num_inconsistent": num_inconsistent}, + ) + + +def _run_mt_bench_preset( + *, + cfg: RunConfig, + res_folder: Path, + result_name: str, + questions_df: pd.DataFrame, + completions_a: pd.DataFrame, + completions_b: pd.DataFrame, + judge_chat_model, + resolved_prompt: ResolvedJudgePrompt, + started_at_utc: datetime, +) -> pd.Series: + protocol = _task_protocol(cfg.task) + prefs, annotations, combined_metadata = judge_mt_bench_with_preset( + judge_chat_model=judge_chat_model, + judge_model=cfg.judge.model, + questions=questions_df, + completions_a=completions_a, + completions_b=completions_b, + model_a=cfg.model.name, + model_b=cfg.model.baseline, + turns_mode=protocol.judge.turns_mode, + swap_mode=cfg.judge.swap_mode, + truncate_input_chars=cfg.generation.truncate_judge_input_chars, + use_tqdm=cfg.run.use_tqdm, + reference_categories=protocol.judge.reference_categories, + prompt_preset=cfg.judge.prompt_preset or resolved_prompt.preset_name, + provide_explanation=cfg.judge.provide_explanation, + system_file=cfg.judge.system_prompt_file, + user_file=cfg.judge.user_prompt_file, + strip_thinking_before_judging=cfg.judge.strip_thinking_before_judging, + ) + return _finalize_mt_bench_run( + cfg=cfg, + res_folder=res_folder, + result_name=result_name, + prefs=prefs, + annotations=annotations, + combined_metadata=combined_metadata, + resolved_prompt=resolved_prompt, + questions_df=questions_df, + completions_a=completions_a, + completions_b=completions_b, + started_at_utc=started_at_utc, + ) + + def run_mt_bench_benchmark(cfg: RunConfig): - """Prepare one run directory and execute the YAML-selected MT-Bench runner.""" - baseline = cfg.model.baseline or native_pairwise_baseline(cfg.task) - if not isinstance(baseline, str): - raise ValueError("MT-Bench requires a flat native baseline.") + """Run the registered MT-Bench generation, judging, and reporting lifecycle.""" + run_started_at = datetime.now(UTC) + protocol = _task_protocol(cfg.task) + if cfg.model.baseline is None: + baseline = native_pairwise_baseline(cfg.task) + cfg.model.baseline = baseline if isinstance(baseline, str) else None + if cfg.model.baseline is None: + raise ValueError( + f"--model_B is required for dataset '{cfg.task}'; " + "no dataset-native baseline registered." + ) result_name = ( - f"{cfg.task}-{cfg.model.name}-{baseline}-{cfg.judge.model}-" + f"{cfg.task}-{cfg.model.name}-{cfg.model.baseline}-{cfg.judge.model}-" f"{cfg.judge.swap_mode}" ).replace("/", "_") - run_timestamp = datetime.now(UTC).strftime("%Y%m%d_%H%M%S") - result_folder = prepare_run_directory( + run_timestamp = run_started_at.strftime("%Y%m%d_%H%M%S") + res_folder = prepare_run_directory( cfg, Path(cfg.run.result_folder) / f"{result_name}-{run_timestamp}", ) - return run_mt_bench( - cfg, - cfg.run.ignore_cache, - res_folder=result_folder, + + questions_df = load_instructions( + cfg.task, n_instructions=cfg.generation.n_instructions + ) + logger.info( + "Generating multi-turn completions for MT-Bench with %s and %s.", + cfg.model.name, + cfg.model.baseline, + ) + completions_a, completions_b = _generate_mt_bench_completions( + cfg=cfg, + questions_df=questions_df, + ) + resolved_prompt = resolve_run_judge_prompt(cfg.task, cfg.judge, multi_turn=True) + if resolved_prompt.delegated and not cfg.judge.provide_explanation: + logger.info( + "MT-Bench keeps the original FastChat-style explanation-plus-verdict " + "prompt when delegated to FastChat compatibility mode." + ) + judge_model_kwargs = cfg.judge.model_kwargs( + base_engine_kwargs=cfg.model.engine_kwargs, + fallback_chat_template=cfg.model.chat_template, + ) + if resolved_prompt.delegated and cfg.judge.temperature is None: + judge_model_kwargs.setdefault( + "temperature", protocol.judge.fastchat_temperature + ) + judge_chat_model = make_model(model=cfg.judge.model, **judge_model_kwargs) + if resolved_prompt.delegated: + return _run_mt_bench_fastchat( + cfg=cfg, + res_folder=res_folder, + result_name=result_name, + questions_df=questions_df, + completions_a=completions_a, + completions_b=completions_b, + judge_chat_model=judge_chat_model, + resolved_prompt=resolved_prompt, + fastchat_prompt_preset=protocol.judge.fastchat_prompt_preset, + started_at_utc=run_started_at, + ) + return _run_mt_bench_preset( + cfg=cfg, + res_folder=res_folder, result_name=result_name, + questions_df=questions_df, + completions_a=completions_a, + completions_b=completions_b, + judge_chat_model=judge_chat_model, + resolved_prompt=resolved_prompt, + started_at_utc=run_started_at, ) diff --git a/tests/test_mt_bench_downloads.py b/tests/test_mt_bench_downloads.py index 1580199..b142245 100644 --- a/tests/test_mt_bench_downloads.py +++ b/tests/test_mt_bench_downloads.py @@ -3,7 +3,7 @@ import pandas as pd import pytest -import judgearena.benchmarks.mt_bench.mt_bench_utils as mt_bench_utils +import judgearena.benchmarks.mt_bench.runner as mt_bench_runner import judgearena.datasets.fluency as fluency_mod import judgearena.datasets.mt_bench as mt_bench import judgearena.utils.io as utils_io @@ -156,7 +156,7 @@ def test_generate_mt_bench_completions_uses_pregenerated_baseline(monkeypatch): generated_models = [] monkeypatch.setattr( - mt_bench_utils, "cache_function_dataframe", lambda fun, **_kwargs: fun() + mt_bench_runner, "cache_function_dataframe", lambda fun, **_kwargs: fun() ) def fake_generate_multiturn(**kwargs): @@ -169,9 +169,9 @@ def fake_generate_multiturn(**kwargs): } ) - monkeypatch.setattr(mt_bench_utils, "generate_multiturn", fake_generate_multiturn) + monkeypatch.setattr(mt_bench_runner, "generate_multiturn", fake_generate_multiturn) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_mt_bench_model_answers", lambda model, n_instructions=None: ( pd.DataFrame( @@ -197,10 +197,9 @@ def fake_generate_multiturn(**kwargs): generation={"n_instructions": 2}, ) - completions_a, completions_b = mt_bench_utils._generate_mt_bench_completions( + completions_a, completions_b = mt_bench_runner._generate_mt_bench_completions( cfg=cfg, questions_df=questions_df, - ignore_cache=False, ) assert generated_models == ["VLLM/example/model-a"] @@ -216,7 +215,7 @@ def test_generate_mt_bench_completions_reports_missing_baseline_rows(monkeypatch ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_mt_bench_model_answers", lambda model, n_instructions=None: pd.DataFrame( { @@ -235,10 +234,9 @@ def test_generate_mt_bench_completions_reports_missing_baseline_rows(monkeypatch ) with pytest.raises(ValueError, match="missing 1 question"): - mt_bench_utils._generate_mt_bench_completions( + mt_bench_runner._generate_mt_bench_completions( cfg=cfg, questions_df=questions_df, - ignore_cache=False, ) @@ -250,7 +248,7 @@ def fake_write_run_metadata(**kwargs): return tmp_path / "run-metadata.v1.json" monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "write_run_metadata_safely", fake_write_run_metadata, ) @@ -261,7 +259,7 @@ def fake_write_run_metadata(**kwargs): ) started_at = datetime(2026, 1, 2, 3, 4, tzinfo=UTC) - mt_bench_utils._save_mt_bench_results( + mt_bench_runner._save_mt_bench_results( cfg=cfg, res_folder=tmp_path, result_name="mt-bench-test", @@ -273,13 +271,10 @@ def fake_write_run_metadata(**kwargs): judge_user_prompt_template="user", ) - assert (tmp_path / "config.yaml").exists() assert (tmp_path / "mt-bench-test-annotations.csv").exists() - # The results JSON is now written by report.save() in _finalize_mt_bench_run, - # not by _save_mt_bench_results (which writes config / annotations / run metadata). assert ( captured["entrypoint"] - == "judgearena.benchmarks.mt_bench.mt_bench_utils.run_mt_bench" + == "judgearena.benchmarks.mt_bench.runner.run_mt_bench_benchmark" ) assert captured["input_payloads"] == {"instruction_index": [1]} assert captured["judge_system_prompt"] == "system" @@ -297,14 +292,14 @@ def test_run_mt_bench_resolves_native_baseline_and_judge_controls( captured = {} monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_instructions", lambda dataset, n_instructions=None: questions_df, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_generate_mt_bench_completions", - lambda cfg, questions_df, ignore_cache: ( + lambda cfg, questions_df: ( pd.DataFrame( {"completion_turn_1": ["A1"], "completion_turn_2": ["A2"]}, index=questions_df.index, @@ -320,14 +315,14 @@ def fake_make_model(**kwargs): captured["make_model"] = kwargs return object() - monkeypatch.setattr(mt_bench_utils, "make_model", fake_make_model) + monkeypatch.setattr(mt_bench_runner, "make_model", fake_make_model) def fake_run_mt_bench_fastchat(**kwargs): captured["fastchat"] = kwargs return pd.Series([0.0], dtype=float) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_run_mt_bench_fastchat", fake_run_mt_bench_fastchat, ) @@ -348,12 +343,7 @@ def fake_run_mt_bench_fastchat(**kwargs): run={"result_folder": str(tmp_path)}, ) - mt_bench_utils.run_mt_bench( - cfg, - ignore_cache=False, - res_folder=tmp_path, - result_name="mt-bench-test", - ) + mt_bench_runner.run_mt_bench_benchmark(cfg) assert cfg.model.baseline == "gpt-4" assert captured["make_model"]["max_model_len"] == 65536 @@ -373,14 +363,14 @@ def test_run_mt_bench_defaults_to_delegated_fastchat(monkeypatch, tmp_path): captured = {} monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_instructions", lambda dataset, n_instructions=None: questions_df, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_generate_mt_bench_completions", - lambda cfg, questions_df, ignore_cache: ( + lambda cfg, questions_df: ( pd.DataFrame( {"completion_turn_1": ["A1"], "completion_turn_2": ["A2"]}, index=questions_df.index, @@ -396,19 +386,19 @@ def fake_make_model(**kwargs): captured["make_model"] = kwargs return object() - monkeypatch.setattr(mt_bench_utils, "make_model", fake_make_model) + monkeypatch.setattr(mt_bench_runner, "make_model", fake_make_model) def fake_run_mt_bench_fastchat(**kwargs): captured["fastchat"] = kwargs return pd.Series([0.0], dtype=float) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_run_mt_bench_fastchat", fake_run_mt_bench_fastchat, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_run_mt_bench_preset", lambda **_kwargs: pytest.fail("preset path should not run"), ) @@ -421,12 +411,7 @@ def fake_run_mt_bench_fastchat(**kwargs): run={"result_folder": str(tmp_path)}, ) - mt_bench_utils.run_mt_bench( - cfg, - ignore_cache=False, - res_folder=tmp_path, - result_name="mt-bench-test", - ) + mt_bench_runner.run_mt_bench_benchmark(cfg) assert cfg.model.baseline == "gpt-4" assert captured["make_model"]["temperature"] == 0.0 @@ -443,14 +428,14 @@ def test_run_mt_bench_concrete_prompt_preset_uses_preset_judging(monkeypatch, tm ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_instructions", lambda dataset, n_instructions=None: questions_df, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_generate_mt_bench_completions", - lambda cfg, questions_df, ignore_cache: ( + lambda cfg, questions_df: ( pd.DataFrame( {"completion_turn_1": ["A1"], "completion_turn_2": ["A2"]}, index=questions_df.index, @@ -471,14 +456,14 @@ def fake_run_mt_bench_preset(**kwargs): return pd.Series([0.0], dtype=float) captured = {} - monkeypatch.setattr(mt_bench_utils, "make_model", fake_make_model) + monkeypatch.setattr(mt_bench_runner, "make_model", fake_make_model) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_run_mt_bench_preset", fake_run_mt_bench_preset, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_run_mt_bench_fastchat", lambda **_kwargs: pytest.fail("fastchat path should not run"), ) @@ -491,12 +476,7 @@ def fake_run_mt_bench_preset(**kwargs): run={"result_folder": str(tmp_path)}, ) - mt_bench_utils.run_mt_bench( - cfg, - ignore_cache=False, - res_folder=tmp_path, - result_name="mt-bench-test", - ) + mt_bench_runner.run_mt_bench_benchmark(cfg) assert captured["preset"]["resolved_prompt"].preset_name == ( "default_with_explanation" @@ -512,10 +492,10 @@ def test_generate_mt_bench_completions_forwards_thinking_controls(monkeypatch): captured: dict[str, dict] = {} monkeypatch.setattr( - mt_bench_utils, "cache_function_dataframe", lambda fun, **_kwargs: fun() + mt_bench_runner, "cache_function_dataframe", lambda fun, **_kwargs: fun() ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_mt_bench_model_answers", lambda model, n_instructions=None: None, ) @@ -530,7 +510,7 @@ def fake_generate_multiturn(**kwargs): } ) - monkeypatch.setattr(mt_bench_utils, "generate_multiturn", fake_generate_multiturn) + monkeypatch.setattr(mt_bench_runner, "generate_multiturn", fake_generate_multiturn) cfg = RunConfig( task="mt-bench", @@ -547,10 +527,9 @@ def fake_generate_multiturn(**kwargs): generation={"n_instructions": 1}, ) - mt_bench_utils._generate_mt_bench_completions( + mt_bench_runner._generate_mt_bench_completions( cfg=cfg, questions_df=questions_df, - ignore_cache=False, ) thinking_call = captured["VLLM/Qwen/Qwen3.5-9B"] @@ -572,14 +551,14 @@ def test_run_mt_bench_forwards_strip_thinking_to_fastchat_judge(monkeypatch, tmp captured: dict[str, dict] = {} monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "load_instructions", lambda dataset, n_instructions=None: questions_df, ) monkeypatch.setattr( - mt_bench_utils, + mt_bench_runner, "_generate_mt_bench_completions", - lambda cfg, questions_df, ignore_cache: ( + lambda cfg, questions_df: ( pd.DataFrame( {"completion_turn_1": ["A1"], "completion_turn_2": ["A2"]}, index=questions_df.index, @@ -590,16 +569,16 @@ def test_run_mt_bench_forwards_strip_thinking_to_fastchat_judge(monkeypatch, tmp ), ), ) - monkeypatch.setattr(mt_bench_utils, "make_model", lambda **kwargs: object()) + monkeypatch.setattr(mt_bench_runner, "make_model", lambda **kwargs: object()) monkeypatch.setattr( - mt_bench_utils, "_finalize_mt_bench_run", lambda **kwargs: kwargs["prefs"] + mt_bench_runner, "_finalize_mt_bench_run", lambda **kwargs: kwargs["prefs"] ) def fake_judge(**kwargs): captured["judge"] = kwargs return pd.Series([0.0], dtype=float), [], [], 0 - monkeypatch.setattr(mt_bench_utils, "judge_mt_bench_pairwise_fastchat", fake_judge) + monkeypatch.setattr(mt_bench_runner, "judge_mt_bench_pairwise_fastchat", fake_judge) cfg = RunConfig( task="mt-bench", @@ -609,12 +588,7 @@ def fake_judge(**kwargs): run={"result_folder": str(tmp_path)}, ) - mt_bench_utils.run_mt_bench( - cfg, - ignore_cache=False, - res_folder=tmp_path, - result_name="mt-bench-test", - ) + mt_bench_runner.run_mt_bench_benchmark(cfg) assert captured["judge"]["strip_thinking_before_judging"] is True assert captured["judge"]["reference_categories"] == (