Skip to content
Open
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
6 changes: 6 additions & 0 deletions judgearena/artifacts/metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
from pathlib import Path
from typing import Any

from judgearena.usage import current_run_usage

METADATA_FILENAME = "run-metadata.v1.json"
METADATA_SCHEMA_VERSION = "judgearena-run-metadata/v1"
_REQUIREMENT_NAME_RE = re.compile(r"^\s*([A-Za-z0-9][A-Za-z0-9_.-]*)")
Expand Down Expand Up @@ -295,6 +297,10 @@ def write_run_metadata(
if task_definition is not None:
metadata["task_definition"] = task_definition

usage = current_run_usage()
if usage is not None:
metadata["usage"] = usage.to_dict()

git_hash = _get_git_hash(start_path=Path(__file__).resolve().parent)
if git_hash:
metadata["git_hash"] = git_hash
Expand Down
1 change: 1 addition & 0 deletions judgearena/benchmarks/mt_bench/pairwise_judging.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ def infer_pairwise_judgments_by_prompt_groups(
chat_model=judge_chat_model,
inputs=prompt_inputs,
use_tqdm=use_tqdm,
stage="judging",
)
for item_index, output, prompt_kwargs in zip(
idxs, outputs, batch_kwargs, strict=True
Expand Down
7 changes: 6 additions & 1 deletion judgearena/benchmarks/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

from judgearena.benchmarks.registry import resolve_benchmark
from judgearena.log import get_logger
from judgearena.usage import track_usage

if TYPE_CHECKING:
from judgearena.config import RunConfig
Expand All @@ -17,4 +18,8 @@ def run_benchmark(cfg: RunConfig) -> object:
"""Run a task through its registered benchmark adapter."""
resolved = resolve_benchmark(cfg.task)
logger.info("Using %s benchmark adapter for %s.", resolved.adapter.name, cfg.task)
return resolved.adapter.runner(cfg, resolved.task)
with track_usage() as usage_tracker:
try:
return resolved.adapter.runner(cfg, resolved.task)
finally:
usage_tracker.render_summary()
1 change: 1 addition & 0 deletions judgearena/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,7 @@ def annotate_battles(
inputs=inputs,
use_tqdm=use_tqdm,
return_top_logprobs=collect_top_logprobs,
stage="judging",
)
if not collect_top_logprobs:
judge_results = [InferenceResult(text=text) for text in judge_results]
Expand Down
17 changes: 12 additions & 5 deletions judgearena/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,10 @@ def generate_instructions(
completions = do_inference(
chat_model=chat_model,
inputs=inputs,
use_tqdm=use_tqdm,
# This path historically used one synchronous batch regardless of the
# progress setting; preserve that execution behavior here.
use_tqdm=False,
stage="generation",
)
df_outputs = pd.DataFrame(
data={
Expand Down Expand Up @@ -88,6 +91,7 @@ def _infer_grouped_by_temperature(
chat_model=group_model,
inputs=group_inputs,
use_tqdm=use_tqdm,
stage="generation",
)
for i, out in zip(idxs, group_outs, strict=True):
outputs[i] = out
Expand Down Expand Up @@ -152,6 +156,7 @@ def generate_multiturn(
chat_model=chat_model,
inputs=turn1_inputs,
use_tqdm=use_tqdm,
stage="generation",
)

turn2_inputs = []
Expand Down Expand Up @@ -206,6 +211,7 @@ def generate_multiturn(
chat_model=chat_model,
inputs=turn2_inputs,
use_tqdm=use_tqdm,
stage="generation",
)

return pd.DataFrame(
Expand All @@ -225,18 +231,19 @@ def generate_base(
use_tqdm: bool = False,
**engine_kwargs,
) -> pd.DataFrame:
model = make_model(model, max_tokens=max_tokens, **engine_kwargs)
chat_model = make_model(model, max_tokens=max_tokens, **engine_kwargs)

inputs = [
truncate(instruction, max_len=truncate_input_chars)
for instruction in instructions
]

completions = model.batch(
completions = do_inference(
chat_model=chat_model,
inputs=inputs,
max_tokens=max_tokens,
use_tqdm=use_tqdm,
stage="generation",
)
completions = [x.content if hasattr(x, "content") else x for x in completions]

df_outputs = pd.DataFrame(
data={
Expand Down
116 changes: 108 additions & 8 deletions judgearena/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import os
import time
import warnings
from dataclasses import dataclass
from dataclasses import dataclass, replace

from langchain_community.llms import LlamaCpp
from langchain_openai import ChatOpenAI
Expand All @@ -16,6 +16,7 @@

from judgearena.constants import VLLM_REASONING_END_STR, VLLM_REASONING_START_STR
from judgearena.log import get_logger
from judgearena.usage import RequestUsage, RunUsage, record_usage
from judgearena.utils.io import safe_parse_int

logger = get_logger(__name__)
Expand Down Expand Up @@ -447,11 +448,11 @@ async def ainvoke(self, input_item, **invoke_kwargs):

@dataclass(frozen=True)
class InferenceResult:
"""A text completion with the first generated token's top logprobs, when
the backend was asked for (and returned) them."""
"""A text completion and optional provider response details."""

text: str
first_token_top_logprobs: dict[str, float] | None = None
usage: RequestUsage | None = None


def _first_token_top_logprobs(response) -> dict[str, float] | None:
Expand All @@ -466,24 +467,97 @@ def _first_token_top_logprobs(response) -> dict[str, float] | None:
}


def _to_inference_result(response) -> InferenceResult:
def _optional_number(value, number_type):
if value is None:
return None
try:
return number_type(value)
except (TypeError, ValueError):
return None


def _request_usage(
response,
*,
stage: str,
default_model: str | None,
) -> RequestUsage:
response_metadata = getattr(response, "response_metadata", None) or {}
raw_usage = response_metadata.get("token_usage") or {}
prompt_details = raw_usage.get("prompt_tokens_details") or {}
completion_details = raw_usage.get("completion_tokens_details") or {}

input_tokens = _optional_number(raw_usage.get("prompt_tokens"), int)
output_tokens = _optional_number(raw_usage.get("completion_tokens"), int)
total_tokens = _optional_number(raw_usage.get("total_tokens"), int)
if total_tokens is None and input_tokens is not None and output_tokens is not None:
total_tokens = input_tokens + output_tokens

model = response_metadata.get("model_name") or default_model
return RequestUsage(
stage=stage,
model=str(model) if model is not None else None,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
reasoning_tokens=_optional_number(
completion_details.get("reasoning_tokens"), int
),
cached_tokens=_optional_number(prompt_details.get("cached_tokens"), int),
cost_usd=_optional_number(raw_usage.get("cost"), float),
)


def _model_name(chat_model) -> str | None:
for attribute in ("model_name", "model", "model_path", "name"):
value = getattr(chat_model, attribute, None)
if isinstance(value, str) and value:
return value
return None


def _to_inference_result(
response,
*,
stage: str,
default_model: str | None,
) -> InferenceResult:
if isinstance(response, InferenceResult):
return response
if response.usage is not None:
return response
return replace(
response,
usage=RequestUsage(stage=stage, model=default_model),
)
if hasattr(response, "content"):
return InferenceResult(
text=response.content,
first_token_top_logprobs=_first_token_top_logprobs(response),
usage=_request_usage(
response,
stage=stage,
default_model=default_model,
),
)
return InferenceResult(text=response)
return InferenceResult(
text=response,
usage=RequestUsage(stage=stage, model=default_model),
)


def do_inference(
chat_model, inputs, use_tqdm: bool = False, return_top_logprobs: bool = False
chat_model,
inputs,
use_tqdm: bool = False,
return_top_logprobs: bool = False,
*,
stage: str = "unspecified",
):
"""Run inference over *inputs*, returning a list of text completions.

With ``return_top_logprobs=True``, returns ``InferenceResult`` objects
carrying the first token's top logprobs where the backend provided them.
``stage`` labels request usage for the run-level summary.

Retries on rate-limit/server errors with exponential backoff. The async
path (``use_tqdm=True``) retries individual calls; the batch path splits
Expand Down Expand Up @@ -573,7 +647,33 @@ def batch_with_retry(batch_inputs, max_retries=5, base_delay=1.0):

# Langchain chat models return AIMessage objects, barebones models plain
# strings, and ChatVLLM with logprobs enabled InferenceResult objects.
res = [_to_inference_result(x) for x in res]
res = [
_to_inference_result(
response,
stage=stage,
default_model=_model_name(chat_model),
)
for response in res
]
request_usage = [result.usage for result in res if result.usage is not None]
record_usage(request_usage)

batch_summary = RunUsage(tuple(request_usage)).summary()
if batch_summary["requests_with_token_usage"]:
cost = batch_summary["cost_usd"]
input_tokens = batch_summary["input_tokens"]
output_tokens = batch_summary["output_tokens"]
cost_text = (
f", ${float(cost):.6f}" if cost is not None else ", cost unavailable"
)
logger.info(
"Model usage (%s): %d request(s), %s input / %s output tokens%s.",
stage,
batch_summary["requests"],
f"{int(input_tokens):,}" if input_tokens is not None else "unknown",
f"{int(output_tokens):,}" if output_tokens is not None else "unknown",
cost_text,
)
if return_top_logprobs:
return res
return [r.text for r in res]
Expand Down
Loading
Loading