dvwn/nl2sql-api
0
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)