Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 19 additions & 2 deletions judgearena/benchmarks/elo/calibration.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,9 @@
from judgearena.arenas_utils import extract_turn_text
from judgearena.benchmarks.elo.rating import winner_to_pref
from judgearena.evaluate import judge_and_parse_prefs
from judgearena.inference import JudgementInferenceCache
from judgearena.log import get_logger
from judgearena.models import make_model
from judgearena.models import prepare_model
from judgearena.prompts.parsing import PairScore
from judgearena.prompts.registry import ResolvedJudgePrompt

Expand Down Expand Up @@ -62,6 +63,8 @@ def calibrate_pairscore_temperature(
prompt: ResolvedJudgePrompt,
truncate_input_chars: int | None,
default_temperature: float,
arena: str,
inference_cache: JudgementInferenceCache | None = None,
) -> float | None:
"""Judge sampled human battles and return a fitted PairScore temperature."""
if not enabled:
Expand Down Expand Up @@ -103,7 +106,12 @@ def calibrate_pairscore_temperature(
for index in calibration_battles.index
]

calibration_judge = make_model(model=judge_model, **dict(judge_model_kwargs))
source_rows = source_battles.loc[calibration_battles.index]
calibration_judge = prepare_model(
model=judge_model,
cache=inference_cache,
**dict(judge_model_kwargs),
)
annotations, _, _ = judge_and_parse_prefs(
judge_chat_model=calibration_judge,
instructions=instructions,
Expand All @@ -115,6 +123,15 @@ def calibrate_pairscore_temperature(
prompt_preset=prompt.preset_name,
parse=prompt.parser,
truncate_input_chars=truncate_input_chars,
cache_row_metadata=[
{
"instruction_id": f"{arena}:{row.question_id}",
"model_a": row.model_a,
"model_b": row.model_b,
"orientation": "direct",
}
for row in source_rows.itertuples()
],
)

score_differences: list[float] = []
Expand Down
96 changes: 39 additions & 57 deletions judgearena/benchmarks/elo/runner.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,4 @@
import hashlib
from datetime import UTC, datetime
from functools import partial
from pathlib import Path
from typing import TYPE_CHECKING

Expand All @@ -19,10 +17,13 @@
from judgearena.benchmarks.elo.rating import (
arena_anchor_battles,
prefs_to_battle_results,
sampling_cache_token,
select_seeded_random_arena_battles,
)
from judgearena.benchmarks.execution import build_generation_kwargs
from judgearena.benchmarks.execution import (
build_completion_cache,
build_generation_kwargs,
build_judgement_cache,
)
from judgearena.benchmarks.scoring import build_metrics, calculate_metrics
from judgearena.datasets import load_battles
from judgearena.evaluate import (
Expand All @@ -33,10 +34,9 @@
)
from judgearena.generate import generate_instructions
from judgearena.log import get_logger
from judgearena.models import build_default_judge_model_kwargs, make_model
from judgearena.models import build_default_judge_model_kwargs, prepare_model
from judgearena.reports import EloReport
from judgearena.tasks.schema import EloProtocol, ResolvedTaskSpec
from judgearena.utils import cache_function_dataframe

if TYPE_CHECKING:
from judgearena.config import RunConfig
Expand Down Expand Up @@ -108,6 +108,11 @@ def run_elo(cfg: "RunConfig", task: ResolvedTaskSpec | None = None) -> dict:
extract_turn_text(row["conversation_a"][0])
for _, row in df_battles.iterrows()
],
index=(
arena + ":" + df_battles["question_id"].astype(str)
if "question_id" in df_battles
else df_battles.index.astype(str)
),
name="instruction",
)
logger.debug("First instruction:\n%s", instructions.iloc[0][:300])
Expand All @@ -121,48 +126,13 @@ def run_elo(cfg: "RunConfig", task: ResolvedTaskSpec | None = None) -> dict:
# dropped battle_thinking_token_budget).
extra_kwargs = build_generation_kwargs(cfg, cfg.model.name, role="A")
use_tqdm = False
gen_fun = partial(
generate_instructions,
completions_df = generate_instructions(
instructions=instructions,
model=cfg.model.name,
truncate_input_chars=cfg.generation.truncate_all_input_chars,
use_tqdm=use_tqdm,
inference_cache=build_completion_cache(cfg),
**extra_kwargs,
)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nice that we get rid of this!


def replace_slash(s: str) -> str:
return s.replace("/", "_")

languages_str = (
"-".join(sorted(selected_languages)) if selected_languages else "all"
)
extra_kwargs_str = (
"_".join(f"{k}={v}" for k, v in sorted(extra_kwargs.items()))
if extra_kwargs
else ""
)
cache_token = sampling_cache_token(
sampling_metadata,
n_instructions=cfg.generation.n_instructions,
n_instructions_per_language=cfg.elo.n_instructions_per_language,
)
cache_suffix = (
f"{arena}_{replace_slash(cfg.model.name)}_"
f"{cache_token}_"
f"{languages_str}_{cfg.generation.truncate_all_input_chars}_{extra_kwargs['max_tokens']}"
+ (f"_{extra_kwargs_str}" if extra_kwargs_str else "")
)
if len(cache_suffix) > 100:
cache_hash = hashlib.sha256(cache_suffix.encode()).hexdigest()[:16]
logger.debug(
"Cache suffix too long (%d chars), using hash: %s (full: %s)",
len(cache_suffix),
cache_hash,
cache_suffix,
)
cache_suffix = cache_hash
completions_df = cache_function_dataframe(
lambda: gen_fun(instructions=instructions, model=cfg.model.name),
ignore_cache=cfg.run.ignore_cache,
cache_name=f"elo/{cache_suffix}",
).set_index("instruction_index")
completions = completions_df.loc[:, "completion"]

Expand Down Expand Up @@ -210,8 +180,9 @@ def replace_slash(s: str) -> str:
)

