Team Ai
Apppublic

ritvik360/nl2sql-bench

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
grader.py215 linesDownload Raw Back to server
1"""2nl2sql-bench/server/grader.py3==============================4Deterministic, programmatic reward grader.5 6No LLM-as-judge. Every reward is computed by comparing the agent's SQL7execution results against a ground-truth result set.8 9Reward decomposition (sums to 1.0 for a perfect first-attempt answer):10  +0.10  syntax_ok        — query runs without SQLite error11  +0.20  columns_match    — returned column names match ground truth exactly12  +0.20  row_count_match  — number of returned rows matches13  +0.50  exact_match      — full result set equals ground truth (order-aware14                            for ORDER BY queries, order-agnostic otherwise)15 16Step penalty:17  -0.05 per step beyond the first (encourages solving in fewer steps),18  clamped so the minimum is always 0.0.19 20All rewards are floats in [0.0, 1.0].21"""22 23from __future__ import annotations24 25import sqlite326from typing import Any, Dict, List, Optional, Tuple27 28 29# ── Result normalisation ───────────────────────────────────────────────────30 31def _normalise_value(v: Any) -> Any:32    """Round floats for comparison so 1.2300000001 == 1.23."""33    if isinstance(v, float):34        return round(v, 4)35    if isinstance(v, str):36        return v.strip()37    return v38 39 40def _normalise_row(row: Dict[str, Any]) -> Dict[str, Any]:41    return {k: _normalise_value(v) for k, v in row.items()}42 43 44def _normalise_rows(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]:45    return [_normalise_row(r) for r in rows]46 47 48# ── SQL execution ──────────────────────────────────────────────────────────49 50def execute_query(51    conn: sqlite3.Connection,52    query: str,53    max_rows: int = 200,54) -> Tuple[Optional[List[Dict[str, Any]]], Optional[str]]:55    """56    Execute a SQL query safely.57 58    Returns (rows, error_string).59    rows is None on error.60    """61    query = query.strip().rstrip(";")62    if not query:63        return None, "Empty query."64 65    # Block write operations — the environment is read-only from the agent's view.66    forbidden = ("insert", "update", "delete", "drop", "alter",67                 "create", "replace", "truncate", "pragma")68    first_word = query.split()[0].lower() if query.split() else ""69    if first_word in forbidden:70        return None, (71            f"Operation '{first_word.upper()}' is not allowed. "72            "Only SELECT queries are permitted."73        )74 75    try:76        cur = conn.execute(query)77        cols = [d[0] for d in cur.description] if cur.description else []78        rows = [dict(zip(cols, row)) for row in cur.fetchmany(max_rows)]79        return rows, None80    except sqlite3.Error as exc:81        return None, str(exc)82 83 84# ── Grading logic ──────────────────────────────────────────────────────────85 86class GradeResult:87    __slots__ = (88        "reward", "syntax_ok", "columns_match",89        "row_count_match", "exact_match", "step_penalty",90        "breakdown",91    )92 93    def __init__(94        self,95        reward: float,96        syntax_ok: bool,97        columns_match: bool,98        row_count_match: bool,99        exact_match: bool,100        step_penalty: float,101    ) -> None:102        self.reward          = reward103        self.syntax_ok       = syntax_ok104        self.columns_match   = columns_match105        self.row_count_match = row_count_match106        self.exact_match     = exact_match107        self.step_penalty    = step_penalty108        self.breakdown = {109            "syntax_ok":       0.10 if syntax_ok else 0.0,110            "columns_match":   0.20 if (syntax_ok and columns_match) else 0.0,111            "row_count_match": 0.20 if (syntax_ok and row_count_match) else 0.0,112            "exact_match":     0.50 if (syntax_ok and exact_match) else 0.0,113            "step_penalty":    -step_penalty,114        }115 116    def __repr__(self) -> str:  # pragma: no cover117        return (118            f"GradeResult(reward={self.reward:.3f}, "119            f"exact={self.exact_match}, cols={self.columns_match}, "120            f"rows={self.row_count_match}, syntax={self.syntax_ok})"121        )122 123 124def grade(125    actual_rows: Optional[List[Dict[str, Any]]],126    ground_truth_rows: List[Dict[str, Any]],127    error: Optional[str],128    step: int,129    order_sensitive: bool = False,130) -> GradeResult:131    """132    Grade the agent's query result against ground truth.133 134    Parameters135    ----------136    actual_rows       : Rows returned by the agent's query (None on error).137    ground_truth_rows : Expected rows (pre-computed at task load time).138    error             : SQLite error string (None if query ran successfully).139    step              : Current step number (1-indexed) for penalty calculation.140    order_sensitive   : If True, row order matters (queries with ORDER BY).141    """142    # ── Syntax ──────────────────────────────────────────────────────────143    syntax_ok = error is None and actual_rows is not None144 145    if not syntax_ok:146        return GradeResult(147            reward=0.0,148            syntax_ok=False,149            columns_match=False,150            row_count_match=False,151            exact_match=False,152            step_penalty=0.0,153        )154 155    gt_norm   = _normalise_rows(ground_truth_rows)156    act_norm  = _normalise_rows(actual_rows)157 158    gt_cols   = set(gt_norm[0].keys()) if gt_norm else set()159    act_cols  = set(act_norm[0].keys()) if act_norm else set()160    columns_match   = act_cols == gt_cols161    row_count_match = len(act_norm) == len(gt_norm)162 163    # Exact match: if order matters, compare list; otherwise compare sorted sets164    if columns_match and row_count_match:165        if order_sensitive:166            exact_match = act_norm == gt_norm167        else:168            # Sort rows by their string representation for order-agnostic compare169            def _sort_key(r: Dict) -> str:170                return str(sorted(r.items()))171            exact_match = (172                sorted(act_norm, key=_sort_key) == sorted(gt_norm, key=_sort_key)173            )174    else:175        exact_match = False176 177    # ── Score assembly ────────────────────────────────────────────────178    raw = (179        0.10                             # syntax180        + (0.20 if columns_match else 0.0)181        + (0.20 if row_count_match else 0.0)182        + (0.50 if exact_match else 0.0)183    )184 185    penalty     = max(0.0, step - 1) * 0.05186    reward      = float(max(0.0, min(1.0, raw - penalty)))187 188    return GradeResult(189        reward=reward,190        syntax_ok=syntax_ok,191        columns_match=columns_match,192        row_count_match=row_count_match,193        exact_match=exact_match,194        step_penalty=penalty,195    )196 197 198# ── Convenience: pre-compute ground truth rows ─────────────────────────────199 200def compute_ground_truth(201    conn: sqlite3.Connection,202    sql: str,203) -> List[Dict[str, Any]]:204    """Execute the ground-truth SQL and return normalised rows."""205    rows, error = execute_query(conn, sql)206    if error or rows is None:207        raise ValueError(f"Ground-truth SQL failed: {error}\nSQL: {sql}")208    return _normalise_rows(rows)209 210 211def has_order_by(sql: str) -> bool:212    """Heuristic: does the top-level query have an ORDER BY?"""213    # Simple check sufficient for our controlled task SQL214    return "ORDER BY" in sql.upper()215