Spaces:
Running on Zero
Running on Zero
| import re | |
| import pandas as pd | |
| import sqlglot | |
| from datasets import load_dataset | |
| from config import DATASET_ID, GENERATION_EVAL_SAMPLES, OUTPUT_DIR, SEED | |
| from inference import generate_sql, load_model | |
| def normalize_sql(sql): | |
| """Normalize SQL for a lightweight exact-match comparison.""" | |
| sql = sql.strip().rstrip(";") | |
| sql = re.sub(r"\s+", " ", sql) | |
| return sql.lower() | |
| def is_valid_sql(sql): | |
| """Return True when sqlglot can parse the generated SQL.""" | |
| try: | |
| sqlglot.parse_one(sql) | |
| return True | |
| except Exception: | |
| return False | |
| def main(): | |
| test_dataset = load_dataset(DATASET_ID, split="test").shuffle(seed=SEED) | |
| test_dataset = test_dataset.select( | |
| range(min(GENERATION_EVAL_SAMPLES, len(test_dataset))) | |
| ) | |
| model, tokenizer = load_model(OUTPUT_DIR) | |
| results = [] | |
| for i, example in enumerate(test_dataset): | |
| prediction = generate_sql( | |
| model, | |
| tokenizer, | |
| example["sql_context"], | |
| example["sql_prompt"], | |
| ) | |
| target = example["sql"].strip() | |
| exact_match = normalize_sql(prediction) == normalize_sql(target) | |
| syntax_valid = is_valid_sql(prediction) | |
| results.append( | |
| { | |
| "index": i, | |
| "question": example["sql_prompt"], | |
| "target_sql": target, | |
| "predicted_sql": prediction, | |
| "exact_match": exact_match, | |
| "syntax_valid": syntax_valid, | |
| } | |
| ) | |
| print( | |
| f"[{i + 1}/{len(test_dataset)}] " | |
| f"Exact={exact_match} | Valid={syntax_valid}" | |
| ) | |
| df = pd.DataFrame(results) | |
| df.to_csv("evaluation_results.csv", index=False) | |
| print("\nEvaluation summary") | |
| print("------------------") | |
| print(f"Examples: {len(df)}") | |
| print(f"Exact-match: {df['exact_match'].mean():.3f}") | |
| print(f"SQL syntax rate: {df['syntax_valid'].mean():.3f}") | |
| print("\nSaved detailed results to evaluation_results.csv") | |
| if __name__ == "__main__": | |
| main() |