def run_judge() -> pd.DataFrame:
judge_chat_model = make_model(
judge_chat_model = prepare_model(
model=cfg.judge.model,
cache=build_judgement_cache(cfg),
**judge_extra_kwargs,
)
annotations, annotations_reversed, prefs = judge_and_parse_prefs(
Expand All @@ -227,6 +198,25 @@ def run_judge() -> pd.DataFrame:
parse=resolved_prompt.parser,
truncate_input_chars=cfg.generation.truncate_judge_input_chars,
use_tqdm=use_tqdm,
cache_row_metadata=[
{
"instruction_id": str(instructions.index[index]),
"model_a": (
cfg.model.name
if our_model_is_position_a[index]
else opponent_models[index]
),
"model_b": (
opponent_models[index]
if our_model_is_position_a[index]
else cfg.model.name
),
"orientation": (
"direct" if our_model_is_position_a[index] else "reversed"
),
}
for index in range(n)
],
)
if annotations_reversed is None:
row_annotations = list(annotations)
Expand Down Expand Up @@ -259,17 +249,7 @@ def run_judge() -> pd.DataFrame:
)
return frame

# Stripping reasoning traces changes the judged text but not the cached
# completions, so it must be part of the judge cache key. Only append when
# enabled so prior (non-stripped) runs keep their existing cache hashes.
judge_cache_suffix = f"judge_{cache_suffix}"
if cfg.judge.strip_thinking_before_judging:
judge_cache_suffix += "_stripthinking"
df_judge = cache_function_dataframe(
run_judge,
ignore_cache=cfg.run.ignore_cache,
cache_name=f"elo/{judge_cache_suffix}",
)
df_judge = run_judge()

# Restore position arrays and prefs from cache (in case loaded from disk)
use_model_a_as_opponent = df_judge["use_model_a_as_opponent"].to_numpy()
Expand Down Expand Up @@ -308,6 +288,8 @@ def run_judge() -> pd.DataFrame:
prompt=resolved_prompt,
truncate_input_chars=cfg.generation.truncate_judge_input_chars,
default_temperature=cfg.elo.soft_elo_temperature,
arena=arena,
inference_cache=build_judgement_cache(cfg),
)

# Reparse cached responses at this run's temperature, keeping the selected
Expand Down
19 changes: 17 additions & 2 deletions judgearena/benchmarks/execution.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,14 @@

from __future__ import annotations

from pathlib import Path
from typing import TYPE_CHECKING, Literal

