Team Ai
Apppublic

dvwn/nl2sql-api

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
evaluation_mode.py205 linesDownload Raw Back to scripts
1# Path: src/scripts/evaluation_mode.py2# Evaluation script for Hugging Face SQL generation.3import json4import sqlglot5from pathlib import Path6import pandas as pd7from src.database.db_manager import get_db_connection8from src.nl2sql.hf_engine import get_models9from src.nl2sql.sql_agent import nl2sql_agent10from src.scripts.taxonomy_report import print_taxonomyReport11 12TEST_CASES_PATH = Path("src/scripts/test_cases.json")13 14def _normalize_dataframe(dataframe: pd.DataFrame) -> list:15    # Normalize dataframe to ensure accurate comparison16    """17    Standardize dataframes for Execution Accuracy (EX).18    - Converts dataframe to a list of tuples to ignore column names.19    - Rounds floating points to 4 decimal places to avoid precision mismatch.20    - Sorts the final list to ensure order-agnostic comparison.21    """22    if dataframe is None or dataframe.empty:23        return []24    25    normalized = dataframe.copy()26 27    for column in normalized.columns:28        normalized[column] = normalized[column].map(29            lambda x: round(float(x), 4)30            if pd.api.types.is_numeric_dtype(type(x)) and isinstance(x, float)31            else x32            #lambda value: round(float(value), 6)33            #if isinstance(value, (float, int))34            #else value35        )36 37    # Convert to list of tuples for order-agnostic comparison38    data_tuples = [tuple(row) for row in normalized.to_numpy()]39 40    # Sort to ensure order agnoticism41    try:42        data_tuples.sort(key=lambda x: str(x))43    except Exception as e:44        pass45        46    return data_tuples47 48# Semantic safety net49def extract_tables(sql: str) -> set:50    """51    Extract a set of all table names used in a SQL query.52    Used to catch false positives where EX passes but the model queried the wrong tables.53    """54    if not sql:55        return set()56    try:57        parsed = sqlglot.parse_one(sql, read=None)58 59        # Find all table expressions & extract names, ignore aliases60        return set(table.name.lower() for table in parsed.find_all(sqlglot.exp.Table) if table.name)61    except Exception as e:62        return set()63 64# EX: Compare generated SQL results with expected results65def calculate_ex(df_generated: pd.DataFrame, df_gold: pd.DataFrame) -> bool:66    """67    Execution Accuracy (EX): Compare generated SQL results with expected results.68    """69    if df_generated is None or df_gold is None:70        return False71 72    #if normalized_generated.shape != normalized_gold.shape:73    #        return False74    75    try:76        normalized_generated = _normalize_dataframe(df_generated)77        normalized_gold = _normalize_dataframe(df_gold)78        79        return normalized_generated == normalized_gold80 81    except Exception as error:82        print(f"EX Evaluation Error: {error}")83        return False84 85def calculate_esm(generated_sql: str, gold_sql: str) -> bool:86    """87    Exact Set Match (ESM): Compare AST structure using sqlglot.88    - Ignores formatting, capitalization, and minor syntactic sugar.89    """90    if not generated_sql or not gold_sql:91        return False92    93    try:94        # Parse both SQL queries into expressions95        generated_exp = sqlglot.parse_one(generated_sql, read=None)96        gold_exp = sqlglot.parse_one(gold_sql, read=None)97 98        # Compare the expressions for structural equivalence99        return generated_exp == gold_exp100    except Exception as error:101        print(f"ESM Evaluation Error: {error}")102        return False103 104def run_evaluation(model_id: str):105    print(f"\nRunning SQL evaluation for model: {model_id}")106    print("\n" + "-" *50)107 108    if not TEST_CASES_PATH.exists():109        print(f"Error: Could not find test cases at {TEST_CASES_PATH}")110        return111    112    with TEST_CASES_PATH.open("r", encoding="utf-8") as handle:113        test_cases = json.load(handle)114 115    results = []116    ex_count = 0117    esm_count = 0118 119    print(f"Running evaluation on {len(test_cases)} test cases...\n")120 121    for case in test_cases:122        id = case.get("id")123        question = case.get("question")124        gold_sql = case.get("gold_sql")125        taxonomy = case.get("taxonomy", "Unknown")126        print(f"Testing ID {id}: {question[:40]}...")127 128        # Implement agent to handle RAG retrieval and SQL generation129        agent_response = nl2sql_agent(user_question=question, model_id=model_id)130        generated_sql = agent_response.get("query", "")131 132        # ESM Evaluation133        esm_result = calculate_esm(generated_sql, gold_sql)134        if esm_result:135            esm_count += 1136 137        # EX Evaluation138        ex_result = False139        connection = get_db_connection()140        if connection is None:141            raise RuntimeError("Unable to connect to the SQLite database.")142 143        try:144            df_generated = pd.read_sql_query(generated_sql, connection)145            df_gold = pd.read_sql_query(gold_sql, connection)146 147            # Trap the False Positive (empty set): weak test case148            if df_gold.empty:149                print(f"[!]WARNING: Gold SQL for ID {id} returned an emoty response.")150 151            ex_result = calculate_ex(df_generated, df_gold)152 153            # Semantic safety net check154            if ex_result:155                gen_tables = extract_tables(generated_sql)156                gold_tables = extract_tables(gold_sql)157 158                if gen_tables != gold_tables:159                    print(f"[X] FALSE POSITIVE (ID{id}): Data matched, tables not")160                    print(f"\nGenerated SQL tables: {gen_tables} | Gold SQL tables: {gold_tables}")161                    ex_result = False162        163            if ex_result:164                ex_count += 1165        except Exception as error:166            print(f"Error executing SQL for ID {id}: {error}")167        finally:168            connection.close()169 170        results.append({171            "id": id,172            "question": question,173            "taxonomy": taxonomy,174            "ex_pass": ex_result,175            "esm_pass": esm_result,176            "generated_sql": generated_sql,177            "gold_sql": gold_sql178        })179    180    # Summary Statistics181    total = len(test_cases)182    ex_accuracy = (ex_count / total) * 100 if total > 0 else 0183    esm_accuracy = (esm_count / total) * 100 if total > 0 else 0184 185    print("\nEVALUATION SUMMARY")186    print("-" * 40)187    print(f"Model Evaluated: {model_id}")188    print(f"Total Test Cases: {total}")189    print(f"Execution Accuracy (EX): {ex_accuracy:.2f}% ({ex_count}/{total})")190    print(f"Exact Set Match (ESM): {esm_accuracy:.2f}% ({esm_count}/{total})")191 192    safe_model_name = model_id.replace("/", "_").replace(":", "_")193    output_file = Path(f"sql_eval_{safe_model_name}.json")194    with output_file.open("w", encoding="utf-8") as handle:195        json.dump(results, handle, indent=4)196 197    print_taxonomyReport(results)198 199if __name__ == "__main__":200    from dotenv import load_dotenv, find_dotenv201    load_dotenv(find_dotenv())202 203    models_to_test = get_models()204    for model in models_to_test:205        run_evaluation(model)