"""
LLMSQL Evaluation Module
=========================
Provides the `evaluate()` function to benchmark Text-to-SQL model outputs
on the LLMSQL benchmark.
See the documentation for full usage details.
"""
from datetime import datetime, timezone
from typing import Any
import uuid
from rich.progress import track
from llmsql.config.config import (
DEFAULT_LLMSQL_VERSION,
get_repo_id,
)
from llmsql.loggers.logging_config import log
from llmsql.utils.evaluation_utils import (
connect_sqlite,
evaluate_sample,
resolve_prediction_coverage,
)
from llmsql.utils.inference_utils import _maybe_download, resolve_workdir_path
from llmsql.utils.leaderboard_utils import (
build_leaderboard_record,
write_leaderboard_yaml,
)
from llmsql.utils.rich_utils import log_mismatch, print_summary
from llmsql.utils.utils import load_jsonl, load_jsonl_dict_by_key, save_json_report
[docs]
def evaluate(
outputs: str | list[dict[int, str | int]],
*,
version: str = DEFAULT_LLMSQL_VERSION,
workdir_path: str | None = None,
save_report: str | None = None,
show_mismatches: bool = True,
max_mismatches: int = 5,
model_name: str | None = None,
save_leaderboard_yaml: str | None = None,
run_metadata: dict[str, Any] | None = None,
) -> dict:
"""
Evaluate predicted SQL queries against the LLMSQL benchmark.
Args:
version: LLMSQL version
outputs: Either a JSONL file path or a list of dicts.
workdir_path: Directory to store downloaded benchmark files. If omitted, a
temporary directory is created automatically.
save_report: Optional manual save path. If None → auto-generated.
show_mismatches: Print mismatches while evaluating.
max_mismatches: Max mismatches to print.
model_name: Name of the evaluated model (e.g. ``Qwen/Qwen3-0.6B``).
Stored in the JSON report and in the leaderboard YAML.
save_leaderboard_yaml: Optional path to additionally save the
results in the leaderboard ``run.yaml`` format (see the
``leaderboard/`` folder). If None, no YAML is written.
run_metadata: Optional dict deep-merged into the leaderboard YAML to
fill fields that cannot be detected automatically, e.g.
``{"type": "open-source", "inference": {"backend": "vllm",
"arguments": {"num_fewshots": 5}}}``.
Returns:
dict: Metrics and mismatches, plus coverage counters ``expected``,
``answered``, ``missing`` and ``duplicates``. ``accuracy`` is
computed over the predictions, ``accuracy_over_benchmark`` counts
unanswered questions as wrong.
Raises:
ValueError: if a prediction references a ``question_id`` that is not
part of the benchmark.
"""
# Determine input type
input_mode = "jsonl_path" if isinstance(outputs, str) else "dict_list"
workdir = resolve_workdir_path(workdir_path)
repo_id = get_repo_id(version)
questions_path = _maybe_download(repo_id, "questions.jsonl", workdir)
db_path = _maybe_download(repo_id, "sqlite_tables.db", workdir)
# --- Load benchmark questions ---
questions = load_jsonl_dict_by_key(questions_path, key="question_id")
# --- Load predictions (path or list) ---
if isinstance(outputs, str):
outputs_list = load_jsonl(outputs)
elif isinstance(outputs, list):
outputs_list = outputs
else:
raise TypeError(
"outputs must be file path or list of dicts in format {'question_id': int, 'completion': str}"
)
# --- Drop duplicates / measure coverage against the benchmark ---
# Scoring the predictions file directly made partial runs (crash, --limit,
# filtered file) look like full results, and counted duplicate ids twice.
outputs_list, coverage = resolve_prediction_coverage(outputs_list, questions)
if coverage["duplicates"]:
log.warning(
f"{coverage['duplicates']} duplicate question_id(s) in the predictions "
f"file; keeping only the first prediction for each."
)
if coverage["missing"]:
log.warning(
f"Coverage {coverage['answered']}/{coverage['expected']} question(s): "
f"{coverage['missing']} benchmark question(s) have no prediction, so "
f"accuracy is computed over {coverage['answered']} answer(s) only."
)
# --- Connect to DB ---
conn = connect_sqlite(db_path)
# --- Evaluation loop ---
metrics = {
"total": 0,
"matches": 0,
"exact_string_matches": 0,
"pred_none": 0,
"gold_none": 0,
"sql_errors": 0,
}
mismatches: list[dict] = []
for item in track(outputs_list, description="Evaluating"):
metrics["total"] += 1
is_match, mismatch_info, m = evaluate_sample(item, questions, conn)
metrics["matches"] += is_match
metrics["pred_none"] += m["pred_none"]
metrics["gold_none"] += m["gold_none"]
metrics["sql_errors"] += m["sql_error"]
metrics["exact_string_matches"] += m["exact_string_match"]
if mismatch_info:
mismatches.append(mismatch_info)
if show_mismatches and len(mismatches) <= max_mismatches:
log_mismatch(**mismatch_info)
print_summary(
metrics["total"],
metrics["matches"],
metrics["pred_none"],
metrics["gold_none"],
metrics["sql_errors"],
metrics["exact_string_matches"],
coverage,
)
# --- Build report structure ---
accuracy = metrics["matches"] / metrics["total"] if metrics["total"] else 0.0
report = {
"model_name": model_name,
"version": version,
**metrics,
**coverage,
"accuracy": accuracy,
# Missing answers count as wrong, so a partial run cannot silently
# report a full-looking accuracy.
"accuracy_over_benchmark": (
metrics["matches"] / coverage["expected"] if coverage["expected"] else 0.0
),
"exact_string_match_accuracy": (
metrics["exact_string_matches"] / metrics["total"]
if metrics["total"]
else 0.0
),
"mismatches": mismatches,
"timestamp": datetime.now(timezone.utc).isoformat(),
"input_mode": input_mode,
}
# --- Auto-generate report filename (if not provided) ---
if save_report is None:
save_report = f"evaluation_results_{uuid.uuid4()}.json"
save_json_report(save_report, report)
if save_leaderboard_yaml is not None:
# The leaderboard ranks models on the whole benchmark: unanswered
# questions count as wrong, so a partial run cannot rank too high.
record = build_leaderboard_record(
accuracy=report["accuracy_over_benchmark"],
total=coverage["expected"],
version=version,
model_name=model_name,
answers_path=outputs if isinstance(outputs, str) else None,
run_metadata=run_metadata,
)
write_leaderboard_yaml(save_leaderboard_yaml, record)
conn.close()
return report