from judgearena.inference import CompletionInferenceCache, JudgementInferenceCache
from judgearena.models import (
build_default_judge_model_kwargs,
is_thinking_model,
make_model,
prepare_model,
)

if TYPE_CHECKING:
Expand Down Expand Up @@ -35,8 +37,9 @@ def build_generation_kwargs(

def build_judge(cfg: RunConfig):
"""Construct the configured judge consistently across benchmark runners."""
return make_model(
return prepare_model(
model=cfg.judge.model,
cache=build_judgement_cache(cfg),
**build_default_judge_model_kwargs(
cfg.judge.model,
cfg.model.engine_kwargs,
Expand All @@ -45,3 +48,15 @@ def build_judge(cfg: RunConfig):
),
),
)


def build_completion_cache(cfg: RunConfig) -> CompletionInferenceCache | None:
if cfg.run.store_root is None:
return None
return CompletionInferenceCache(Path(cfg.run.store_root), cfg.task)


def build_judgement_cache(cfg: RunConfig) -> JudgementInferenceCache | None:
if cfg.run.store_root is None:
return None
return JudgementInferenceCache(Path(cfg.run.store_root), cfg.task)
9 changes: 9 additions & 0 deletions judgearena/benchmarks/meta_eval/annotate.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,15 @@ def annotate_sample(
parse=parser,
truncate_input_chars=cfg.generation.truncate_judge_input_chars,
use_tqdm=cfg.run.use_tqdm,
cache_row_metadata=[
{
"instruction_id": battle["battle_id"],
"model_a": battle["model_a"],
"model_b": battle["model_b"],
"orientation": "direct",
}
for _, battle in df_sample.iterrows()
],
)

n_battles = len(df_sample)
Expand Down
4 changes: 4 additions & 0 deletions judgearena/benchmarks/mt_bench/fastchat_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -362,6 +362,8 @@ def judge_mt_bench_pairwise_fastchat(
items=items,
use_tqdm=use_tqdm,
swap_answers=False,
model_a=model_a,
model_b=model_b,
)

g2_judgments: list[str] | None = None
Expand All @@ -371,6 +373,8 @@ def judge_mt_bench_pairwise_fastchat(
items=items,
use_tqdm=use_tqdm,
swap_answers=True,
model_a=model_a,
model_b=model_b,
)

annotations: list[dict[str, Any]] = []
Expand Down
13 changes: 13 additions & 0 deletions judgearena/benchmarks/mt_bench/pairwise_judging.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,8 @@ def infer_pairwise_judgments_by_prompt_groups(
items: list[MTBenchJudgeItem],
use_tqdm: bool,
swap_answers: bool,
model_a: str,
model_b: str,
) -> tuple[list[str], list[dict[str, str]]]:
judgments: list[str] = [""] * len(items)
used_prompt_kwargs: list[dict[str, str]] = [{} for _ in items]
Expand All @@ -97,6 +99,17 @@ def infer_pairwise_judgments_by_prompt_groups(
inputs=prompt_inputs,
use_tqdm=use_tqdm,
stage="judging",
cache_row_metadata=[
{
"instruction_id": (
f"{items[item_index].question_id}:turn-{items[item_index].turn}"
),
"model_a": model_b if swap_answers else model_a,
"model_b": model_a if swap_answers else model_b,
"orientation": "reversed" if swap_answers else "direct",
}
for item_index in idxs
],
)
for item_index, output, prompt_kwargs in zip(
idxs, outputs, batch_kwargs, strict=True
Expand Down
4 changes: 4 additions & 0 deletions judgearena/benchmarks/mt_bench/preset_judging.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,8 @@ def judge_mt_bench_with_preset(
items=items,
use_tqdm=use_tqdm,
swap_answers=False,
model_a=model_a,
model_b=model_b,
)

annotations: list[dict[str, Any]] = []
Expand Down Expand Up @@ -246,6 +248,8 @@ def _append_results(
items=items,
use_tqdm=use_tqdm,
swap_answers=True,
model_a=model_a,
model_b=model_b,
)
)
_append_results(swapped_judgments, swapped_prompt_kwargs, swapped=True)
Expand Down
Loading